Spaces:
Sleeping
Sleeping
| """ | |
| LoRA fine-tuning for PixelDiT using precomputed image+caption embeddings. | |
| Flow matching loss: predict velocity (noise - x), noisy image via shifted schedule. | |
| Only LoRA weights update — base model is fully frozen. | |
| Usage: | |
| # Precompute first: | |
| python scripts/precompute_lora_data.py --images /data/my_images --out /data/lora_cache | |
| # Train: | |
| python scripts/train_lora.py --data /data/lora_cache --out lora_out/ --epochs 100 | |
| # Inference with trained LoRA: | |
| from peft import PeftModel | |
| from pixeldit.modeling_pixeldit_hf import PixelDiTModel | |
| model = PixelDiTModel.from_pretrained("madtune/pixeldit-diffusers", subfolder="transformer") | |
| model = PeftModel.from_pretrained(model, "lora_out/") | |
| """ | |
| import argparse | |
| import json | |
| import os | |
| import sys | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from torch.utils.data import DataLoader, Dataset | |
| from tqdm import tqdm | |
| # ---- Flow schedule (matching NVIDIA's training config) ---------------------- | |
| _T = 1000 | |
| _TXT_MAX = 300 | |
| _GEMMA_ID = "Efficient-Large-Model/gemma-2-2b-it" | |
| _SELECT_IDX = [0] + list(range(-(_TXT_MAX - 1), 0)) | |
| def _build_flow_schedule(flow_shift: float): | |
| # sigmas[t] ≈ shift * (t/1000) / (1 + (shift-1) * (t/1000)) | |
| # matches FlowMatchEulerDiscreteScheduler with the same shift value | |
| betas = np.linspace(1.0, 0.001, _T, dtype=np.float64) | |
| sigmas_raw = 1.0 - betas | |
| sigmas = flow_shift * sigmas_raw / (1 + (flow_shift - 1) * sigmas_raw) | |
| alphas = 1.0 - sigmas | |
| return torch.from_numpy(sigmas).float(), torch.from_numpy(alphas).float() | |
| def q_sample(x, t, noise, alphas, sigmas): | |
| a = alphas[t].view(-1, 1, 1, 1) | |
| s = sigmas[t].view(-1, 1, 1, 1) | |
| return a * x + s * noise | |
| # ---- Dataset ---------------------------------------------------------------- | |
| class LoraDataset(Dataset): | |
| def __init__(self, data_dir): | |
| meta_path = os.path.join(data_dir, "meta.json") | |
| self.meta = {} | |
| if os.path.exists(meta_path): | |
| with open(meta_path, "r", encoding="utf-8") as f: | |
| self.meta = json.load(f) | |
| print( | |
| f"[dataset] encoder={self.meta.get('encoder', 'unknown')} " | |
| f"n={self.meta.get('n_samples')} " | |
| f"emb_dim={self.meta.get('emb_dim')} " | |
| f"trigger={self.meta.get('trigger') or 'none'}" | |
| ) | |
| self.imgs = np.load(os.path.join(data_dir, "lora_images.npy"), mmap_mode="r") | |
| self.embs = np.load(os.path.join(data_dir, "lora_embs.npy"), mmap_mode="r") | |
| self.masks = np.load(os.path.join(data_dir, "lora_masks.npy"), mmap_mode="r") | |
| captions_path = os.path.join(data_dir, "captions.json") | |
| self.captions = None | |
| if os.path.exists(captions_path): | |
| with open(captions_path, "r", encoding="utf-8") as f: | |
| self.captions = json.load(f) | |
| if len(self.captions) != len(self.imgs): | |
| raise RuntimeError( | |
| f"{captions_path} has {len(self.captions)} captions for {len(self.imgs)} images" | |
| ) | |
| assert len(self.imgs) == len(self.embs) == len(self.masks) | |
| # sanity check: catch all-zeros from a failed precompute run | |
| sample = self.embs[:min(4, len(self.embs))] | |
| if np.count_nonzero(sample) == 0: | |
| raise RuntimeError( | |
| f"lora_embs.npy in {data_dir} are ALL ZEROS — " | |
| "re-run precompute_lora_data.py to regenerate." | |
| ) | |
| def __len__(self): | |
| return len(self.imgs) | |
| def img_size(self): | |
| return self.meta.get("img_size", 512) | |
| def __getitem__(self, idx): | |
| img = torch.from_numpy(self.imgs[idx].astype(np.float32)) # [3, H, H] | |
| mask = torch.from_numpy(self.masks[idx].astype(np.float32)) # [300] | |
| return idx, img, mask # return idx so we can look up learnable embeddings | |
| # ---- Learnable Embeddings --------------------------------------------------- | |
| class LearnableEmbeddings(nn.Module): | |
| """Wrapper to make embeddings learnable parameters during training.""" | |
| def __init__(self, embeddings: torch.Tensor): | |
| super().__init__() | |
| # embeddings: [N, 300, 2304] float32 | |
| self.embs = nn.Parameter(embeddings) | |
| def forward(self, indices: torch.Tensor): | |
| """Return embeddings for given indices. indices: [B]""" | |
| return self.embs[indices] # [B, 300, 2304] | |
| def encode_gemma_batch(text_encoder, tokenizer, captions, device): | |
| toks = tokenizer( | |
| captions, | |
| max_length=_TXT_MAX, | |
| padding="max_length", | |
| truncation=True, | |
| return_tensors="pt", | |
| ).to(device) | |
| emb = text_encoder( | |
| input_ids=toks.input_ids, | |
| attention_mask=toks.attention_mask, | |
| ).last_hidden_state | |
| return emb[:, _SELECT_IDX, :] | |
| # ---- Main ------------------------------------------------------------------- | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--data", required=True, help="precomputed cache dir") | |
| ap.add_argument("--out", default="lora_out/", help="output dir for LoRA weights") | |
| ap.add_argument("--model", default="madtune/pixeldit-diffusers") | |
| ap.add_argument("--epochs", type=int, default=100) | |
| ap.add_argument("--batch", type=int, default=2) | |
| ap.add_argument("--accum", type=int, default=4, help="gradient accumulation steps") | |
| ap.add_argument("--lr", type=float, default=1e-4) | |
| ap.add_argument("--lora_r", type=int, default=16) | |
| ap.add_argument("--lora_alpha", type=int, default=16) | |
| ap.add_argument("--train_text_encoder", action="store_true", | |
| help="train a Gemma LoRA together with the PixelDiT transformer LoRA") | |
| ap.add_argument("--text_encoder_model", default=_GEMMA_ID, help="Gemma model id/path") | |
| ap.add_argument("--text_device", default=None, | |
| help="device for Gemma LoRA training; default follows --device") | |
| ap.add_argument("--text_lora_r", type=int, default=8) | |
| ap.add_argument("--text_lora_alpha", type=int, default=8) | |
| ap.add_argument("--text_lora_targets", default="q_proj,k_proj,v_proj,o_proj", | |
| help="comma-separated Gemma module names for PEFT LoRA") | |
| ap.add_argument("--text_lr", type=float, default=None, | |
| help="Gemma LoRA LR; default is lr * 0.25") | |
| ap.add_argument("--cfg_drop", type=float, default=0.1, help="CFG dropout probability") | |
| ap.add_argument("--device", default="cuda:0") | |
| ap.add_argument("--save_every", type=int, default=10) | |
| ap.add_argument("--flow_shift", type=float, default=4.0, | |
| help="flow schedule shift — must match inference scheduler (default 4.0 for 1024px)") | |
| ap.add_argument("--grad_ckpt", action="store_true", | |
| help="enable gradient checkpointing to reduce VRAM at the cost of speed") | |
| ap.add_argument("--timestep_logit_std", type=float, default=1.0, | |
| help="std of logit-normal timestep sampling (higher = more uniform; 0 = pure midpoint)") | |
| ap.add_argument("--loss_weighting", default="sigma_sqrt", choices=["sigma_sqrt", "none"], | |
| help="sigma_sqrt upweights low-noise steps (identity/detail); none = uniform (Flux default: sigma_sqrt)") | |
| args = ap.parse_args() | |
| os.makedirs(args.out, exist_ok=True) | |
| device = torch.device(args.device) | |
| text_device = torch.device(args.text_device if args.text_device else args.device) | |
| # 1. Load model + inject LoRA | |
| print("Loading PixelDiTModel...") | |
| try: | |
| from peft import get_peft_model, LoraConfig | |
| except ImportError: | |
| print("peft not installed — run: pip install peft") | |
| sys.exit(1) | |
| from diffusers.pipelines.pixeldit import PixelDiTModel | |
| model = PixelDiTModel.from_pretrained(args.model, subfolder="transformer") | |
| lora_cfg = LoraConfig( | |
| r = args.lora_r, | |
| lora_alpha = args.lora_alpha, | |
| target_modules = ["qkv_x", "qkv_y", "proj_x", "proj_y"], | |
| lora_dropout = 0.05, | |
| bias = "none", | |
| ) | |
| model = get_peft_model(model, lora_cfg) | |
| model.print_trainable_parameters() | |
| if args.grad_ckpt: | |
| model.enable_input_require_grads() | |
| model.gradient_checkpointing_enable() | |
| print("[+] Gradient checkpointing enabled") | |
| model = model.to(device).train() | |
| print(f"[devices] PixelDiT={device} Gemma={text_device if args.train_text_encoder else 'off'}") | |
| text_encoder = None | |
| tokenizer = None | |
| # null embedding for CFG dropout — zeros | |
| null_emb = torch.zeros(1, 300, 2304, device=device) | |
| null_mask = torch.zeros(1, 300, device=device) | |
| # 2. Dataset + loader | |
| dataset = LoraDataset(args.data) | |
| img_size = dataset.img_size | |
| loader = DataLoader(dataset, batch_size=args.batch, shuffle=True, | |
| num_workers=2, pin_memory=True, drop_last=True) | |
| print(f"Dataset: {len(dataset)} samples img_size={img_size} batch={args.batch} accum={args.accum} steps/epoch={len(loader)}") | |
| # 2b. Text conditioning path | |
| learnable_embs = None | |
| opt_groups = [{"params": [p for p in model.parameters() if p.requires_grad]}] | |
| if args.train_text_encoder: | |
| if dataset.meta.get("encoder") not in (None, "gemma"): | |
| raise RuntimeError("--train_text_encoder requires a Gemma precompute cache") | |
| if dataset.captions is None: | |
| raise RuntimeError( | |
| "--train_text_encoder requires captions.json. " | |
| "Re-run scripts/precompute_lora_data.py with this patched version." | |
| ) | |
| print("Loading Gemma text encoder + LoRA...") | |
| from transformers import AutoTokenizer, AutoModelForCausalLM | |
| tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_model) | |
| tokenizer.padding_side = "right" | |
| text_encoder = ( | |
| AutoModelForCausalLM.from_pretrained( | |
| args.text_encoder_model, torch_dtype=torch.bfloat16 | |
| ) | |
| .get_decoder() | |
| ) | |
| if hasattr(text_encoder, "config"): | |
| text_encoder.config.use_cache = False | |
| text_lora_cfg = LoraConfig( | |
| r=args.text_lora_r, | |
| lora_alpha=args.text_lora_alpha, | |
| target_modules=[m.strip() for m in args.text_lora_targets.split(",") if m.strip()], | |
| lora_dropout=0.05, | |
| bias="none", | |
| ) | |
| text_encoder = get_peft_model(text_encoder, text_lora_cfg) | |
| if args.grad_ckpt and hasattr(text_encoder, "gradient_checkpointing_enable"): | |
| text_encoder.gradient_checkpointing_enable() | |
| text_encoder.enable_input_require_grads() | |
| text_encoder = text_encoder.to(text_device).train() | |
| text_encoder.print_trainable_parameters() | |
| opt_groups.append({ | |
| "params": [p for p in text_encoder.parameters() if p.requires_grad], | |
| "lr": args.text_lr if args.text_lr is not None else args.lr * 0.25, | |
| }) | |
| else: | |
| print("Loading embeddings as learnable parameters...") | |
| all_embs = torch.from_numpy(dataset.embs[:].astype(np.float32)) | |
| learnable_embs = LearnableEmbeddings(all_embs).to(device) | |
| print(f" {learnable_embs.embs.numel()} embedding params") | |
| opt_groups.append({"params": learnable_embs.parameters(), "lr": args.lr * 0.1}) | |
| # 3. Optimizer | |
| opt = torch.optim.AdamW(opt_groups, lr=args.lr, weight_decay=1e-2) | |
| total_steps = args.epochs * len(loader) // args.accum | |
| sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=total_steps, eta_min=args.lr * 0.1) | |
| sigmas_cpu, alphas_cpu = _build_flow_schedule(args.flow_shift) | |
| alphas = alphas_cpu.to(device) | |
| sigmas = sigmas_cpu.to(device) | |
| print(f"Flow shift: {args.flow_shift}") | |
| best_loss = float("inf") | |
| step = 0 | |
| def save_lora(path, loss): | |
| os.makedirs(path, exist_ok=True) | |
| transformer_dir = os.path.join(path, "transformer") | |
| model.save_pretrained(transformer_dir) | |
| if text_encoder is not None: | |
| text_encoder.save_pretrained(os.path.join(path, "text_encoder")) | |
| if learnable_embs is not None: | |
| torch.save(learnable_embs.state_dict(), os.path.join(path, "learnable_embs.pt")) | |
| meta = { | |
| "data_dir": args.data, | |
| "model": args.model, | |
| "img_size": img_size, | |
| "epochs": args.epochs, | |
| "batch": args.batch, | |
| "accum": args.accum, | |
| "lr": args.lr, | |
| "lora_r": args.lora_r, | |
| "lora_alpha": args.lora_alpha, | |
| "cfg_drop": args.cfg_drop, | |
| "flow_shift": args.flow_shift, | |
| "timestep_logit_std": args.timestep_logit_std, | |
| "loss_weighting": args.loss_weighting, | |
| "loss": loss, | |
| "has_learnable_embeddings": learnable_embs is not None, | |
| "has_text_encoder_lora": text_encoder is not None, | |
| "text_encoder_model": args.text_encoder_model if text_encoder is not None else None, | |
| "text_device": str(text_device) if text_encoder is not None else None, | |
| "text_lora_r": args.text_lora_r if text_encoder is not None else None, | |
| "text_lora_alpha": args.text_lora_alpha if text_encoder is not None else None, | |
| "text_lora_targets": args.text_lora_targets if text_encoder is not None else None, | |
| "text_lr": args.text_lr if args.text_lr is not None else args.lr * 0.25, | |
| "precompute": dataset.meta, | |
| } | |
| with open(os.path.join(path, "training_meta.json"), "w", encoding="utf-8") as f: | |
| json.dump(meta, f, indent=2) | |
| for epoch in range(args.epochs): | |
| total_loss, n = 0.0, 0 | |
| bar = tqdm(loader, desc=f"epoch {epoch+1}/{args.epochs}") | |
| opt.zero_grad() | |
| for i, (indices, imgs, _masks) in enumerate(bar): | |
| indices = indices.to(device) # [B] batch indices | |
| imgs = imgs.to(device) # [B, 3, H, H] | |
| B = imgs.shape[0] | |
| if text_encoder is not None: | |
| batch_captions = [dataset.captions[int(idx)] for idx in indices.detach().cpu().tolist()] | |
| with torch.autocast(device_type=text_device.type, dtype=torch.bfloat16): | |
| embs = encode_gemma_batch(text_encoder, tokenizer, batch_captions, text_device) | |
| embs = embs.to(device) | |
| else: | |
| embs = learnable_embs(indices) | |
| # CFG dropout: replace some embeddings with the null (empty) embedding | |
| if args.cfg_drop > 0: | |
| drop = (torch.rand(B, device=device) < args.cfg_drop).view(B, 1, 1) | |
| embs = torch.where(drop, null_emb.expand(B, -1, -1).to(embs.dtype), embs) | |
| # Logit-normal timestep sampling — concentrates gradient signal | |
| # around mid-noise where the model does the most meaningful work. | |
| noise = torch.randn_like(imgs) | |
| u = torch.sigmoid(torch.randn(B, device=device) * args.timestep_logit_std) | |
| t = (u * _T).long().clamp(0, _T - 1) | |
| x_t = q_sample(imgs, t, noise, alphas, sigmas) | |
| target = noise - imgs # velocity: direction from data → noise | |
| with torch.autocast(device_type="cuda", dtype=torch.bfloat16): | |
| pred = model(x_t.bfloat16(), t, embs.bfloat16()) | |
| # per-sample MSE, then apply sigma weighting before reducing | |
| loss_per = F.mse_loss(pred.float(), target, reduction="none").mean(dim=(1, 2, 3)) | |
| # sigma_sqrt: weight = 1/sigma² — upweights low-noise steps where | |
| # identity and fine detail are learned; matches Flux Dev training. | |
| if args.loss_weighting == "sigma_sqrt": | |
| sig = sigmas[t].clamp(min=1e-3) # [B] | |
| w = (1.0 / sig ** 2).to(loss_per.device) | |
| loss = (w * loss_per).mean() | |
| else: | |
| loss = loss_per.mean() | |
| (loss / args.accum).backward() | |
| total_loss += loss.item() | |
| n += 1 | |
| if (i + 1) % args.accum == 0: | |
| all_params = [p for p in model.parameters() if p.requires_grad] | |
| if text_encoder is not None: | |
| all_params += [p for p in text_encoder.parameters() if p.requires_grad] | |
| if learnable_embs is not None: | |
| all_params += list(learnable_embs.parameters()) | |
| nn.utils.clip_grad_norm_(all_params, 1.0) | |
| opt.step() | |
| sched.step() | |
| opt.zero_grad() | |
| step += 1 | |
| bar.set_postfix(loss=f"{total_loss/n:.4f}", lr=f"{sched.get_last_lr()[0]:.2e}") | |
| # apply gradients from any tail batches that didn't fill a full accum window | |
| if n % args.accum != 0: | |
| all_params = [p for p in model.parameters() if p.requires_grad] | |
| if text_encoder is not None: | |
| all_params += [p for p in text_encoder.parameters() if p.requires_grad] | |
| if learnable_embs is not None: | |
| all_params += list(learnable_embs.parameters()) | |
| nn.utils.clip_grad_norm_(all_params, 1.0) | |
| opt.step() | |
| sched.step() | |
| opt.zero_grad() | |
| step += 1 | |
| avg = total_loss / n | |
| print(f" epoch {epoch+1} loss={avg:.4f}") | |
| is_last = (epoch + 1) == args.epochs | |
| is_save = (epoch + 1) % args.save_every == 0 | |
| is_best = avg < best_loss | |
| if is_best: | |
| best_loss = avg | |
| save_lora(os.path.join(args.out, "best"), avg) | |
| print(f" best saved → {args.out}/best/") | |
| if is_save or is_best or is_last: | |
| ckpt_name = f"ckpt_epoch_{epoch+1:03d}" | |
| save_lora(os.path.join(args.out, ckpt_name), avg) | |
| print(f" checkpoint → {args.out}/{ckpt_name}/") | |
| print(f"\nDone. Best loss: {best_loss:.4f}") | |
| print(f"LoRA checkpoints saved to {args.out}/") | |
| print("\nEach checkpoint contains:") | |
| print(" transformer/adapter_model.safetensors - PixelDiT transformer LoRA") | |
| if args.train_text_encoder: | |
| print(" text_encoder/adapter_model.safetensors - Gemma text-encoder LoRA") | |
| else: | |
| print(" learnable_embs.pt - fine-tuned per-image embeddings [N, 300, 2304]") | |
| print(" training_meta.json - metadata") | |
| print("\nTo use in inference:") | |
| print(f" python generate.py --prompt 'corrychase, your prompt here' --lora {args.out}/best --device {args.device}") | |
| if __name__ == "__main__": | |
| main() | |