#!/usr/bin/env python3 """Train and compare single-block Self-Forcing Predictor initializations.""" from __future__ import annotations import argparse import csv import hashlib import json import math import os import random import sys import time from pathlib import Path from typing import Any def _preparse_gpu() -> str: parser = argparse.ArgumentParser(add_help=False) parser.add_argument("--gpu", default="2") args, _ = parser.parse_known_args() # Respect torchrun's inherited multi-GPU visibility. Standalone invocations # keep the historical GPU-2 default, while an explicit --gpu still wins. gpu_was_explicit = any( value == "--gpu" or value.startswith("--gpu=") for value in sys.argv[1:] ) if gpu_was_explicit or "CUDA_VISIBLE_DEVICES" not in os.environ: os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu) return str(args.gpu) PHYSICAL_GPU = _preparse_gpu() import torch import torch.nn.functional as F from omegaconf import OmegaConf from safetensors.torch import save_file from torch.optim import AdamW REPO_ROOT = Path(__file__).resolve().parents[1] if str(REPO_ROOT) not in sys.path: sys.path.insert(0, str(REPO_ROOT)) from predictor_training.offline_data import ( OfflinePredictorStore, TOKENS_PER_CHUNK, ) from predictor_training.single_block import ( SingleBlockPredictor, TripleFeatureFusion, initialize_predictor_block, ) from utils.misc import set_seed from utils.wan_wrapper import WanDiffusionWrapper from wan.modules.model import sinusoidal_embedding_1d OTHER_METHODS = ( "random_full", "teacher_identity", "random_identity", "full_zero", ) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--gpu", default=PHYSICAL_GPU) parser.add_argument( "--dataset_root", type=Path, default=Path("outputs/predictor_offline_100_all_blocks"), ) parser.add_argument( "--checkpoint_path", type=Path, default=Path("checkpoints/self_forcing_dmd.pt"), ) parser.add_argument( "--config_path", type=Path, default=Path("configs/self_forcing_sid.yaml"), ) parser.add_argument( "--output_dir", type=Path, default=Path("outputs/single_block_init_sweep"), ) parser.add_argument( "--teacher_layers", type=int, nargs="*", default=None, help="Teacher source layers. Omit to sweep 0..29.", ) parser.add_argument("--max_steps", type=int, default=1000) parser.add_argument("--batch_size", type=int, default=32) parser.add_argument("--eval_batch_size", type=int, default=10) parser.add_argument("--eval_every", type=int, default=100) parser.add_argument("--log_every", type=int, default=20) parser.add_argument("--save_every", type=int, default=100) parser.add_argument("--train_prompts", type=int, default=80) parser.add_argument("--val_prompts", type=int, default=20) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--fusion_lr", type=float, default=1e-4) parser.add_argument("--block_lr", type=float, default=1e-5) parser.add_argument("--weight_decay", type=float, default=0.01) parser.add_argument("--hidden_weight", type=float, default=0.1) parser.add_argument("--flow_weight", type=float, default=1.0) parser.add_argument( "--gate_mode", choices=("baseline", "learned", "constant"), default="baseline", ) parser.add_argument("--gate_hidden_dim", type=int, default=128) parser.add_argument("--gate_initial_bias", type=float, default=4.6) parser.add_argument("--gate_floor", type=float, default=0.0) parser.add_argument("--gate_lr", type=float, default=None) parser.add_argument("--gate_freeze_steps", type=int, default=0) parser.add_argument("--constant_gate", type=float, default=1.0) parser.add_argument("--grad_clip", type=float, default=1.0) parser.add_argument("--fusion_warmup_steps", type=int, default=100) parser.add_argument("--block_freeze_steps", type=int, default=100) parser.add_argument("--block_warmup_steps", type=int, default=100) parser.add_argument( "--gradient_checkpointing", action=argparse.BooleanOptionalAction, default=True, ) parser.add_argument( "--run_other_initializations", action=argparse.BooleanOptionalAction, default=True, ) parser.add_argument( "--save_final_weights", action=argparse.BooleanOptionalAction, default=True, ) args = parser.parse_args() if args.max_steps < 1: parser.error("--max_steps must be positive") if args.train_prompts < 1 or args.val_prompts < 1: parser.error("Train and validation prompt counts must be positive") if args.train_prompts + args.val_prompts > 100: parser.error("The offline dataset contains 100 prompts") if args.batch_size > args.train_prompts: parser.error("--batch_size cannot exceed --train_prompts") if args.eval_batch_size > args.val_prompts: args.eval_batch_size = args.val_prompts if not 0.0 <= args.constant_gate <= 1.0: parser.error("--constant_gate must be in [0, 1]") if not 0.0 <= args.gate_floor < 1.0: parser.error("--gate_floor must be in [0, 1)") if args.gate_freeze_steps < 0: parser.error("--gate_freeze_steps must be non-negative") if args.gate_lr is None: args.gate_lr = args.fusion_lr return args def resolve(path: Path) -> Path: path = path.expanduser() return path.resolve() if path.is_absolute() else (REPO_ROOT / path).resolve() def atomic_json(path: Path, value: Any) -> None: path.parent.mkdir(parents=True, exist_ok=True) temporary = path.with_suffix(path.suffix + ".tmp") temporary.write_text( json.dumps(value, indent=2, ensure_ascii=False) + "\n", encoding="utf-8", ) os.replace(temporary, path) def append_jsonl(path: Path, value: dict[str, Any]) -> None: path.parent.mkdir(parents=True, exist_ok=True) with path.open("a", encoding="utf-8") as handle: handle.write(json.dumps(value, sort_keys=True) + "\n") def parameter_norm(parameters: list[torch.nn.Parameter]) -> float: total = 0.0 for parameter in parameters: total += float(parameter.detach().float().square().sum()) return math.sqrt(total) def gradient_norm(parameters: list[torch.nn.Parameter]) -> float: total = 0.0 for parameter in parameters: if parameter.grad is not None: total += float(parameter.grad.detach().float().square().sum()) return math.sqrt(total) def build_shared_nonblock_state(seed: int) -> dict[str, dict[str, torch.Tensor]]: set_seed(seed) fusion = TripleFeatureFusion(1536) residual = torch.nn.Linear(1536, 1536) torch.nn.init.zeros_(residual.weight) torch.nn.init.zeros_(residual.bias) return { "fusion": { key: value.detach().clone() for key, value in fusion.state_dict().items() }, "residual_out": { key: value.detach().clone() for key, value in residual.state_dict().items() }, } def load_teacher( checkpoint_path: Path, config_path: Path, device: torch.device, ) -> torch.nn.Module: config = OmegaConf.merge( OmegaConf.load(REPO_ROOT / "configs/default_config.yaml"), OmegaConf.load(config_path), ) wrapper = WanDiffusionWrapper( **getattr(config, "model_kwargs", {}), is_causal=True, ) checkpoint = torch.load( checkpoint_path, map_location="cpu", weights_only=False, mmap=True, ) wrapper.load_state_dict(checkpoint["generator_ema"], strict=True) del checkpoint wrapper.to(device=device, dtype=torch.bfloat16) wrapper.eval().requires_grad_(False) model = wrapper.model if len(model.blocks) != 30 or model.dim != 1536: raise ValueError( f"Expected Wan 1.3B with 30×1536 blocks, got " f"{len(model.blocks)}×{model.dim}" ) if model.freqs.device != device: model.freqs = model.freqs.to(device) return model class BatchSchedule: """Precompute identical prompt/group batches for every initialization.""" def __init__( self, prompt_ids: list[int], batch_size: int, steps: int, seed: int, ) -> None: self.entries: list[tuple[int, int, list[int]]] = [] rng = random.Random(seed) groups = [ (chunk, target_step) for chunk in range(1, 7) for target_step in range(1, 4) ] while len(self.entries) < steps: epoch_groups = groups.copy() rng.shuffle(epoch_groups) for chunk, target_step in epoch_groups: prompts = rng.sample(prompt_ids, batch_size) self.entries.append((chunk, target_step, prompts)) if len(self.entries) == steps: break def fingerprint(self) -> str: payload = json.dumps(self.entries, separators=(",", ":")).encode() return hashlib.sha256(payload).hexdigest() def lr_values( step: int, max_steps: int, fusion_lr: float, block_lr: float, fusion_warmup_steps: int, block_freeze_steps: int, block_warmup_steps: int, ) -> tuple[float, float]: if step < fusion_warmup_steps: fusion_factor = float(step + 1) / max(1, fusion_warmup_steps) else: progress = (step - fusion_warmup_steps) / max( 1, max_steps - fusion_warmup_steps ) fusion_factor = 0.5 * (1.0 + math.cos(math.pi * min(1.0, progress))) if step < block_freeze_steps: block_factor = 0.0 elif step < block_freeze_steps + block_warmup_steps: block_factor = float(step - block_freeze_steps + 1) / max( 1, block_warmup_steps ) else: progress = ( step - block_freeze_steps - block_warmup_steps ) / max(1, max_steps - block_freeze_steps - block_warmup_steps) block_factor = 0.5 * ( 1.0 + math.cos(math.pi * min(1.0, progress)) ) return fusion_lr * fusion_factor, block_lr * block_factor def move_batch( batch: dict[str, Any], device: torch.device, ) -> dict[str, Any]: output = { key: value.to( device=device, dtype=( torch.bfloat16 if value.is_floating_point() and key != "timestep" else value.dtype ), non_blocking=False, ) for key, value in batch.items() if isinstance(value, torch.Tensor) } output.update( { "prompt_ids": batch["prompt_ids"], "chunk": batch["chunk"], "anchor_step": batch["anchor_step"], "target_step": batch["target_step"], } ) return output def frozen_inputs( batch: dict[str, Any], teacher: torch.nn.Module, device: torch.device, ) -> dict[str, torch.Tensor]: noisy = batch["noisy_latent"] batch_size = noisy.shape[0] with torch.no_grad(), torch.autocast( device_type="cuda", dtype=torch.bfloat16 ): current_tokens = teacher.patch_embedding( noisy.permute(0, 2, 1, 3, 4) ).flatten(2).transpose(1, 2) timestep = batch["timestep"] time_embedding = teacher.time_embedding( sinusoidal_embedding_1d( teacher.freq_dim, timestep.flatten() ).type_as(current_tokens) ) timestep_modulation = teacher.time_projection( time_embedding ).unflatten(1, (6, teacher.dim)).unflatten( dim=0, sizes=timestep.shape ) head_embedding = time_embedding.unflatten( dim=0, sizes=timestep.shape ).unsqueeze(2) condition_per_frame = time_embedding.unflatten( dim=0, sizes=timestep.shape ) tokens_per_frame = 30 * 52 condition_tokens = ( condition_per_frame[:, :, None, :] .expand(batch_size, timestep.shape[1], tokens_per_frame, teacher.dim) .reshape(batch_size, -1, teacher.dim) ) grid_sizes = torch.tensor( [[3, 30, 52]] * batch_size, dtype=torch.long, device="cpu", ) return { "current_tokens": current_tokens, "timestep_modulation": timestep_modulation, "head_embedding": head_embedding, "condition_tokens": condition_tokens, "grid_sizes": grid_sizes, "freqs": teacher.freqs, } def hidden_to_flow( pred_hidden: torch.Tensor, head_embedding: torch.Tensor, grid_sizes: torch.Tensor, teacher: torch.nn.Module, ) -> torch.Tensor: with torch.autocast(device_type="cuda", dtype=torch.bfloat16): head_tokens = teacher.head(pred_hidden, head_embedding) flow_channels_first = torch.stack( teacher.unpatchify(head_tokens, grid_sizes) ) return flow_channels_first.permute(0, 2, 1, 3, 4) def forward_predictor( model: SingleBlockPredictor, batch: dict[str, Any], teacher: torch.nn.Module, device: torch.device, ) -> tuple[torch.Tensor, torch.Tensor]: frozen = frozen_inputs(batch, teacher, device) with torch.autocast(device_type="cuda", dtype=torch.bfloat16): pred_hidden = model( current_tokens=frozen["current_tokens"], anchor_hidden=batch["anchor_hidden"], previous_hidden=batch["previous_hidden"], timestep_modulation=frozen["timestep_modulation"], grid_sizes=frozen["grid_sizes"], freqs=frozen["freqs"], history_k=batch["history_k"], history_v=batch["history_v"], cross_k=batch["cross_k"], cross_v=batch["cross_v"], current_start=batch["chunk"] * TOKENS_PER_CHUNK, condition_tokens=frozen["condition_tokens"], anchor_distance=batch.get("anchor_distance"), ) pred_flow = hidden_to_flow( pred_hidden, frozen["head_embedding"], frozen["grid_sizes"], teacher, ) return pred_hidden, pred_flow @torch.inference_mode() def evaluate( model: SingleBlockPredictor, store: OfflinePredictorStore, val_prompt_ids: list[int], batch_size: int, teacher: torch.nn.Module, device: torch.device, hidden_weight: float, flow_weight: float, ) -> dict[str, float]: model.eval() hidden_squared = 0.0 hidden_elements = 0 flow_squared = 0.0 flow_elements = 0 gate_sum = 0.0 gate_square_sum = 0.0 gate_count = 0 gate_histogram = torch.zeros(10, dtype=torch.int64) started = time.perf_counter() for chunk in range(1, 7): for target_step in range(1, 4): for start in range(0, len(val_prompt_ids), batch_size): prompt_ids = val_prompt_ids[start : start + batch_size] batch = move_batch( store.batch(prompt_ids, chunk, target_step), device ) pred_hidden, pred_flow = forward_predictor( model, batch, teacher, device ) gate = model.fusion.last_gate if gate is None: raise RuntimeError("Fusion did not expose gate values") gate_float = gate.float() gate_sum += float(gate_float.sum()) gate_square_sum += float(gate_float.square().sum()) gate_count += gate_float.numel() gate_histogram += torch.histc( gate_float, bins=10, min=0.0, max=1.0 ).to(device="cpu", dtype=torch.int64) hidden_error = pred_hidden.float() - batch["target_hidden"].float() flow_error = pred_flow.float() - batch["target_flow"].float() hidden_squared += float(hidden_error.square().sum()) hidden_elements += hidden_error.numel() flow_squared += float(flow_error.square().sum()) flow_elements += flow_error.numel() del batch, pred_hidden, pred_flow, hidden_error, flow_error hidden_mse = hidden_squared / hidden_elements flow_mse = flow_squared / flow_elements gate_mean = gate_sum / gate_count gate_variance = max(0.0, gate_square_sum / gate_count - gate_mean**2) model.train() return { "hidden_mse": hidden_mse, "flow_mse": flow_mse, "total_loss": hidden_weight * hidden_mse + flow_weight * flow_mse, "gate_mean": gate_mean, "gate_std": math.sqrt(gate_variance), "gate_histogram_10_bins": gate_histogram.tolist(), "gate_count": gate_count, "eval_time_s": time.perf_counter() - started, } def normalized_auc( evaluations: list[dict[str, Any]], key: str, max_steps: int ) -> float: if len(evaluations) < 2: return float(evaluations[0][key]) area = 0.0 for left, right in zip(evaluations, evaluations[1:]): width = int(right["step"]) - int(left["step"]) area += width * (float(left[key]) + float(right[key])) * 0.5 return area / max_steps def save_predictor_weights( model: SingleBlockPredictor, path: Path ) -> None: tensors = { key: value.detach().to(device="cpu").contiguous() for key, value in model.state_dict().items() } temporary = path.with_suffix(path.suffix + ".tmp") save_file(tensors, temporary) os.replace(temporary, path) def make_model( teacher: torch.nn.Module, source_layer: int, method: str, shared_nonblock_state: dict[str, dict[str, torch.Tensor]], seed: int, gradient_checkpointing: bool, device: torch.device, gate_mode: str = "baseline", gate_hidden_dim: int = 128, gate_initial_bias: float = 4.6, gate_floor: float = 0.0, constant_gate: float = 1.0, input_variant: str = "self_forcing", atc_previous_scope: str = "chunk", atc_freq_dim: int = 256, atc_mlp_hidden_dim: int = 3072, atc_gate_hidden_dim: int = 512, atc_transport_residual_scale: float = 0.1, atc_gate_initial_probability: float = 0.3, ) -> SingleBlockPredictor: # Random block baselines use the same seed independently of the shared # fusion initialization. set_seed(seed) block = initialize_predictor_block( teacher.blocks[source_layer], method, ) model = SingleBlockPredictor( block, dim=1536, gradient_checkpointing=gradient_checkpointing, input_variant=input_variant, gate_mode=gate_mode, gate_hidden_dim=gate_hidden_dim, gate_initial_bias=gate_initial_bias, gate_floor=gate_floor, constant_gate=constant_gate, atc_previous_scope=atc_previous_scope, atc_freq_dim=atc_freq_dim, atc_mlp_hidden_dim=atc_mlp_hidden_dim, atc_gate_hidden_dim=atc_gate_hidden_dim, atc_transport_residual_scale=atc_transport_residual_scale, atc_gate_initial_probability=atc_gate_initial_probability, ) fusion_state = shared_nonblock_state["fusion"] if input_variant == "disca": # Keep all overlapping initialization exactly matched to the # Self-Forcing baseline, while physically removing the third D-wide # previous-chunk input channel from proj_in. disca_state = { key: value for key, value in fusion_state.items() if key.startswith(("current_norm.", "anchor_norm.", "proj_out.")) } disca_state["proj_in.weight"] = fusion_state["proj_in.weight"][ :, : 2 * teacher.dim ].clone() disca_state["proj_in.bias"] = fusion_state["proj_in.bias"].clone() model.fusion.load_state_dict(disca_state, strict=True) elif input_variant == "self_forcing": incompatible = model.fusion.load_state_dict(fusion_state, strict=False) unexpected = list(incompatible.unexpected_keys) missing = [ key for key in incompatible.missing_keys if not key.startswith("gate.") ] if unexpected or missing: raise RuntimeError( f"Shared fusion state mismatch: missing={missing}, " f"unexpected={unexpected}" ) elif input_variant != "atc": raise ValueError(f"Unknown input variant: {input_variant}") model.residual_out.load_state_dict( shared_nonblock_state["residual_out"], strict=True ) return model.to(device=device) def run_experiment( *, name: str, method: str, source_layer: int, args: argparse.Namespace, store: OfflinePredictorStore, teacher: torch.nn.Module, train_prompt_ids: list[int], val_prompt_ids: list[int], schedule: BatchSchedule, shared_nonblock_state: dict[str, dict[str, torch.Tensor]], device: torch.device, ) -> dict[str, Any]: run_dir = args.output_dir / name metrics_path = run_dir / "metrics.json" if metrics_path.exists(): existing = json.loads(metrics_path.read_text(encoding="utf-8")) if existing.get("status") == "complete": print(f"[run] skip completed {name}", flush=True) return existing run_dir.mkdir(parents=True, exist_ok=True) log_path = run_dir / "train_log.jsonl" model = make_model( teacher, source_layer, method, shared_nonblock_state, args.seed, args.gradient_checkpointing, device, args.gate_mode, args.gate_hidden_dim, args.gate_initial_bias, args.gate_floor, args.constant_gate, ) fusion_parameters = model.fusion_parameters_without_gate() gate_parameters = model.gate_parameters() block_parameters = model.block_parameters() optimizer_groups = [ {"params": fusion_parameters, "lr": args.fusion_lr}, ] if gate_parameters: optimizer_groups.append({"params": gate_parameters, "lr": args.gate_lr}) block_group_index = len(optimizer_groups) optimizer_groups.append({"params": block_parameters, "lr": args.block_lr}) optimizer = AdamW( optimizer_groups, betas=(0.9, 0.95), weight_decay=args.weight_decay, ) initial_block_norm = parameter_norm(block_parameters) evaluations: list[dict[str, Any]] = [] start_step = 0 latest_path = run_dir / "training_latest.pt" if latest_path.exists(): state = torch.load( latest_path, map_location="cpu", weights_only=False ) model.load_state_dict(state["model"], strict=True) optimizer.load_state_dict(state["optimizer"]) evaluations = state["evaluations"] start_step = int(state["step"]) print(f"[run] resume {name} at step {start_step}", flush=True) model.set_block_trainable(start_step >= args.block_freeze_steps) config = { "name": name, "initialization_method": method, "source_layer": source_layer, "source_definition": ( "block weights, clean-history K/V projection/input, and text K/V " "all use this generator_ema Teacher layer" ), "seed": args.seed, "train_prompt_ids": train_prompt_ids, "val_prompt_ids": val_prompt_ids, "max_steps": args.max_steps, "batch_size": args.batch_size, "eval_batch_size": args.eval_batch_size, "fusion_lr": args.fusion_lr, "block_lr": args.block_lr, "fusion_warmup_steps": args.fusion_warmup_steps, "block_freeze_steps": args.block_freeze_steps, "block_warmup_steps": args.block_warmup_steps, "weight_decay": args.weight_decay, "hidden_weight": args.hidden_weight, "flow_weight": args.flow_weight, "gradient_checkpointing": args.gradient_checkpointing, "gate_mode": args.gate_mode, "gate_hidden_dim": args.gate_hidden_dim, "gate_initial_bias": args.gate_initial_bias, "gate_floor": args.gate_floor, "gate_lr": args.gate_lr, "gate_freeze_steps": args.gate_freeze_steps, "constant_gate": args.constant_gate, "batch_schedule_sha256": schedule.fingerprint(), "initial_block_parameter_norm": initial_block_norm, "trainable_parameters": sum( parameter.numel() for parameter in model.parameters() ), } atomic_json(run_dir / "config.json", config) print( f"[run] {name}: method={method}, source={source_layer}, " f"start={start_step}", flush=True, ) if not evaluations: initial_eval = evaluate( model, store, val_prompt_ids, args.eval_batch_size, teacher, device, args.hidden_weight, args.flow_weight, ) evaluations.append({"step": 0, **initial_eval}) print( f"[eval] {name} step=0 flow={initial_eval['flow_mse']:.8f} " f"hidden={initial_eval['hidden_mse']:.8f}", flush=True, ) optimizer.zero_grad(set_to_none=True) model.train() run_started = time.perf_counter() for step in range(start_step, args.max_steps): block_enabled = step >= args.block_freeze_steps if any(parameter.requires_grad != block_enabled for parameter in block_parameters): model.set_block_trainable(block_enabled) fusion_lr, block_lr = lr_values( step, args.max_steps, args.fusion_lr, args.block_lr, args.fusion_warmup_steps, args.block_freeze_steps, args.block_warmup_steps, ) optimizer.param_groups[0]["lr"] = fusion_lr if gate_parameters: gate_lr, _ = lr_values( max(0, step - args.gate_freeze_steps), max(1, args.max_steps - args.gate_freeze_steps), args.gate_lr, args.block_lr, args.fusion_warmup_steps, 0, args.block_warmup_steps, ) if step < args.gate_freeze_steps: gate_lr = 0.0 optimizer.param_groups[1]["lr"] = gate_lr else: gate_lr = 0.0 optimizer.param_groups[block_group_index]["lr"] = block_lr chunk, target_step, prompt_ids = schedule.entries[step] batch = move_batch( store.batch(prompt_ids, chunk, target_step), device ) step_started = time.perf_counter() pred_hidden, pred_flow = forward_predictor( model, batch, teacher, device ) hidden_loss = F.mse_loss( pred_hidden.float(), batch["target_hidden"].float() ) flow_loss = F.mse_loss( pred_flow.float(), batch["target_flow"].float() ) loss = ( args.hidden_weight * hidden_loss + args.flow_weight * flow_loss ) loss.backward() if step < args.gate_freeze_steps: for parameter in gate_parameters: parameter.grad = None fusion_grad_norm = gradient_norm(fusion_parameters) block_grad_norm = gradient_norm(block_parameters) total_grad_norm = torch.nn.utils.clip_grad_norm_( model.parameters(), args.grad_clip ) optimizer.step() optimizer.zero_grad(set_to_none=True) completed_step = step + 1 if completed_step == 1 or completed_step % args.log_every == 0: torch.cuda.synchronize() record = { "step": completed_step, "train_total_loss": float(loss.detach()), "train_hidden_mse": float(hidden_loss.detach()), "train_flow_mse": float(flow_loss.detach()), "fusion_grad_norm": fusion_grad_norm, "block_grad_norm": block_grad_norm, "total_grad_norm_before_clip": float(total_grad_norm), "fusion_lr": fusion_lr, "gate_lr": gate_lr, "block_lr": block_lr, "chunk": chunk, "target_step": target_step, "step_time_s": time.perf_counter() - step_started, "peak_gpu_gib": torch.cuda.max_memory_allocated() / (1024**3), } append_jsonl(log_path, record) print( f"[train] {name} {completed_step}/{args.max_steps} " f"flow={record['train_flow_mse']:.8f} " f"time={record['step_time_s']:.2f}s " f"mem={record['peak_gpu_gib']:.1f}G", flush=True, ) should_eval = ( completed_step % args.eval_every == 0 or completed_step == args.max_steps ) if should_eval: validation = evaluate( model, store, val_prompt_ids, args.eval_batch_size, teacher, device, args.hidden_weight, args.flow_weight, ) evaluations.append({"step": completed_step, **validation}) print( f"[eval] {name} step={completed_step} " f"flow={validation['flow_mse']:.8f} " f"hidden={validation['hidden_mse']:.8f}", flush=True, ) should_save = ( completed_step % args.save_every == 0 or completed_step == args.max_steps ) if should_save: temporary = latest_path.with_suffix(".pt.tmp") torch.save( { "model": { key: value.detach().cpu() for key, value in model.state_dict().items() }, "optimizer": optimizer.state_dict(), "evaluations": evaluations, "step": completed_step, }, temporary, ) os.replace(temporary, latest_path) del ( batch, pred_hidden, pred_flow, hidden_loss, flow_loss, loss, total_grad_norm, ) final = evaluations[-1] result = { "status": "complete", **config, "final_val_hidden_mse": final["hidden_mse"], "final_val_flow_mse": final["flow_mse"], "final_val_total_loss": final["total_loss"], "val_hidden_mse_auc": normalized_auc( evaluations, "hidden_mse", args.max_steps ), "val_flow_mse_auc": normalized_auc( evaluations, "flow_mse", args.max_steps ), "val_total_loss_auc": normalized_auc( evaluations, "total_loss", args.max_steps ), "evaluations": evaluations, "training_time_s": time.perf_counter() - run_started, "final_block_parameter_norm": parameter_norm(block_parameters), } if args.save_final_weights: save_predictor_weights(model, run_dir / "predictor_final.safetensors") atomic_json(metrics_path, result) if latest_path.exists(): latest_path.unlink() del model, optimizer torch.cuda.empty_cache() return result def write_summary(output_dir: Path, results: list[dict[str, Any]]) -> None: rows = [ { "name": result["name"], "initialization_method": result["initialization_method"], "source_layer": result["source_layer"], "final_val_flow_mse": result["final_val_flow_mse"], "final_val_hidden_mse": result["final_val_hidden_mse"], "final_val_total_loss": result["final_val_total_loss"], "val_flow_mse_auc": result["val_flow_mse_auc"], "val_hidden_mse_auc": result["val_hidden_mse_auc"], "val_total_loss_auc": result["val_total_loss_auc"], "training_time_s": result["training_time_s"], } for result in results ] rows.sort(key=lambda row: float(row["final_val_flow_mse"])) atomic_json(output_dir / "summary.json", rows) temporary = output_dir / "summary.csv.tmp" with temporary.open("w", newline="", encoding="utf-8") as handle: writer = csv.DictWriter(handle, fieldnames=list(rows[0])) writer.writeheader() writer.writerows(rows) os.replace(temporary, output_dir / "summary.csv") def main() -> None: args = parse_args() args.dataset_root = resolve(args.dataset_root) args.checkpoint_path = resolve(args.checkpoint_path) args.config_path = resolve(args.config_path) args.output_dir = resolve(args.output_dir) args.output_dir.mkdir(parents=True, exist_ok=True) layers = ( list(range(30)) if args.teacher_layers is None or len(args.teacher_layers) == 0 else sorted(set(args.teacher_layers)) ) if any(layer < 0 or layer >= 30 for layer in layers): raise ValueError(f"Invalid Teacher layers: {layers}") set_seed(args.seed) torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = True torch.set_float32_matmul_precision("high") device = torch.device("cuda") train_prompt_ids = list(range(args.train_prompts)) val_prompt_ids = list( range(80, 80 + args.val_prompts) if args.train_prompts == 80 else range(args.train_prompts, args.train_prompts + args.val_prompts) ) prompt_ids = sorted(set(train_prompt_ids + val_prompt_ids)) schedule = BatchSchedule( train_prompt_ids, args.batch_size, args.max_steps, args.seed, ) shared_nonblock_state = build_shared_nonblock_state(args.seed) atomic_json( args.output_dir / "sweep_config.json", { **{ key: str(value) if isinstance(value, Path) else value for key, value in vars(args).items() }, "teacher_layers": layers, "train_prompt_ids": train_prompt_ids, "val_prompt_ids": val_prompt_ids, "batch_schedule_sha256": schedule.fingerprint(), "other_initializations": list(OTHER_METHODS), }, ) print("[setup] loading frozen generator_ema", flush=True) teacher = load_teacher( args.checkpoint_path, args.config_path, device, ) print("[setup] loading common offline trajectories into RAM", flush=True) store = OfflinePredictorStore(args.dataset_root, prompt_ids) results: list[dict[str, Any]] = [] for layer in layers: name = f"{args.gate_mode}_layer_{layer:02d}" metrics_path = args.output_dir / name / "metrics.json" if metrics_path.exists(): existing = json.loads(metrics_path.read_text(encoding="utf-8")) if existing.get("status") == "complete": results.append(existing) print(f"[sweep] already complete: {name}", flush=True) continue store.load_layer_cache(layer, teacher, device) results.append( run_experiment( name=name, method="teacher_full", source_layer=layer, args=args, store=store, teacher=teacher, train_prompt_ids=train_prompt_ids, val_prompt_ids=val_prompt_ids, schedule=schedule, shared_nonblock_state=shared_nonblock_state, device=device, ) ) write_summary(args.output_dir, results) teacher_results = [ result for result in results if result["initialization_method"] == "teacher_full" ] if not teacher_results: raise RuntimeError("No Teacher layer experiment completed") best_teacher = min( teacher_results, key=lambda result: float(result["final_val_flow_mse"]), ) best_layer = int(best_teacher["source_layer"]) atomic_json( args.output_dir / "best_teacher_layer.json", { "source_layer": best_layer, "selection_metric": "final_val_flow_mse", "value": best_teacher["final_val_flow_mse"], "run": best_teacher["name"], }, ) print( f"[sweep] best Teacher layer={best_layer}, " f"flow={best_teacher['final_val_flow_mse']:.8f}", flush=True, ) if args.run_other_initializations: store.load_layer_cache(best_layer, teacher, device) for method in OTHER_METHODS: name = f"{method}_source_{best_layer:02d}" results.append( run_experiment( name=name, method=method, source_layer=best_layer, args=args, store=store, teacher=teacher, train_prompt_ids=train_prompt_ids, val_prompt_ids=val_prompt_ids, schedule=schedule, shared_nonblock_state=shared_nonblock_state, device=device, ) ) write_summary(args.output_dir, results) write_summary(args.output_dir, results) print( f"[sweep] complete: {len(results)} experiments -> " f"{args.output_dir / 'summary.csv'}", flush=True, ) if __name__ == "__main__": main()