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