torch-dimensions / viewer /make_samples.py
Celsia's picture
Upload folder using huggingface_hub
ecc81b3 verified
Raw
History Blame Contribute Delete
2.98 kB
"""Regenerate the viewer's bundled sample specs.
python viewer/make_samples.py
The samples used to be artifacts with no recipe: someone produced them once,
they were committed, and when the spec format moved they silently described a
library version that no longer existed. Now they are a build step, and the
models below are the source of truth for what the viewer ships with.
Each sample is chosen to exercise a different rendering path — sparse cells,
a paired schedule, dimensional stacking at rank 4, and the kernel family's
simultaneous contraction (which has no travelling wavefront at all).
"""
from __future__ import annotations
import json
from pathlib import Path
import torch
import torch_dimensions as td
OUT = Path(__file__).parent / "src" / "samples"
def sparse(shape: tuple[int, ...], keep: float, seed: int = 0) -> torch.Tensor:
g = torch.Generator().manual_seed(seed)
valid = torch.rand(shape, generator=g) < keep
valid.reshape(-1)[0] = True
return valid
def lstm_2d_sparse():
lat = td.Lattice(shape=(6, 8), names=("h", "w"), valid=sparse((6, 8), 0.75), time=True)
return td.LSTM(32, 6, lat, d_input=1, bidirectional=("h", "w"))
def mamba_3d():
lat = td.Lattice(shape=(4, 5, 6), names=("depth", "row", "col"), time=True)
plan = td.ScanPlan.paired(lat.axis_names, n_layers=12, bidirectional=("depth", "row", "col"))
return td.Mamba(48, lattice=lat, plan=plan, d_input=1, d_state=16)
def s4d_4d():
lat = td.Lattice(
shape=(3, 4, 5, 4),
names=("depth", "row", "col", "group"),
valid=sparse((3, 4, 5, 4), 0.7, seed=3),
time=True,
)
return td.S4D(32, 8, lat, d_input=1, d_state=32)
def cafa_hybrid():
"""The kernel family: every layer contracts both spatial axes at once and
sweeps only time. There is no per-axis wavefront to draw, and until the
spec described families honestly the viewer drew one anyway."""
lat = td.Lattice(
shape=(5, 7), names=("station", "sensor"), valid=sparse((5, 7), 0.8, seed=1), time=True
)
return td.LSTM(32, 5, lat, d_input=1, method=td.cafa)
def vit_joint():
"""The joint family: a ViT's patch grid, where every cell is one token and
no axis is swept. Nothing travels, so the scene pulses instead."""
return td.ViT(96, 6, image=(32, 32), patch=4, in_channels=3, n_heads=4, names=("row", "col"))
SAMPLES = {
"lstm_2d_sparse": lstm_2d_sparse,
"mamba_3d": mamba_3d,
"s4d_4d": s4d_4d,
"cafa_hybrid": cafa_hybrid,
"vit_joint": vit_joint,
}
def main() -> None:
OUT.mkdir(parents=True, exist_ok=True)
for name, build in SAMPLES.items():
torch.manual_seed(0)
spec = td.spec(build())
(OUT / f"{name}.json").write_text(json.dumps(spec, indent=1, sort_keys=True) + "\n")
n = spec["model"]["n_params"]
print(f"{name:16s} v{spec['version']} {spec['nd_method']['family']:7s} {n:,} params")
if __name__ == "__main__":
main()