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