Download scripts/run_single_block_init_sweep.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 37.8 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/run_single_block_init_sweep.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/scripts/run_single_block_init_sweep.py
-
curl -L -o run_single_block_init_sweep.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/run_single_block_init_sweep.py
37.8 kB
| #!/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 | |
| 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() | |