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