#!/usr/bin/env python3 """Probe one real frozen-Predictor + Confidence-token-head optimizer step.""" from __future__ import annotations import argparse import json import os import sys import time import traceback from pathlib import Path from typing import Any def _preparse_gpu() -> str: parser = argparse.ArgumentParser(add_help=False) parser.add_argument("--gpu", required=True) 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 import safe_open from safetensors.torch import load_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.confidence import ConfidenceTokenHead from predictor_training.lazy_offline_data import ( LazyLayer17Dataset, collate_lazy_samples, ) from predictor_training.offline_data import TOKENS_PER_CHUNK from predictor_training.single_block import ( SingleBlockPredictor, initialize_predictor_block, ) from scripts.run_single_block_init_sweep import frozen_inputs, load_teacher from scripts.train_layer17_stage1_lazy_ddp import move_training_batch from utils.misc import set_seed def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--gpu", default=PHYSICAL_GPU) parser.add_argument("--batch_size", type=int, required=True) parser.add_argument("--prompt_start", type=int, default=0) parser.add_argument("--chunk", type=int, default=6) parser.add_argument("--target_step", type=int, default=3) parser.add_argument( "--dataset_root", type=Path, default=REPO_ROOT / "offline_training_datasets" / "predictor_offline_layer17_1000p_21f_seed0_no_chunk0", ) parser.add_argument( "--predictor_weights", type=Path, default=REPO_ROOT / "training_runs" / "layer17_atc_chunk_stage1_1000p_4gpu_b16_2000steps" / "checkpoint_step_2000" / "predictor.safetensors", ) parser.add_argument( "--predictor_input_variant", choices=("auto", "self_forcing", "disca", "atc"), default="auto", ) parser.add_argument( "--checkpoint_path", type=Path, default=REPO_ROOT / "checkpoints/self_forcing_dmd.pt", ) parser.add_argument( "--config_path", type=Path, default=REPO_ROOT / "configs/self_forcing_sid.yaml", ) parser.add_argument("--output", type=Path, required=True) parser.add_argument("--learning_rate", type=float, default=3e-4) parser.add_argument("--weight_decay", type=float, default=0.01) parser.add_argument("--seed", type=int, default=0) args = parser.parse_args() if args.batch_size < 1: parser.error("--batch_size must be positive") if args.prompt_start < 0 or args.prompt_start + args.batch_size > 900: parser.error("probe prompts must stay in the training split 0..899") if not 1 <= args.chunk <= 6: parser.error("--chunk must be in 1..6") if not 1 <= args.target_step <= 3: parser.error("--target_step must be in 1..3") for name in ( "dataset_root", "predictor_weights", "checkpoint_path", "config_path", "output", ): path = getattr(args, name).expanduser() setattr( args, name, path.resolve() if path.is_absolute() else (REPO_ROOT / path).resolve(), ) return args def atomic_json(path: Path, value: dict[str, 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, sort_keys=True) + "\n", encoding="utf-8" ) os.replace(temporary, path) def read_predictor_config( path: Path, expected_input_variant: str = "auto" ) -> dict[str, Any]: with safe_open(path, framework="pt", device="cpu") as handle: metadata = handle.metadata() or {} raw = metadata.get("predictor_config") if raw is None: if expected_input_variant != "self_forcing": raise ValueError( f"Missing predictor_config metadata: {path}; legacy concat " "checkpoints require --predictor_input_variant self_forcing" ) return { "source_layer": 17, "input_variant": "self_forcing", "gate_mode": "baseline", "metadata_source": "explicit_legacy_concat_override", } config = json.loads(raw) actual = str(config.get("input_variant", "self_forcing")) if expected_input_variant != "auto" and actual != expected_input_variant: raise ValueError( f"Predictor input variant mismatch: expected={expected_input_variant} " f"actual={actual}" ) return config def load_predictor( teacher: torch.nn.Module, weights: Path, device: torch.device, expected_input_variant: str = "auto", ) -> tuple[SingleBlockPredictor, dict[str, Any]]: config = read_predictor_config(weights, expected_input_variant) source_layer = int(config.get("source_layer", 17)) predictor = SingleBlockPredictor( block=initialize_predictor_block( teacher.blocks[source_layer], "teacher_full" ), dim=teacher.dim, gradient_checkpointing=False, input_variant=str(config.get("input_variant", "self_forcing")), atc_previous_scope=config.get("atc_previous_scope", "chunk"), atc_freq_dim=int(config.get("atc_freq_dim", 256)), atc_mlp_hidden_dim=int(config.get("atc_mlp_hidden_dim", 3072)), atc_gate_hidden_dim=int(config.get("atc_gate_hidden_dim", 512)), atc_transport_residual_scale=float( config.get("atc_transport_residual_scale", 0.1) ), atc_gate_initial_probability=float( config.get("atc_gate_initial_probability", 0.3) ), atc_collect_diagnostics=False, ) predictor.load_state_dict(load_file(str(weights), device="cpu"), strict=True) predictor.to(device=device, dtype=torch.bfloat16) predictor.eval().requires_grad_(False) return predictor, config @torch.no_grad() def extract_predictor_features( predictor: 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): output = predictor( 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, return_features=True, condition_tokens=frozen["condition_tokens"], anchor_distance=batch["anchor_distance"], ) if not isinstance(output, tuple): raise RuntimeError("Predictor did not return (pred_hidden, transformed)") return output def hidden_nrmse(predicted: torch.Tensor, target: torch.Tensor) -> torch.Tensor: error_energy = (predicted.float() - target.float()).square().sum(dim=(1, 2)) target_energy = target.float().square().sum(dim=(1, 2)) return torch.sqrt(error_energy / target_energy.clamp_min(1e-8)) def memory_gib(value: int) -> float: return value / 2**30 def main() -> None: args = parse_args() started = time.perf_counter() result: dict[str, Any] = { "success": False, "physical_gpu": str(args.gpu), "batch_size": args.batch_size, "prompt_ids": [args.prompt_start, args.prompt_start + args.batch_size - 1], "split": {"train": "0..899", "validation": "900..999"}, "chunk": args.chunk, "target_step": args.target_step, "predictor_weights": str(args.predictor_weights), } try: torch.cuda.set_device(0) device = torch.device("cuda", 0) set_seed(args.seed) torch.set_num_threads(4) torch.set_num_interop_threads(1) torch.backends.cuda.matmul.allow_tf32 = True torch.set_float32_matmul_precision("high") print(f"[probe] gpu={args.gpu} batch={args.batch_size} loading teacher", flush=True) teacher = load_teacher(args.checkpoint_path, args.config_path, device) predictor, predictor_config = load_predictor( teacher, args.predictor_weights, device, expected_input_variant=args.predictor_input_variant, ) result["predictor_config"] = predictor_config head = ConfidenceTokenHead(num_steps=3, dropout=0.1).to(device=device) result["head_parameters"] = sum( parameter.numel() for parameter in head.parameters() ) optimizer = AdamW( head.parameters(), lr=args.learning_rate, betas=(0.9, 0.95), weight_decay=args.weight_decay, ) print(f"[probe] gpu={args.gpu} batch={args.batch_size} loading samples", flush=True) dataset = LazyLayer17Dataset(args.dataset_root, layer_id=17) cpu_batch = collate_lazy_samples( [ dataset[(prompt_id, args.chunk, args.target_step)] for prompt_id in range( args.prompt_start, args.prompt_start + args.batch_size ) ] ) torch.cuda.reset_peak_memory_stats() batch = move_training_batch(cpu_batch, teacher, device) del cpu_batch torch.cuda.synchronize() feature_started = time.perf_counter() pred_hidden, transformed = extract_predictor_features( predictor, batch, teacher, device ) target_log = torch.log( hidden_nrmse(pred_hidden, batch["target_hidden"]) + 1e-6 ).detach() torch.cuda.synchronize() result["predictor_forward_s"] = time.perf_counter() - feature_started chunk_position = torch.full( (args.batch_size,), (args.chunk - 1) / 5.0, dtype=torch.float32, device=device, ) step_id = torch.full( (args.batch_size,), args.target_step, dtype=torch.long, device=device, ) head_started = time.perf_counter() optimizer.zero_grad(set_to_none=True) with torch.autocast(device_type="cuda", dtype=torch.bfloat16): predicted_log = head( transformed_hidden=transformed, pred_hidden=pred_hidden, anchor_hidden=batch["anchor_hidden"], chunk_position=chunk_position, step_id=step_id, ) loss = F.smooth_l1_loss(predicted_log, target_log) loss.backward() grad_norm = torch.nn.utils.clip_grad_norm_(head.parameters(), 1.0) optimizer.step() torch.cuda.synchronize() free_bytes, total_bytes = torch.cuda.mem_get_info() result.update( success=True, loss=float(loss.detach()), grad_norm=float(grad_norm), head_step_s=time.perf_counter() - head_started, peak_allocated_gib=memory_gib(torch.cuda.max_memory_allocated()), peak_reserved_gib=memory_gib(torch.cuda.max_memory_reserved()), final_free_gib=memory_gib(free_bytes), gpu_total_gib=memory_gib(total_bytes), elapsed_s=time.perf_counter() - started, ) print( f"[probe] PASS gpu={args.gpu} batch={args.batch_size} " f"peak={result['peak_allocated_gib']:.2f}GiB " f"reserved={result['peak_reserved_gib']:.2f}GiB", flush=True, ) except torch.cuda.OutOfMemoryError as error: result.update( error_type="CUDAOutOfMemoryError", error=str(error), peak_allocated_gib=memory_gib(torch.cuda.max_memory_allocated()), peak_reserved_gib=memory_gib(torch.cuda.max_memory_reserved()), elapsed_s=time.perf_counter() - started, ) print(f"[probe] OOM gpu={args.gpu} batch={args.batch_size}: {error}", flush=True) except Exception as error: result.update( error_type=type(error).__name__, error=str(error), traceback=traceback.format_exc(), elapsed_s=time.perf_counter() - started, ) atomic_json(args.output, result) raise atomic_json(args.output, result) if not result["success"]: raise SystemExit(2) if __name__ == "__main__": main()