#!/usr/bin/env python3 """Train Clover Image Tiny v2 as a context-aware 9-channel inpainting model. The trainer warm-starts from the existing Clover inpainting checkpoint and distills the official SD 1.5 inpainting U-Net while retaining a ground-truth diffusion objective. Masks are synthesized on the fly with free-form, multi-region, object-like, and outpainting geometry. """ from __future__ import annotations import argparse import json import random import shutil from pathlib import Path from typing import Any import numpy as np import torch import torch.nn.functional as F from accelerate import Accelerator from datasets import load_dataset from diffusers import ( AutoencoderKL, DDPMScheduler, StableDiffusionInpaintPipeline, UNet2DConditionModel, get_scheduler, ) from PIL import Image, ImageOps from transformers import CLIPTextModel, CLIPTokenizer from inpainting.masks import apply_mask, mask_area_fraction, random_mask from inpainting.model import make_inpainting_unet from inpainting.objective import min_snr_weights, spatial_loss_weights, weighted_mse def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--pretrained_model_name_or_path", required=True) parser.add_argument("--revision") parser.add_argument("--initial_inpaint_model") parser.add_argument("--initial_inpaint_revision") parser.add_argument("--teacher_model_name_or_path") parser.add_argument("--teacher_revision") parser.add_argument("--dataset_name", required=True) parser.add_argument("--dataset_config_name") parser.add_argument("--dataset_revision") parser.add_argument("--dataset_split", default="train") parser.add_argument("--image_column", default="image") parser.add_argument("--caption_column", default="caption") parser.add_argument("--resolution", type=int, default=512) parser.add_argument("--train_batch_size", type=int, default=1) parser.add_argument("--max_train_steps", type=int, default=12000) parser.add_argument("--learning_rate", type=float, default=5e-6) parser.add_argument("--lr_scheduler", default="cosine") parser.add_argument("--lr_warmup_steps", type=int, default=500) parser.add_argument("--gradient_accumulation_steps", type=int, default=4) parser.add_argument("--gradient_checkpointing", action="store_true") parser.add_argument("--mixed_precision", choices=("no", "fp16", "bf16"), default="bf16") parser.add_argument("--seed", type=int, default=20260811) parser.add_argument("--mask_min_area", type=float, default=0.04) parser.add_argument("--mask_max_area", type=float, default=0.65) parser.add_argument("--caption_dropout_probability", type=float, default=0.10) parser.add_argument("--teacher_loss_weight", type=float, default=0.75) parser.add_argument("--ground_truth_loss_weight", type=float, default=0.25) parser.add_argument("--context_loss_weight", type=float, default=0.25) parser.add_argument("--masked_loss_weight", type=float, default=2.5) parser.add_argument("--boundary_loss_weight", type=float, default=2.0) parser.add_argument("--boundary_radius", type=int, default=2) parser.add_argument("--snr_gamma", type=float, default=5.0) parser.add_argument("--random_flip", action="store_true") parser.add_argument("--max_train_samples", type=int) parser.add_argument("--validation_samples", type=int, default=128) parser.add_argument("--num_workers", type=int, default=4) parser.add_argument("--checkpointing_steps", type=int, default=500) parser.add_argument("--checkpoints_total_limit", type=int, default=3) parser.add_argument("--resume_from_checkpoint") parser.add_argument("--output_dir", type=Path, required=True) parser.add_argument("--push_to_hub", action="store_true") parser.add_argument("--hub_model_id") return parser.parse_args() def model_kwargs(revision: str | None) -> dict[str, Any]: return {"revision": revision} if revision else {} def _frozen_weight_dtype(mixed_precision: str) -> torch.dtype: if mixed_precision == "fp16": return torch.float16 if mixed_precision == "bf16": return torch.bfloat16 return torch.float32 def _normalize_caption(value: Any) -> str: if isinstance(value, list): value = value[0] if value else "" return " ".join(str(value or "").split()) def _crop_image( image: Image.Image, *, resolution: int, rng: random.Random, random_flip: bool, ) -> Image.Image: image = ImageOps.exif_transpose(image).convert("RGB") width, height = image.size scale = rng.uniform(0.78, 1.0) crop_side = max(1, round(min(width, height) * scale)) left = rng.randint(0, max(0, width - crop_side)) top = rng.randint(0, max(0, height - crop_side)) image = image.crop((left, top, left + crop_side, top + crop_side)) image = image.resize((resolution, resolution), Image.Resampling.LANCZOS) if random_flip and rng.random() < 0.5: image = ImageOps.mirror(image) return image def _image_to_tensor(image: Image.Image) -> torch.Tensor: array = np.asarray(image.convert("RGB"), dtype=np.float32) / 127.5 - 1.0 return torch.from_numpy(array).permute(2, 0, 1).contiguous() def make_collate_fn( *, tokenizer: CLIPTokenizer, resolution: int, min_area: float, max_area: float, seed: int, caption_dropout_probability: float, random_flip: bool, ): worker_rng: random.Random | None = None def collate(examples: list[dict[str, Any]]) -> dict[str, Any]: nonlocal worker_rng if worker_rng is None: # torch.initial_seed differs across DataLoader workers and remains # stable for the lifetime of each worker process. worker_rng = random.Random(seed ^ torch.initial_seed()) images: list[torch.Tensor] = [] masked_images: list[torch.Tensor] = [] masks: list[torch.Tensor] = [] captions: list[str] = [] mask_areas: list[float] = [] for example in examples: image = example["image"] if not isinstance(image, Image.Image): image = Image.fromarray(np.asarray(image)) image = _crop_image( image, resolution=resolution, rng=worker_rng, random_flip=random_flip, ) mask = random_mask( (resolution, resolution), worker_rng, min_area=min_area, max_area=max_area, ) masked = apply_mask(image, mask) images.append(_image_to_tensor(image)) masked_images.append(_image_to_tensor(masked)) masks.append( torch.from_numpy(np.asarray(mask, dtype=np.float32) / 255.0).unsqueeze(0) ) caption = _normalize_caption(example["caption"]) if worker_rng.random() < caption_dropout_probability: caption = "" captions.append(caption) mask_areas.append(mask_area_fraction(mask)) tokenized = tokenizer( captions, max_length=tokenizer.model_max_length, padding="max_length", truncation=True, return_tensors="pt", ) return { "pixel_values": torch.stack(images), "masked_pixel_values": torch.stack(masked_images), "mask": torch.stack(masks), "input_ids": tokenized.input_ids, "mask_area": torch.tensor(mask_areas, dtype=torch.float32), } return collate def _load_student(args: argparse.Namespace) -> UNet2DConditionModel: if args.initial_inpaint_model: student = UNet2DConditionModel.from_pretrained( args.initial_inpaint_model, subfolder="unet", low_cpu_mem_usage=False, **model_kwargs(args.initial_inpaint_revision), ) if student.config.in_channels != 9: raise ValueError("initial inpainting U-Net must have nine input channels") return student return make_inpainting_unet( args.pretrained_model_name_or_path, revision=args.revision, ) def _save_pipeline( *, args: argparse.Namespace, unet: UNet2DConditionModel, output_dir: Path, ) -> None: output_dir.mkdir(parents=True, exist_ok=True) pipeline = StableDiffusionInpaintPipeline.from_pretrained( args.pretrained_model_name_or_path, revision=args.revision, unet=unet, ) pipeline.save_pretrained(output_dir, safe_serialization=True) (output_dir / "inpainting-config.json").write_text( json.dumps( { "recipe_version": 2, "unet_in_channels": 9, "mask_semantics": "white=regenerate, black=preserve", "base_model": args.pretrained_model_name_or_path, "base_revision": args.revision, "initial_inpaint_model": args.initial_inpaint_model, "initial_inpaint_revision": args.initial_inpaint_revision, "teacher_model": args.teacher_model_name_or_path, "teacher_revision": args.teacher_revision, "objective": "teacher distillation + ground-truth diffusion", "mask_distribution": "brush, multi-brush, box, ellipse, polygon, multi-region, outpaint", }, indent=2, ) + "\n" ) def _prune_checkpoints(checkpoint_root: Path, limit: int) -> None: checkpoints = sorted( (path for path in checkpoint_root.glob("checkpoint-*") if path.is_dir()), key=lambda path: int(path.name.rsplit("-", 1)[-1]), ) for path in checkpoints[:-limit]: shutil.rmtree(path) def _validate_args(args: argparse.Namespace) -> None: if not 0.0 < args.mask_min_area < args.mask_max_area < 1.0: raise ValueError("mask area bounds must satisfy 0 < min < max < 1") if not 0.0 <= args.caption_dropout_probability < 1.0: raise ValueError("caption dropout probability must be in [0, 1)") if args.teacher_loss_weight < 0 or args.ground_truth_loss_weight < 0: raise ValueError("loss weights must be non-negative") if args.teacher_loss_weight > 0 and not args.teacher_model_name_or_path: raise ValueError("teacher model is required when teacher loss is enabled") if args.teacher_loss_weight + args.ground_truth_loss_weight <= 0: raise ValueError("at least one training objective must be enabled") if args.push_to_hub and not args.hub_model_id: raise ValueError("--hub_model_id is required with --push_to_hub") if args.resolution % 8: raise ValueError("resolution must be divisible by 8") def main() -> None: args = parse_args() _validate_args(args) args.output_dir.mkdir(parents=True, exist_ok=True) checkpoint_root = args.output_dir / "checkpoints" metrics_path = args.output_dir / "training-metrics.jsonl" accelerator = Accelerator( gradient_accumulation_steps=args.gradient_accumulation_steps, mixed_precision=args.mixed_precision, ) torch.manual_seed(args.seed) random.seed(args.seed) np.random.seed(args.seed % (2**32)) kwargs = model_kwargs(args.revision) weight_dtype = _frozen_weight_dtype(args.mixed_precision) dataset = load_dataset( args.dataset_name, args.dataset_config_name, split=args.dataset_split, revision=args.dataset_revision, ) dataset = dataset.shuffle(seed=args.seed) if args.validation_samples < 0 or args.validation_samples >= len(dataset): raise ValueError("validation_samples must be smaller than the dataset") if args.validation_samples: dataset = dataset.select(range(len(dataset) - args.validation_samples)) if args.max_train_samples: dataset = dataset.select(range(min(args.max_train_samples, len(dataset)))) if args.image_column not in dataset.column_names: raise ValueError(f"Missing image column {args.image_column!r}: {dataset.column_names}") if args.caption_column not in dataset.column_names: raise ValueError(f"Missing caption column {args.caption_column!r}: {dataset.column_names}") dataset = dataset.rename_columns( {args.image_column: "image", args.caption_column: "caption"} ) tokenizer = CLIPTokenizer.from_pretrained( args.pretrained_model_name_or_path, subfolder="tokenizer", **kwargs, ) text_encoder = CLIPTextModel.from_pretrained( args.pretrained_model_name_or_path, subfolder="text_encoder", dtype=weight_dtype, **kwargs, ) vae = AutoencoderKL.from_pretrained( args.pretrained_model_name_or_path, subfolder="vae", torch_dtype=weight_dtype, **kwargs, ) student = _load_student(args) noise_scheduler = DDPMScheduler.from_pretrained( args.pretrained_model_name_or_path, subfolder="scheduler", **kwargs, ) teacher = None if args.teacher_loss_weight > 0: teacher = UNet2DConditionModel.from_pretrained( args.teacher_model_name_or_path, subfolder="unet", low_cpu_mem_usage=True, torch_dtype=weight_dtype, use_safetensors=True, variant="fp16", **model_kwargs(args.teacher_revision), ) if teacher.config.in_channels != 9: raise ValueError("teacher inpainting U-Net must have nine input channels") if args.gradient_checkpointing: student.enable_gradient_checkpointing() text_encoder.requires_grad_(False).eval() vae.requires_grad_(False).eval() if teacher is not None: teacher.requires_grad_(False).eval() collate_fn = make_collate_fn( tokenizer=tokenizer, resolution=args.resolution, min_area=args.mask_min_area, max_area=args.mask_max_area, seed=args.seed, caption_dropout_probability=args.caption_dropout_probability, random_flip=args.random_flip, ) dataloader = torch.utils.data.DataLoader( dataset, shuffle=True, collate_fn=collate_fn, batch_size=args.train_batch_size, num_workers=args.num_workers, pin_memory=True, persistent_workers=args.num_workers > 0, ) optimizer = torch.optim.AdamW( student.parameters(), lr=args.learning_rate, betas=(0.9, 0.999), weight_decay=1e-2, eps=1e-8, ) lr_scheduler = get_scheduler( args.lr_scheduler, optimizer=optimizer, num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, num_training_steps=args.max_train_steps * accelerator.num_processes, ) student, optimizer, dataloader, lr_scheduler = accelerator.prepare( student, optimizer, dataloader, lr_scheduler ) vae.to(accelerator.device) text_encoder.to(accelerator.device) if teacher is not None: teacher.to(accelerator.device) global_step = 0 if args.resume_from_checkpoint: checkpoint = Path(args.resume_from_checkpoint) if args.resume_from_checkpoint == "latest": candidates = sorted( checkpoint_root.glob("checkpoint-*"), key=lambda path: int(path.name.rsplit("-", 1)[-1]), ) if not candidates: raise ValueError("No checkpoint is available to resume") checkpoint = candidates[-1] accelerator.load_state(checkpoint) global_step = int(checkpoint.name.rsplit("-", 1)[-1]) accelerator.print(f"Resumed from {checkpoint} at step {global_step}") if accelerator.is_main_process: summary = vars(args).copy() summary["output_dir"] = str(args.output_dir) summary["dataset_size"] = len(dataset) summary["effective_batch_size"] = ( args.train_batch_size * args.gradient_accumulation_steps * accelerator.num_processes ) (args.output_dir / "training-summary.json").write_text( json.dumps(summary, indent=2, default=str) + "\n" ) data_iterator = iter(dataloader) while global_step < args.max_train_steps: try: batch = next(data_iterator) except StopIteration: data_iterator = iter(dataloader) batch = next(data_iterator) with accelerator.accumulate(student): pixel_values = batch["pixel_values"].to( accelerator.device, dtype=weight_dtype, non_blocking=True ) masked_pixel_values = batch["masked_pixel_values"].to( accelerator.device, dtype=weight_dtype, non_blocking=True ) mask = batch["mask"].to( accelerator.device, dtype=weight_dtype, non_blocking=True ) with torch.no_grad(): latents = vae.encode(pixel_values).latent_dist.sample() latents = latents * vae.config.scaling_factor masked_latents = vae.encode(masked_pixel_values).latent_dist.sample() masked_latents = masked_latents * vae.config.scaling_factor encoder_hidden_states = text_encoder( batch["input_ids"].to(accelerator.device) )[0] noise = torch.randn_like(latents) timesteps = torch.randint( 0, noise_scheduler.config.num_train_timesteps, (latents.shape[0],), device=latents.device, ).long() noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps) latent_mask = F.interpolate(mask, size=latents.shape[-2:], mode="nearest") model_input = torch.cat([noisy_latents, latent_mask, masked_latents], dim=1) model_pred = student( model_input, timesteps, encoder_hidden_states=encoder_hidden_states, ).sample prediction_type = noise_scheduler.config.prediction_type if prediction_type == "epsilon": ground_truth_target = noise elif prediction_type == "v_prediction": ground_truth_target = noise_scheduler.get_velocity(latents, noise, timesteps) else: raise ValueError(f"Unsupported prediction type: {prediction_type}") spatial_weights = spatial_loss_weights( latent_mask, context_weight=args.context_loss_weight, masked_weight=args.masked_loss_weight, boundary_weight=args.boundary_loss_weight, boundary_radius=args.boundary_radius, ) sample_weights = min_snr_weights( noise_scheduler.alphas_cumprod, timesteps, gamma=args.snr_gamma, prediction_type=prediction_type, ) ground_truth_loss = weighted_mse( model_pred, ground_truth_target, spatial_weights=spatial_weights, sample_weights=sample_weights, ) teacher_loss = torch.zeros((), device=accelerator.device) if teacher is not None: with torch.no_grad(): teacher_target = teacher( model_input, timesteps, encoder_hidden_states=encoder_hidden_states, ).sample teacher_loss = weighted_mse( model_pred, teacher_target, spatial_weights=spatial_weights, sample_weights=sample_weights, ) loss = ( args.ground_truth_loss_weight * ground_truth_loss + args.teacher_loss_weight * teacher_loss ) accelerator.backward(loss) if accelerator.sync_gradients: accelerator.clip_grad_norm_(student.parameters(), 1.0) optimizer.step() lr_scheduler.step() optimizer.zero_grad(set_to_none=True) if accelerator.sync_gradients: global_step += 1 metrics = { "step": global_step, "loss": loss.detach().float().item(), "teacher_loss": teacher_loss.detach().float().item(), "ground_truth_loss": ground_truth_loss.detach().float().item(), "lr": lr_scheduler.get_last_lr()[0], "mask_area": batch["mask_area"].float().mean().item(), } if accelerator.is_main_process and (global_step == 1 or global_step % 25 == 0): accelerator.print( f"step={global_step}/{args.max_train_steps} " f"loss={metrics['loss']:.5f} " f"teacher={metrics['teacher_loss']:.5f} " f"ground_truth={metrics['ground_truth_loss']:.5f} " f"mask={metrics['mask_area']:.3f} " f"lr={metrics['lr']:.3e}" ) with metrics_path.open("a") as handle: handle.write(json.dumps(metrics) + "\n") if global_step % args.checkpointing_steps == 0: accelerator.wait_for_everyone() checkpoint = checkpoint_root / f"checkpoint-{global_step}" accelerator.save_state(checkpoint, safe_serialization=True) if accelerator.is_main_process: _prune_checkpoints(checkpoint_root, args.checkpoints_total_limit) (args.output_dir / "progress.json").write_text( json.dumps(metrics, indent=2) + "\n" ) accelerator.print(f"Saved resumable checkpoint {checkpoint}") accelerator.wait_for_everyone() if accelerator.is_main_process: unwrapped = accelerator.unwrap_model(student).cpu() _save_pipeline(args=args, unet=unwrapped, output_dir=args.output_dir) final_metrics = { "completed_steps": global_step, "final_loss": loss.detach().float().item(), "final_teacher_loss": teacher_loss.detach().float().item(), "final_ground_truth_loss": ground_truth_loss.detach().float().item(), } (args.output_dir / "training-complete.json").write_text( json.dumps(final_metrics, indent=2) + "\n" ) if args.push_to_hub: pipeline = StableDiffusionInpaintPipeline.from_pretrained(args.output_dir) pipeline.push_to_hub(args.hub_model_id) accelerator.print(f"Saved Clover Image Tiny Inpaint v2 to {args.output_dir}") accelerator.end_training() if __name__ == "__main__": main()