Self-Forcing / scripts /build_boundary_conditional_dataset.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw History Blame Contribute Delete
5.5 kB
#!/usr/bin/env python3
"""Assemble the native-feature boundary-to-all conditional-probe dataset.
Self-Forcing and HY-static already use structured temporal/spatial sampling and
are linked read-only. Causal-Forcing is normalized from a fresh recorder run
whose features are stored as [temporal, spatial, channel].
"""
from __future__ import annotations
import argparse
import json
import os
from pathlib import Path
import numpy as np
import torch
ROLE_MAP = {7: "early", 14: "middle", 22: "late", 29: "final"}
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--self_dir", type=Path, required=True)
parser.add_argument("--causal_root", type=Path, required=True)
parser.add_argument("--hy_dir", type=Path, required=True)
parser.add_argument("--output_root", type=Path, required=True)
parser.add_argument("--num_prompts", type=int, default=10)
parser.add_argument("--chunks", type=int, default=4)
parser.add_argument("--steps", type=int, default=4)
return parser.parse_args()
def ensure_link(path: Path, target: Path) -> None:
target = target.resolve()
if path.is_symlink():
if path.resolve() != target:
raise ValueError(f"Existing symlink {path} points to {path.resolve()}, not {target}")
return
if path.exists():
raise FileExistsError(path)
path.symlink_to(target, target_is_directory=True)
def atomic_save(path: Path, value: dict) -> None:
temporary = path.with_suffix(path.suffix + ".tmp")
torch.save(value, temporary)
os.replace(temporary, path)
def coords_from_indices(indices: torch.Tensor) -> np.ndarray:
flat = indices.detach().cpu().numpy().astype(np.int64).reshape(-1)
plane = 30 * 52
temporal = flat // plane
remainder = flat % plane
coords = np.stack([temporal, remainder // 52, remainder % 52], axis=1)
slots = sorted(int(value) for value in np.unique(temporal))
if slots != [0, 1, 2]:
raise ValueError(f"Expected three temporal slots, got {slots}")
reference = coords[coords[:, 0] == slots[-1], 1:]
for slot in slots:
if not np.array_equal(coords[coords[:, 0] == slot, 1:], reference):
raise ValueError(f"Causal slot {slot} does not share the structured spatial grid")
return coords
def normalize_causal(path: Path, chunks: int, steps: int) -> dict:
state = torch.load(path, map_location="cpu", weights_only=False)
features = {}
for layer, role in ROLE_MAP.items():
chunk_rows = []
for chunk in range(chunks):
step_rows = []
for step in range(steps):
value = state["features"][f"{layer}:{chunk}:{step}"]
step_rows.append(value.reshape(-1, value.shape[-1]).to(torch.float16))
chunk_rows.append(torch.stack(step_rows, dim=0))
features[role] = torch.stack(chunk_rows, dim=0).contiguous()
index_key = f"{next(iter(ROLE_MAP))}:0:0"
coords = coords_from_indices(state["feature_indices"][index_key])
token_counts = {int(value.shape[2]) for value in features.values()}
if token_counts != {len(coords)}:
raise ValueError(f"Feature/coordinate mismatch: tokens={token_counts}, coords={len(coords)}")
return {
"prompt_id": int(state["prompt_id"]),
"prompt": state["prompt"],
"seed": int(state["seed"]),
"model_family": "causal_forcing",
"model_variant": "dmd4_boundary_structured",
"features": features,
"timesteps": np.asarray([1000.0, 937.5, 833.3333, 625.0], dtype=np.float32),
"coords": coords,
"grid_shape": np.asarray([3, 30, 52], dtype=np.int64),
}
def main() -> None:
args = parse_args()
output = args.output_root.resolve()
output.mkdir(parents=True, exist_ok=True)
ensure_link(output / "self_forcing", args.self_dir)
ensure_link(output / "hy_worldplay", args.hy_dir)
causal_output = output / "causal_forcing"
causal_output.mkdir(exist_ok=True)
inventory = []
for prompt_id in range(args.num_prompts):
source = (
args.causal_root
/ "runs"
/ f"prompt_{prompt_id:04d}"
/ "feature_snapshots.pt"
)
if not source.exists():
raise FileNotFoundError(source)
item = normalize_causal(source, args.chunks, args.steps)
if item["prompt_id"] != prompt_id:
raise ValueError(f"Prompt mismatch in {source}: {item['prompt_id']}")
destination = causal_output / f"prompt_{prompt_id:04d}.pt"
atomic_save(destination, item)
inventory.append({
"prompt_id": prompt_id,
"source": str(source),
"destination": str(destination),
"tokens": int(item["features"]["early"].shape[2]),
})
manifest = {
"dataset_version": 1,
"chunk_pairing": "boundary_to_all",
"prompt_ids": list(range(args.num_prompts)),
"chunks": args.chunks,
"steps": args.steps,
"self_forcing": str(args.self_dir.resolve()),
"hy_worldplay": str(args.hy_dir.resolve()),
"causal_source": str(args.causal_root.resolve()),
"causal_inventory": inventory,
}
(output / "manifest.json").write_text(
json.dumps(manifest, indent=2, ensure_ascii=False) + "\n",
encoding="utf-8",
)
print(f"[complete] {output} prompts={args.num_prompts}", flush=True)
if __name__ == "__main__":
main()