#!/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()