Download scripts/run_three_block_consecutive_sweep.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 23.6 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/run_three_block_consecutive_sweep.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/scripts/run_three_block_consecutive_sweep.py
-
curl -L -o run_three_block_consecutive_sweep.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/run_three_block_consecutive_sweep.py
23.6 kB
| #!/usr/bin/env python3 | |
| """Train three-block Predictors initialized from consecutive Teacher layers.""" | |
| # ruff: noqa: E402 -- CUDA_VISIBLE_DEVICES must be set before importing torch. | |
| 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.three_block import ThreeBlockPredictor | |
| 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 | |
| Triple = tuple[int, int, int] | |
| def default_triples() -> list[Triple]: | |
| return [(index, index + 1, index + 2) for index in range(28)] | |
| def parse_triple(value: str) -> Triple: | |
| try: | |
| values = tuple(int(item) for item in value.split(",")) | |
| except ValueError as error: | |
| raise argparse.ArgumentTypeError( | |
| f"Triple must look like 0,1,2; got {value!r}" | |
| ) from error | |
| if len(values) != 3: | |
| raise argparse.ArgumentTypeError( | |
| f"Triple must contain exactly three layers; got {value!r}" | |
| ) | |
| triple = (values[0], values[1], values[2]) | |
| if not ( | |
| 0 <= triple[0] | |
| and triple[1] == triple[0] + 1 | |
| and triple[2] == triple[1] + 1 | |
| and triple[2] < 30 | |
| ): | |
| raise argparse.ArgumentTypeError( | |
| "Triple must be three consecutive layers within 0..29; " | |
| f"got {triple}" | |
| ) | |
| return triple | |
| 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/three_block_consecutive_sweep"), | |
| ) | |
| parser.add_argument( | |
| "--triples", | |
| type=parse_triple, | |
| nargs="*", | |
| default=None, | |
| help=( | |
| "Optional subset such as 0,1,2 1,2,3. " | |
| "Omit to train all 28 consecutive triples." | |
| ), | |
| ) | |
| parser.add_argument("--max_triples", 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 | |
| if args.max_triples is not None and args.max_triples < 1: | |
| parser.error("--max_triples must be positive") | |
| return args | |
| def resolve(path: Path) -> Path: | |
| path = path.expanduser() | |
| return path.resolve() if path.is_absolute() else (REPO_ROOT / path).resolve() | |
| def experiment_name(triple: Triple) -> str: | |
| return "triple_" + "_".join(f"{layer:02d}" for layer in triple) | |
| def make_model( | |
| teacher: torch.nn.Module, | |
| triple: Triple, | |
| shared_nonblock_state: dict[str, dict[str, torch.Tensor]], | |
| seed: int, | |
| gradient_checkpointing: bool, | |
| device: torch.device, | |
| ) -> ThreeBlockPredictor: | |
| set_seed(seed) | |
| model = ThreeBlockPredictor( | |
| [teacher.blocks[layer] for layer in triple], | |
| 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: ThreeBlockPredictor, | |
| 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[f"history_k_{index}"] for index in range(3)], | |
| history_vs=[batch[f"history_v_{index}"] for index in range(3)], | |
| cross_ks=[batch[f"cross_k_{index}"] for index in range(3)], | |
| cross_vs=[batch[f"cross_v_{index}"] for index in range(3)], | |
| 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 | |
| def evaluate( | |
| model: ThreeBlockPredictor, | |
| store: OfflinePredictorStore, | |
| triple: Triple, | |
| 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, triple), | |
| 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: ThreeBlockPredictor, 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( | |
| *, | |
| triple: Triple, | |
| 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 = experiment_name(triple) | |
| 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, | |
| triple, | |
| 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": "three_block_predictor", | |
| "initialization_method": "teacher_full", | |
| "source_layers": list(triple), | |
| "triple_kind": "consecutive", | |
| "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}: start={start_step}", flush=True) | |
| if not evaluations: | |
| initial_eval = evaluate( | |
| model, | |
| store, | |
| triple, | |
| 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, triple), 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) | |
| per_block_grad_norms = [ | |
| gradient_norm(list(block.parameters())) for block in model.blocks | |
| ] | |
| 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, | |
| "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), | |
| } | |
| record.update( | |
| { | |
| f"block_{index}_grad_norm": value | |
| for index, value in enumerate(per_block_grad_norms) | |
| } | |
| ) | |
| 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, | |
| triple, | |
| 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], | |
| "source_layer_3": result["source_layers"][2], | |
| "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) | |
| triples = ( | |
| default_triples() | |
| if args.triples is None | |
| else list(dict.fromkeys(args.triples)) | |
| ) | |
| if args.max_triples is not None: | |
| triples = triples[: args.max_triples] | |
| 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", | |
| "architecture": "three_block_predictor", | |
| "initialization_method": "teacher_full", | |
| "triples": [list(triple) for triple in triples], | |
| "num_triples": len(triples), | |
| "default_consecutive_triples": 28, | |
| "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, triple in enumerate(triples, start=1): | |
| name = experiment_name(triple) | |
| 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(triples)} skip {name}", flush=True) | |
| results.append(existing) | |
| continue | |
| print( | |
| f"[data] {index}/{len(triples)} loading layers {triple}", flush=True | |
| ) | |
| store.load_layer_caches(triple, teacher, device) | |
| result = run_experiment( | |
| triple=triple, | |
| 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)} triples -> {args.output_dir / 'summary.csv'}", | |
| flush=True, | |
| ) | |
| if __name__ == "__main__": | |
| main() | |