neonforestmist's picture
Document and package context-aware inpainting v2
95bf78b verified
Raw
History Blame Contribute Delete
23.1 kB
#!/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()