Text-to-Image
PEFT
Safetensors
English
lora

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

Untitled (2)

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()
Downloads last month
671
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for FWKV/FWKV-Image-Hibiscus-LoRA

Base model

FWKV/FWKV-Image
Adapter
(1)
this model

Dataset used to train FWKV/FWKV-Image-Hibiscus-LoRA

Collection including FWKV/FWKV-Image-Hibiscus-LoRA