Download scripts/probe_confidence_token_batch_size.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 13.2 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/probe_confidence_token_batch_size.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/scripts/probe_confidence_token_batch_size.py
-
curl -L -o probe_confidence_token_batch_size.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/scripts/probe_confidence_token_batch_size.py
13.2 kB
| #!/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 | |
| 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() | |