#!/usr/bin/env python3 """Train two-block Predictors over mirror and consecutive Teacher pairs.""" from __future__ import annotations import argparse import csv import json import os 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() 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 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.two_block import TwoBlockPredictor from scripts.run_single_block_init_sweep import ( BatchSchedule, append_jsonl, atomic_json, build_shared_nonblock_state, frozen_inputs, gradient_norm, hidden_to_flow, load_teacher, lr_values, move_batch, normalized_auc, parameter_norm, ) from utils.misc import set_seed def default_pairs() -> list[tuple[int, int]]: mirror = [(index, 29 - index) for index in range(15)] consecutive = [(index, index + 1) for index in range(29)] return list(dict.fromkeys(mirror + consecutive)) def pair_kind(pair: tuple[int, int]) -> str: mirror = pair[0] + pair[1] == 29 consecutive = pair[1] == pair[0] + 1 if mirror and consecutive: return "mirror+consecutive" if mirror: return "mirror" if consecutive: return "consecutive" return "custom" def parse_pair(value: str) -> tuple[int, int]: try: left_text, right_text = value.split(",", maxsplit=1) pair = (int(left_text), int(right_text)) except Exception as error: raise argparse.ArgumentTypeError( f"Pair must look like 0,29; got {value!r}" ) from error if not (0 <= pair[0] < pair[1] < 30): raise argparse.ArgumentTypeError( f"Pair must satisfy 0 <= first < second < 30; got {pair}" ) return pair 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/two_block_pair_sweep") ) parser.add_argument( "--pairs", type=parse_pair, nargs="*", default=None, help="Pairs such as 0,29 1,28. Omit for all mirror+consecutive pairs.", ) parser.add_argument("--max_pairs", type=int, default=None) 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("--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( "--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("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 return args def resolve(path: Path) -> Path: path = path.expanduser() return path.resolve() if path.is_absolute() else (REPO_ROOT / path).resolve() def make_model( teacher: torch.nn.Module, pair: tuple[int, int], shared_nonblock_state: dict[str, dict[str, torch.Tensor]], seed: int, gradient_checkpointing: bool, device: torch.device, ) -> TwoBlockPredictor: set_seed(seed) model = TwoBlockPredictor( [teacher.blocks[pair[0]], teacher.blocks[pair[1]]], dim=teacher.dim, gradient_checkpointing=gradient_checkpointing, ) model.fusion.load_state_dict(shared_nonblock_state["fusion"], strict=True) model.residual_out.load_state_dict( shared_nonblock_state["residual_out"], strict=True ) return model.to(device=device) def forward_predictor( model: TwoBlockPredictor, 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_ks=[batch["history_k_0"], batch["history_k_1"]], history_vs=[batch["history_v_0"], batch["history_v_1"]], cross_ks=[batch["cross_k_0"], batch["cross_k_1"]], cross_vs=[batch["cross_v_0"], batch["cross_v_1"]], current_start=batch["chunk"] * TOKENS_PER_CHUNK, ) 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: TwoBlockPredictor, store: OfflinePredictorStore, pair: tuple[int, int], 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 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_layers(prompt_ids, chunk, target_step, pair), device ) pred_hidden, pred_flow = forward_predictor( model, batch, teacher, device ) 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 model.train() return { "hidden_mse": hidden_mse, "flow_mse": flow_mse, "total_loss": hidden_weight * hidden_mse + flow_weight * flow_mse, "eval_time_s": time.perf_counter() - started, } def save_predictor_weights(model: TwoBlockPredictor, 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 run_experiment( *, pair: tuple[int, 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]: name = f"pair_{pair[0]:02d}_{pair[1]:02d}" 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, pair, shared_nonblock_state, args.seed, args.gradient_checkpointing, device, ) fusion_parameters = model.fusion_parameters() block_parameters = model.block_parameters() optimizer = AdamW( [ {"params": fusion_parameters, "lr": args.fusion_lr}, {"params": block_parameters, "lr": args.block_lr}, ], 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_blocks_trainable(start_step >= args.block_freeze_steps) config = { "name": name, "architecture": "two_block_predictor", "initialization_method": "teacher_full", "source_layers": list(pair), "pair_kind": pair_kind(pair), "source_definition": ( "Predictor block i, its clean-history K/V, and its text K/V use " "generator_ema source_layers[i]" ), "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, "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}: kind={config['pair_kind']} start={start_step}", flush=True ) if not evaluations: initial_eval = evaluate( model, store, pair, 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_blocks_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 optimizer.param_groups[1]["lr"] = block_lr chunk, target_step, prompt_ids = schedule.entries[step] batch = move_batch( store.batch_layers(prompt_ids, chunk, target_step, pair), 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() fusion_grad_norm = gradient_norm(fusion_parameters) block_grad_norm = gradient_norm(block_parameters) block_0_grad_norm = gradient_norm(list(model.blocks[0].parameters())) block_1_grad_norm = gradient_norm(list(model.blocks[1].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, "block_0_grad_norm": block_0_grad_norm, "block_1_grad_norm": block_1_grad_norm, "total_grad_norm_before_clip": float(total_grad_norm), "fusion_lr": fusion_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, ) if completed_step % args.eval_every == 0 or completed_step == args.max_steps: validation = evaluate( model, store, pair, 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, ) if completed_step % args.save_every == 0 or completed_step == args.max_steps: 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 del 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"], "source_layer_1": result["source_layers"][0], "source_layer_2": result["source_layers"][1], "pair_kind": result["pair_kind"], "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"])) if not rows: return destination = output_dir / "summary.csv" temporary = destination.with_suffix(".csv.tmp") with temporary.open("w", encoding="utf-8", newline="") as handle: writer = csv.DictWriter(handle, fieldnames=list(rows[0])) writer.writeheader() writer.writerows(rows) os.replace(temporary, destination) atomic_json(output_dir / "summary.json", rows) 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) pairs = default_pairs() if args.pairs is None else list(dict.fromkeys(args.pairs)) if args.max_pairs is not None: pairs = pairs[: args.max_pairs] train_prompt_ids = list(range(args.train_prompts)) val_prompt_ids = list( range(args.train_prompts, args.train_prompts + args.val_prompts) ) all_prompt_ids = train_prompt_ids + val_prompt_ids device = torch.device("cuda") set_seed(args.seed) torch.set_grad_enabled(True) 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, all_prompt_ids) schedule = BatchSchedule( train_prompt_ids, args.batch_size, args.max_steps, args.seed ) shared_nonblock_state = build_shared_nonblock_state(args.seed) manifest = { "status": "running", "pairs": [list(pair) for pair in pairs], "num_pairs": len(pairs), "mirror_pairs": 15, "consecutive_pairs": 29, "deduplicated_overlap": [14, 15], "max_steps": args.max_steps, "seed": args.seed, "train_prompt_ids": train_prompt_ids, "val_prompt_ids": val_prompt_ids, "batch_schedule_sha256": schedule.fingerprint(), } atomic_json(args.output_dir / "sweep_manifest.json", manifest) results: list[dict[str, Any]] = [] for index, pair in enumerate(pairs, start=1): name = f"pair_{pair[0]:02d}_{pair[1]: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": print(f"[sweep] {index}/{len(pairs)} skip {name}", flush=True) results.append(existing) continue print( f"[data] {index}/{len(pairs)} loading pair {pair[0]},{pair[1]}", flush=True, ) store.load_layer_caches(pair, teacher, device) result = run_experiment( pair=pair, 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, ) results.append(result) write_summary(args.output_dir, results) manifest["status"] = "complete" atomic_json(args.output_dir / "sweep_manifest.json", manifest) write_summary(args.output_dir, results) print( f"[complete] {len(results)} pairs -> {args.output_dir / 'summary.csv'}", flush=True, ) if __name__ == "__main__": main()