Self-Forcing / scripts /run_single_block_init_sweep.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw History Blame Contribute Delete
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
@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()