HuggingEnvs/watercolour-reference-pool
Viewer • Updated • 178 • 757 • 1
How to use FWKV/FWKV-Image-Hibiscus-LoRA with PEFT:
Task type is invalid.
Configuration Parsing Warning:In adapter_config.json: "peft.base_model_name_or_path" must be a string
Configuration Parsing Warning:In adapter_config.json: "peft.task_type" must be a string
Small fun test of FWKV-Image attempt to LoRA. Trained on HuggingEnvs/watercolour-reference-pool
Code used for training:
"""
Run:
pip install -U torch diffusers transformers accelerate peft datasets pillow
python train_fwkv_lora.py
Requires modeling_fwkv_vision.py to be importable (same directory).
"""
import os
import math
import random
import torch
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
from PIL import Image
from transformers import CLIPTokenizer
from datasets import load_dataset
from peft import LoraConfig, get_peft_model
from modeling_fwkv_vision import FWKVVisionModel, FWKVVisionConfig, VAE_SCALING
# Config
BASE_MODEL = "FWKV/FWKV-Image"
DATASET_NAME = "HuggingEnvs/watercolour-reference-pool"
IMAGE_SIZE = 256
STYLE_SUFFIX = ", watercolour wash, soft bleed" # appended to every caption
OUTPUT_DIR = "fwkv-hibiscus-lora"
LORA_R = 8
LORA_ALPHA = 16
LORA_DROPOUT = 0.05
LORA_TARGET_MODULES = ["cross_q", "cross_k", "cross_v", "cross_out"]
BATCH_SIZE = 4
LEARNING_RATE = 1e-4
EPOCHS = 2000 # small dataset -> many epochs, watch for overfitting
GRAD_CLIP = 1.0
SEED = 42
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
# Base model weights are float32; keep everything in float32 to avoid dtype
# mismatches between cast inputs and uncast layers. If you want speed via
# bf16, cast the whole model with model.dit.to(torch.bfloat16) as well.
DTYPE = torch.float32
# Data
class LatentCaptionDataset(Dataset):
"""Pre-encodes every image to a VAE latent once, keeps captions as text."""
def __init__(self, model: FWKVVisionModel, tokenizer: CLIPTokenizer, device: str):
ds = load_dataset(DATASET_NAME, split="train")
ds = ds.filter(lambda r: r["tier"] == "love")
print(f"Loaded {len(ds)} love-tier examples.")
self.captions = [row["subject"] + STYLE_SUFFIX for row in ds]
self.latents = []
vae = model.vae.to(device).eval()
with torch.no_grad():
for row in ds:
img: Image.Image = row["image"].convert("RGB").resize(
(IMAGE_SIZE, IMAGE_SIZE), Image.LANCZOS
)
arr = torch.tensor(
list(img.getdata()), dtype=torch.float32
).view(IMAGE_SIZE, IMAGE_SIZE, 3).permute(2, 0, 1)
arr = (arr / 127.5) - 1.0 # [-1, 1]
arr = arr.unsqueeze(0).to(device)
posterior = vae.encode(arr).latent_dist
latent = posterior.mean * VAE_SCALING # deterministic, matches base training
self.latents.append(latent.squeeze(0).cpu())
vae.to("cpu")
def __len__(self):
return len(self.latents)
def __getitem__(self, idx):
return self.latents[idx], self.captions[idx]
def collate(batch):
latents = torch.stack([b[0] for b in batch])
captions = [b[1] for b in batch]
return latents, captions
# Training
def main():
torch.manual_seed(SEED)
random.seed(SEED)
print(f"Loading base model on {DEVICE} ...")
model = FWKVVisionModel.from_pretrained(BASE_MODEL, trust_remote_code=True)
model.to(DEVICE)
tokenizer = CLIPTokenizer.from_pretrained(model.config.clip_id)
# Freeze everything, then LoRA-wrap only the DiT's cross-attention.
for p in model.parameters():
p.requires_grad_(False)
lora_config = LoraConfig(
r=LORA_R,
lora_alpha=LORA_ALPHA,
target_modules=LORA_TARGET_MODULES,
lora_dropout=LORA_DROPOUT,
bias="none",
)
model.dit = get_peft_model(model.dit, lora_config)
model.dit.print_trainable_parameters()
# Text encoder and VAE stay frozen and in eval mode throughout.
model.text_encoder.eval()
model.vae.eval()
print("Encoding images to latents (one-time)...")
dataset = LatentCaptionDataset(model, tokenizer, DEVICE)
loader = DataLoader(
dataset, batch_size=BATCH_SIZE, shuffle=True, collate_fn=collate, drop_last=True
)
trainable_params = [p for p in model.dit.parameters() if p.requires_grad]
optimizer = torch.optim.AdamW(trainable_params, lr=LEARNING_RATE, weight_decay=0.0)
total_steps = EPOCHS * max(1, len(loader))
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=total_steps)
model.dit.train()
step = 0
for epoch in range(EPOCHS):
epoch_loss = 0.0
for latents, captions in loader:
latents = latents.to(DEVICE, dtype=DTYPE)
with torch.no_grad():
text_tokens, pooled_text = model.encode_text(tokenizer, captions, DEVICE)
text_tokens = text_tokens.to(DTYPE)
pooled_text = pooled_text.to(DTYPE)
x0 = torch.randn_like(latents)
t = torch.rand(latents.shape[0], device=DEVICE, dtype=DTYPE)
t_broadcast = t.view(-1, 1, 1, 1)
xt = (1 - t_broadcast) * x0 + t_broadcast * latents
target = latents - x0
pred = model.dit(xt, t, text_tokens, pooled_text)
loss = F.mse_loss(pred.float(), target.float())
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(trainable_params, GRAD_CLIP)
optimizer.step()
scheduler.step()
epoch_loss += loss.item()
step += 1
avg = epoch_loss / max(1, len(loader))
if epoch % 5 == 0 or epoch == EPOCHS - 1:
print(f"epoch {epoch:4d} step {step:6d} loss {avg:.4f}")
os.makedirs(OUTPUT_DIR, exist_ok=True)
model.dit.save_pretrained(OUTPUT_DIR)
print(f"Saved LoRA adapter to {OUTPUT_DIR}/")
if __name__ == "__main__":
main()
Base model
FWKV/FWKV-Image