Self-Forcing / scripts /build_conditional_probe_dataset.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw
History Blame Contribute Delete
9.52 kB
#!/usr/bin/env python3
"""Normalize three-model hidden snapshots into one offline probe dataset.
The recorder implementations live in their respective projects. This script
only converts their per-prompt snapshots to a small common CPU format; it does
not run a model or manufacture control examples.
"""
from __future__ import annotations
import argparse
import json
import os
import sys
from pathlib import Path
from typing import Any
def _preparse_gpu() -> str:
parser = argparse.ArgumentParser(add_help=False)
parser.add_argument("--gpu", default="0")
args, _ = parser.parse_known_args()
os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu)
return str(args.gpu)
PHYSICAL_GPU = _preparse_gpu()
import numpy as np
import torch
ROLES = {
"self_forcing": {7: "early", 14: "middle", 22: "late", 29: "final"},
"causal_forcing": {7: "early", 14: "middle", 22: "late", 29: "final"},
"hy_worldplay": {13: "early", 26: "middle", 40: "late", 53: "final"},
}
def regular_coords(frames: int = 3, height: int = 30, width: int = 52, max_tokens: int = 240):
total = frames * height * width
if total <= max_tokens:
flat = np.arange(total, dtype=np.int64)
else:
per_frame = max(1, max_tokens // frames)
h_count = min(height, max(1, int(round((per_frame * height / width) ** 0.5))))
w_count = min(width, max(1, per_frame // h_count))
while frames * h_count * w_count > max_tokens and w_count > 1:
w_count -= 1
while frames * h_count * w_count > max_tokens and h_count > 1:
h_count -= 1
hs = np.unique(np.rint(np.linspace(0, height - 1, h_count)).astype(np.int64))
ws = np.unique(np.rint(np.linspace(0, width - 1, w_count)).astype(np.int64))
flat = np.asarray(
[t * height * width + h * width + w for t in range(frames) for h in hs for w in ws],
dtype=np.int64,
)
t = flat // (height * width)
rem = flat % (height * width)
return np.stack([t, rem // width, rem % width], axis=1)
def ensure_stack(values: dict[tuple[int, int], torch.Tensor], layer: int, chunks: int, steps: int):
rows = []
for chunk in range(chunks):
step_rows = []
for step in range(steps):
key = (chunk, step)
if key not in values:
raise ValueError(f"Missing layer={layer} chunk={chunk} step={step}")
step_rows.append(values[key].detach().cpu().to(torch.float16))
rows.append(torch.stack(step_rows, dim=0))
return torch.stack(rows, dim=0).contiguous()
def load_self(path: Path, layers: list[int], chunks: int, steps: int) -> dict[str, Any]:
run = torch.load(path, map_location="cpu", weights_only=False)
features = {}
for layer in layers:
stage = f"block_{layer}_hidden"
values = {}
for key, value in run["records"][stage].items():
c, s = (int(part) for part in key.split(":"))
if c < chunks and s < steps:
values[(c, s)] = value
features[ROLES["self_forcing"][layer]] = ensure_stack(values, layer, chunks, steps)
return {
"prompt_id": int(run["run_index"]),
"prompt": run["prompt"],
"seed": int(run["seed"]),
"model_family": "self_forcing",
"model_variant": "dmd4",
"features": features,
"timesteps": np.asarray([1000.0, 937.5, 833.3333, 625.0], dtype=np.float32),
"coords": regular_coords(),
}
def load_causal(path: Path, layers: list[int], chunks: int, steps: int) -> dict[str, Any]:
run = torch.load(path, map_location="cpu", weights_only=False)
raw = {}
for key, value in run["features"].items():
layer, chunk, step = (int(part) for part in key.split(":"))
if layer in layers and chunk < chunks and step < steps:
raw.setdefault(layer, {})[(chunk, step)] = value
features = {
ROLES["causal_forcing"][layer]: ensure_stack(raw.get(layer, {}), layer, chunks, steps)
for layer in layers
}
return {
"prompt_id": int(run["prompt_id"]),
"prompt": run["prompt"],
"seed": int(run["seed"]),
"model_family": "causal_forcing",
"model_variant": "dmd4",
"features": features,
"timesteps": np.asarray([1000.0, 937.5, 833.3333, 625.0], dtype=np.float32),
"coords": regular_coords(),
}
def load_hy(path: Path, layers: list[int], chunks: int, steps: int) -> dict[str, Any]:
data = np.load(path, allow_pickle=False)
stages = [str(value) for value in data["stages"]]
raw = {}
for index, stage in enumerate(stages):
if not stage.startswith("block_"):
continue
layer = int(stage.split("_")[-1])
chunk = int(data["chunks"][index])
step = int(data["steps"][index])
if layer in layers and chunk < chunks and step < steps:
raw.setdefault(layer, {})[(chunk, step)] = torch.from_numpy(data["features"][index])
features = {
ROLES["hy_worldplay"][layer]: ensure_stack(raw.get(layer, {}), layer, chunks, steps)
for layer in layers
}
return {
"features": features,
"timesteps": np.asarray(data["timesteps"], dtype=np.float32),
"coords": np.asarray(data["coords"], dtype=np.int64),
}
def atomic_save(path: Path, value: dict[str, Any]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
temporary = path.with_suffix(path.suffix + ".tmp")
torch.save(value, temporary)
os.replace(temporary, path)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--self_root", type=Path, required=True)
parser.add_argument("--causal_root", type=Path, required=True)
parser.add_argument("--hy_root", type=Path, required=True)
parser.add_argument("--output_root", type=Path, required=True)
parser.add_argument("--chunks", type=int, default=4)
parser.add_argument("--steps", type=int, default=4)
parser.add_argument("--max_prompts", type=int, default=10)
return parser.parse_args()
def main() -> None:
args = parse_args()
args.output_root.mkdir(parents=True, exist_ok=True)
specs = {
"self_forcing": ([7, 14, 22, 29], args.self_root / "runs", "self"),
"causal_forcing": ([7, 14, 22, 29], args.causal_root / "runs", "causal"),
}
inventory = []
for family, (layers, run_root, prefix) in specs.items():
out_dir = args.output_root / family
out_dir.mkdir(parents=True, exist_ok=True)
for prompt_id in range(args.max_prompts):
if family == "self_forcing":
source = run_root / f"prompt_{prompt_id:02d}.pt"
if not source.exists():
raise FileNotFoundError(source)
item = load_self(source, layers, args.chunks, args.steps)
else:
source = run_root / f"prompt_{prompt_id:04d}" / "feature_snapshots.pt"
if not source.exists():
raise FileNotFoundError(source)
item = load_causal(source, layers, args.chunks, args.steps)
destination = out_dir / f"prompt_{prompt_id:04d}.pt"
atomic_save(destination, item)
inventory.append({
"family": family,
"prompt_id": prompt_id,
"path": str(destination),
"bytes": destination.stat().st_size,
"roles": sorted(item["features"]),
})
hy_files = sorted(args.hy_root.glob("shard_gpu*/runs/prompt_*/forward/final_hidden_snapshots.npz"))
hy_by_prompt = {}
for source in hy_files:
prompt_id = int(source.parts[-3].split("_")[-1])
if prompt_id < args.max_prompts:
hy_by_prompt[prompt_id] = source
out_dir = args.output_root / "hy_worldplay"
out_dir.mkdir(parents=True, exist_ok=True)
for prompt_id in range(args.max_prompts):
source = hy_by_prompt.get(prompt_id)
if source is None:
raise FileNotFoundError(f"HY snapshot for prompt {prompt_id}")
item = load_hy(source, [13, 26, 40, 53], args.chunks, args.steps)
item.update({
"prompt_id": prompt_id,
"model_family": "hy_worldplay",
"model_variant": "ar4",
"seed": 0,
"prompt": f"prompt_{prompt_id:04d}",
})
destination = out_dir / f"prompt_{prompt_id:04d}.pt"
atomic_save(destination, item)
inventory.append({
"family": "hy_worldplay",
"prompt_id": prompt_id,
"path": str(destination),
"bytes": destination.stat().st_size,
"roles": sorted(item["features"]),
})
manifest = {
"dataset_version": 1,
"prompt_ids": list(range(args.max_prompts)),
"chunks": args.chunks,
"steps": args.steps,
"max_tokens": 240,
"roles": ["early", "middle", "late", "final"],
"source_roots": {
"self_forcing": str(args.self_root),
"causal_forcing": str(args.causal_root),
"hy_worldplay": str(args.hy_root),
},
"inventory": inventory,
}
(args.output_root / "manifest.json").write_text(
json.dumps(manifest, indent=2, ensure_ascii=False) + "\n", encoding="utf-8"
)
print(f"[complete] {args.output_root} prompts={args.max_prompts} files={len(inventory)}", flush=True)
if __name__ == "__main__":
main()