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