Add files using upload-large-folder tool
Browse files- code/models/experimental/gr00t/benchmarks/bench_denoise.py +11 -2
- code/models/experimental/gr00t/benchmarks/bench_e2e.py +110 -8
- code/models/experimental/gr00t/benchmarks/bench_load.py +236 -0
- code/models/experimental/gr00t/benchmarks/bench_mk_step.py +291 -0
- code/models/experimental/gr00t/benchmarks/sweep_backbone.py +2 -1
- code/models/experimental/gr00t/tests/tt/test_mk_2cq.py +384 -0
- code/models/experimental/gr00t/tests/tt/test_mk_cpu.py +6 -2
- code/models/experimental/gr00t/tests/tt/test_mk_dit_block.py +398 -0
- code/models/experimental/gr00t/tests/tt/test_mk_dit_step.py +335 -0
- code/models/experimental/gr00t/tt/action_head.py +392 -51
- code/models/experimental/gr00t/tt/device.py +12 -4
- code/models/experimental/gr00t/tt/model.py +256 -25
- code/models/experimental/gr00t/tt/policy.py +70 -3
code/models/experimental/gr00t/benchmarks/bench_denoise.py
CHANGED
|
@@ -225,7 +225,11 @@ def table_estimate_us(cfg: Any, matmul_rows: Sequence[Mapping[str, Any]]) -> Dic
|
|
| 225 |
|
| 226 |
def analyse(version: str, policy: str, layout_name: Optional[str], embodiment: Optional[str]) -> Dict[str, Any]:
|
| 227 |
"""CPU analysis shared by the dry run and the device run."""
|
| 228 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 229 |
cfg, layout = setup["cfg"], setup["layout"]
|
| 230 |
n_steps = cfg.sampler.num_inference_timesteps
|
| 231 |
entries = trace_b_entries(version, setup["embodiment_id"], setup["policy"])
|
|
@@ -382,7 +386,12 @@ def run(
|
|
| 382 |
raise ValueError("warmup must be >= 0 and reps > 0")
|
| 383 |
an = analyse(version, policy, layout_name, embodiment)
|
| 384 |
setup = E2E.resolve_setup(
|
| 385 |
-
version,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 386 |
)
|
| 387 |
ok, reason = E2E.golden_available(version, sample)
|
| 388 |
if not ok:
|
|
|
|
| 225 |
|
| 226 |
def analyse(version: str, policy: str, layout_name: Optional[str], embodiment: Optional[str]) -> Dict[str, Any]:
|
| 227 |
"""CPU analysis shared by the dry run and the device run."""
|
| 228 |
+
# this benchmark measures the Stage-1 (ttnn) trace B -- its streamed bytes and its replay time; the megakernel
|
| 229 |
+
# backend (the TTPolicy default since 2026-09-18) has no separate denoise op sequence to count
|
| 230 |
+
setup = E2E.resolve_setup(
|
| 231 |
+
version, policy=policy, layout_name=layout_name, embodiment=embodiment, dit_backend="ttnn"
|
| 232 |
+
)
|
| 233 |
cfg, layout = setup["cfg"], setup["layout"]
|
| 234 |
n_steps = cfg.sampler.num_inference_timesteps
|
| 235 |
entries = trace_b_entries(version, setup["embodiment_id"], setup["policy"])
|
|
|
|
| 386 |
raise ValueError("warmup must be >= 0 and reps > 0")
|
| 387 |
an = analyse(version, policy, layout_name, embodiment)
|
| 388 |
setup = E2E.resolve_setup(
|
| 389 |
+
version,
|
| 390 |
+
policy=policy,
|
| 391 |
+
trace_layout=trace_layout,
|
| 392 |
+
layout_name=layout_name,
|
| 393 |
+
embodiment=embodiment,
|
| 394 |
+
dit_backend="ttnn", # the Stage-1 denoise trace (see analyse); the device is the default-L1 one
|
| 395 |
)
|
| 396 |
ok, reason = E2E.golden_available(version, sample)
|
| 397 |
if not ok:
|
code/models/experimental/gr00t/benchmarks/bench_e2e.py
CHANGED
|
@@ -20,7 +20,12 @@ copy ms) is reported as ``upload/<input>`` rows and ``upload/total`` (docs/plan/
|
|
| 20 |
|
| 21 |
``--num-cqs`` (default 2) opens the device with that many command queues; ``--cq1-uploads auto|on|off`` and
|
| 22 |
``--embed-gather device|host`` are forwarded to ``Gr00tTT``; ``--profile-encode`` adds the per-step host-encode
|
| 23 |
-
profile (``meta.encode_profile``; also in ``--dry-run``).
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 24 |
|
| 25 |
The model comes from ``tt.model.Gr00tTT`` (WP-14). Until that module exists the import is guarded: ``--dry-run``
|
| 26 |
validates ``--version`` / ``--policy`` / ``--trace-layout`` / ``--layout`` / ``--embodiment`` against ``common.configs``
|
|
@@ -61,6 +66,7 @@ from typing import Any, Callable, Dict, List, Mapping, Optional, Sequence, Tuple
|
|
| 61 |
from models.experimental.gr00t.benchmarks import bench_stage_ops as SO
|
| 62 |
from models.experimental.gr00t.benchmarks import mb1_common as MB
|
| 63 |
from models.experimental.gr00t.tt import device as D
|
|
|
|
| 64 |
|
| 65 |
BENCH_NAME = "bench_e2e"
|
| 66 |
VERSIONS = SO.VERSIONS
|
|
@@ -88,6 +94,8 @@ def resolve_setup(
|
|
| 88 |
trace_layout: str = "per_stage",
|
| 89 |
layout_name: Optional[str] = None,
|
| 90 |
embodiment: Optional[str] = None,
|
|
|
|
|
|
|
| 91 |
) -> Dict[str, Any]:
|
| 92 |
"""Validate the CLI choices against ``common.configs`` / ``tt.policy`` and return the resolved facts
|
| 93 |
(``cfg``, ``layout``, ``policy`` object, ``embodiment_tag`` / ``embodiment_id``, expected trace names)."""
|
|
@@ -102,7 +110,9 @@ def resolve_setup(
|
|
| 102 |
layout = cfg.canonical_layout if layout_name is None else cfg.layout(layout_name)
|
| 103 |
tag = embodiment if embodiment is not None else layout.embodiment_tag
|
| 104 |
emb_id = cfg.embodiment_id(tag) # raises on an unknown tag
|
| 105 |
-
pol = TTPolicy(
|
|
|
|
|
|
|
| 106 |
return {
|
| 107 |
"cfg": cfg,
|
| 108 |
"layout": layout,
|
|
@@ -121,6 +131,8 @@ def resolve_setup(
|
|
| 121 |
"dit_m_pad": layout.dit_m_pad,
|
| 122 |
"noise_shape": list(layout.noise_shape),
|
| 123 |
"policy": pol.to_dict(),
|
|
|
|
|
|
|
| 124 |
"expected_stage_names": list(STAGE_NAMES[trace_layout]),
|
| 125 |
},
|
| 126 |
}
|
|
@@ -190,6 +202,7 @@ def load_model(
|
|
| 190 |
*,
|
| 191 |
embed_gather: str = "device",
|
| 192 |
cq1_uploads: Optional[bool] = None,
|
|
|
|
| 193 |
) -> Any:
|
| 194 |
Gr00tTT = import_gr00t_tt()
|
| 195 |
if Gr00tTT is None:
|
|
@@ -207,6 +220,8 @@ def load_model(
|
|
| 207 |
embed_gather=embed_gather,
|
| 208 |
cq1_uploads=cq1_uploads,
|
| 209 |
)
|
|
|
|
|
|
|
| 210 |
if cache_root is not None:
|
| 211 |
kw["cache_root"] = cache_root
|
| 212 |
return Gr00tTT.from_pretrained(setup["cfg"].version, **kw)
|
|
@@ -235,6 +250,15 @@ def resolve_stage_replayer(model: Any, expected_stages: Sequence[str]) -> Any:
|
|
| 235 |
return tracer
|
| 236 |
|
| 237 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 238 |
def denoise_stage_name(stage_names: Sequence[str]) -> str:
|
| 239 |
"""The trace that holds the 4-step denoise: the last one of either layout (``denoise`` / ``B``)."""
|
| 240 |
if not stage_names:
|
|
@@ -406,6 +430,9 @@ def run(
|
|
| 406 |
embed_gather: str = "device",
|
| 407 |
cq1_uploads: Optional[bool] = None,
|
| 408 |
profile_encode: bool = False,
|
|
|
|
|
|
|
|
|
|
| 409 |
) -> Dict[str, Any]:
|
| 410 |
"""Measure the traced end-to-end latency and its split on an open device; returns the perf payload (not written).
|
| 411 |
Main pass = production path (``tracer.call``), split pass = serialised writes + per-trace sync (module docstring).
|
|
@@ -418,7 +445,13 @@ def run(
|
|
| 418 |
if warmup < 0 or reps <= 0:
|
| 419 |
raise ValueError("warmup must be >= 0 and reps > 0")
|
| 420 |
setup = resolve_setup(
|
| 421 |
-
version,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 422 |
)
|
| 423 |
ok, reason = golden_available(version, sample)
|
| 424 |
if not ok:
|
|
@@ -427,7 +460,14 @@ def run(
|
|
| 427 |
obs, mi = inputs["obs"], inputs["mi"]
|
| 428 |
noise = make_noise(setup["layout"], noise_seed)
|
| 429 |
t_load = time.perf_counter()
|
| 430 |
-
model = load_model(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 431 |
load_s = time.perf_counter() - t_load
|
| 432 |
tracer = resolve_stage_replayer(model, setup["stage_names"])
|
| 433 |
missing = [m for m in PRODUCTION_MEMBERS if not hasattr(tracer, m)]
|
|
@@ -510,6 +550,19 @@ def run(
|
|
| 510 |
per_call_inputs=[s.describe() for s in executor.plan().inputs],
|
| 511 |
per_call_bytes=ex_desc.get("plan", {}).get("per_call_bytes"),
|
| 512 |
vit_fork=getattr(getattr(model.backbone, "tower", None), "fork", None),
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 513 |
)
|
| 514 |
if profile_encode:
|
| 515 |
meta["encode_profile"] = encode_profile(version, inputs, setup["layout"])
|
|
@@ -530,6 +583,9 @@ def run(
|
|
| 530 |
"noise_seed": noise_seed,
|
| 531 |
"embed_gather": embed_gather,
|
| 532 |
"cq1_uploads": cq1_uploads,
|
|
|
|
|
|
|
|
|
|
| 533 |
},
|
| 534 |
method={
|
| 535 |
"main_pass": f"{warmup} warm-up calls, then {reps} calls: encode -> tracer.call (TraceExecutor.replay: "
|
|
@@ -586,10 +642,18 @@ def dry_run(
|
|
| 586 |
host_reps: int = DEFAULT_HOST_REPS,
|
| 587 |
verbose: bool = True,
|
| 588 |
profile_encode: bool = False,
|
|
|
|
|
|
|
| 589 |
) -> Dict[str, Any]:
|
| 590 |
"""Validate the setup and time the host side (encode / decode) when the golden data is present; no device."""
|
| 591 |
setup = resolve_setup(
|
| 592 |
-
version,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 593 |
)
|
| 594 |
stages: Dict[str, Dict[str, Any]] = {}
|
| 595 |
meta = dict(setup["meta"])
|
|
@@ -621,6 +685,8 @@ def dry_run(
|
|
| 621 |
"embodiment": setup["embodiment_tag"],
|
| 622 |
"layout": setup["layout"].name,
|
| 623 |
"host_reps": host_reps,
|
|
|
|
|
|
|
| 624 |
},
|
| 625 |
method={"host": f"{host_reps} host calls after 1 warm-up (encode / decode), median / p90 ms; no device"},
|
| 626 |
stages=stages,
|
|
@@ -655,13 +721,15 @@ def print_summary(payload: Mapping[str, Any]) -> None:
|
|
| 655 |
)
|
| 656 |
m = payload["meta"]
|
| 657 |
print(
|
| 658 |
-
" layout {l} | embodiment {e} | S_pad {s} | shape_key {sk} | policy {p} |
|
| 659 |
-
"cq1_uploads {c} ({n} CQ){skip}".format(
|
| 660 |
l=m.get("layout_name"),
|
| 661 |
e=m.get("embodiment_tag"),
|
| 662 |
s=m.get("s_pad_llm"),
|
| 663 |
sk=m.get("shape_key", "-"),
|
| 664 |
p=(m.get("policy") or {}).get("dtype_policy"),
|
|
|
|
|
|
|
| 665 |
st=m.get("stage_names", m.get("expected_stage_names")),
|
| 666 |
g=m.get("embed_gather", "-"),
|
| 667 |
c=m.get("cq1_uploads", "-"),
|
|
@@ -680,6 +748,24 @@ def build_parser() -> argparse.ArgumentParser:
|
|
| 680 |
)
|
| 681 |
ap.add_argument("--version", required=True, choices=VERSIONS)
|
| 682 |
ap.add_argument("--policy", default="mixed_dit", choices=POLICIES, help="TTPolicy.dtype_policy")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 683 |
ap.add_argument(
|
| 684 |
"--trace-layout", default="per_stage", choices=TRACE_LAYOUTS, help="TTPolicy.trace_layout (plan §4)"
|
| 685 |
)
|
|
@@ -733,13 +819,24 @@ def main(argv: Optional[Sequence[str]] = None) -> int:
|
|
| 733 |
layout_name=args.layout,
|
| 734 |
host_reps=args.host_reps,
|
| 735 |
profile_encode=args.profile_encode,
|
|
|
|
|
|
|
| 736 |
)
|
| 737 |
SO.write_perf_json(payload, Path(args.out))
|
| 738 |
return 0
|
| 739 |
if import_gr00t_tt() is None:
|
| 740 |
raise ImportError("tt.model.Gr00tTT is not importable (WP-14 pending) -- run with --dry-run")
|
| 741 |
SO.ensure_tt_metal_cache()
|
| 742 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 743 |
trace_region_size=args.trace_region_mb << 20,
|
| 744 |
l1_small_size=D.DEFAULT_L1_SMALL_SIZE,
|
| 745 |
num_command_queues=args.num_cqs,
|
|
@@ -761,6 +858,9 @@ def main(argv: Optional[Sequence[str]] = None) -> int:
|
|
| 761 |
embed_gather=args.embed_gather,
|
| 762 |
cq1_uploads=resolve_cq1(args.cq1_uploads),
|
| 763 |
profile_encode=args.profile_encode,
|
|
|
|
|
|
|
|
|
|
| 764 |
)
|
| 765 |
finally:
|
| 766 |
D.close_gr00t_device(device)
|
|
@@ -772,6 +872,8 @@ def main(argv: Optional[Sequence[str]] = None) -> int:
|
|
| 772 |
__all__ = [
|
| 773 |
"BENCH_NAME",
|
| 774 |
"CQ1_MODES",
|
|
|
|
|
|
|
| 775 |
"EMBED_GATHERS",
|
| 776 |
"POLICIES",
|
| 777 |
"PRODUCTION_MEMBERS",
|
|
|
|
| 20 |
|
| 21 |
``--num-cqs`` (default 2) opens the device with that many command queues; ``--cq1-uploads auto|on|off`` and
|
| 22 |
``--embed-gather device|host`` are forwarded to ``Gr00tTT``; ``--profile-encode`` adds the per-step host-encode
|
| 23 |
+
profile (``meta.encode_profile``; also in ``--dry-run``). ``--dit-backend ttnn|megakernel`` (WP-K5) selects the
|
| 24 |
+
denoise backend (``TTPolicy.dit_backend``; default = the policy default, the **megakernel** since 2026-09-18; the
|
| 25 |
+
device is opened through ``tt.model.open_model_device`` so the megakernel gets its ``worker_l1_size``) and
|
| 26 |
+
``--mk-arena-dtype auto|bf16|bfp8_b`` the megakernel's weight arena; ``meta.dit_backend`` / ``meta.mk_arena_dtype`` /
|
| 27 |
+
``meta.megakernel`` / ``meta.weights`` (tensors and bytes ``TTWeights`` uploaded, what the megakernel skipped) record
|
| 28 |
+
what ran. ``--load-ttnn-dit-weights`` keeps the Stage-1 DiT upload under the megakernel (startup A/B).
|
| 29 |
|
| 30 |
The model comes from ``tt.model.Gr00tTT`` (WP-14). Until that module exists the import is guarded: ``--dry-run``
|
| 31 |
validates ``--version`` / ``--policy`` / ``--trace-layout`` / ``--layout`` / ``--embodiment`` against ``common.configs``
|
|
|
|
| 66 |
from models.experimental.gr00t.benchmarks import bench_stage_ops as SO
|
| 67 |
from models.experimental.gr00t.benchmarks import mb1_common as MB
|
| 68 |
from models.experimental.gr00t.tt import device as D
|
| 69 |
+
from models.experimental.gr00t.tt.policy import DEFAULT_DIT_BACKEND, DIT_BACKENDS, MK_ARENA_DTYPES # torch-free
|
| 70 |
|
| 71 |
BENCH_NAME = "bench_e2e"
|
| 72 |
VERSIONS = SO.VERSIONS
|
|
|
|
| 94 |
trace_layout: str = "per_stage",
|
| 95 |
layout_name: Optional[str] = None,
|
| 96 |
embodiment: Optional[str] = None,
|
| 97 |
+
dit_backend: str = DEFAULT_DIT_BACKEND,
|
| 98 |
+
mk_arena_dtype: str = "auto",
|
| 99 |
) -> Dict[str, Any]:
|
| 100 |
"""Validate the CLI choices against ``common.configs`` / ``tt.policy`` and return the resolved facts
|
| 101 |
(``cfg``, ``layout``, ``policy`` object, ``embodiment_tag`` / ``embodiment_id``, expected trace names)."""
|
|
|
|
| 110 |
layout = cfg.canonical_layout if layout_name is None else cfg.layout(layout_name)
|
| 111 |
tag = embodiment if embodiment is not None else layout.embodiment_tag
|
| 112 |
emb_id = cfg.embodiment_id(tag) # raises on an unknown tag
|
| 113 |
+
pol = TTPolicy(
|
| 114 |
+
dtype_policy=policy, trace_layout=trace_layout, dit_backend=dit_backend, mk_arena_dtype=mk_arena_dtype
|
| 115 |
+
) # raises on unknown values
|
| 116 |
return {
|
| 117 |
"cfg": cfg,
|
| 118 |
"layout": layout,
|
|
|
|
| 131 |
"dit_m_pad": layout.dit_m_pad,
|
| 132 |
"noise_shape": list(layout.noise_shape),
|
| 133 |
"policy": pol.to_dict(),
|
| 134 |
+
"dit_backend": pol.dit_backend,
|
| 135 |
+
"mk_arena_dtype": pol.mk_arena_dtype_for(version) if pol.dit_backend == "megakernel" else None,
|
| 136 |
"expected_stage_names": list(STAGE_NAMES[trace_layout]),
|
| 137 |
},
|
| 138 |
}
|
|
|
|
| 202 |
*,
|
| 203 |
embed_gather: str = "device",
|
| 204 |
cq1_uploads: Optional[bool] = None,
|
| 205 |
+
load_ttnn_dit_weights: Optional[bool] = None,
|
| 206 |
) -> Any:
|
| 207 |
Gr00tTT = import_gr00t_tt()
|
| 208 |
if Gr00tTT is None:
|
|
|
|
| 220 |
embed_gather=embed_gather,
|
| 221 |
cq1_uploads=cq1_uploads,
|
| 222 |
)
|
| 223 |
+
if load_ttnn_dit_weights is not None:
|
| 224 |
+
kw["load_ttnn_dit_weights"] = bool(load_ttnn_dit_weights)
|
| 225 |
if cache_root is not None:
|
| 226 |
kw["cache_root"] = cache_root
|
| 227 |
return Gr00tTT.from_pretrained(setup["cfg"].version, **kw)
|
|
|
|
| 250 |
return tracer
|
| 251 |
|
| 252 |
|
| 253 |
+
def _worker_l1_size(device: Any) -> Optional[int]:
|
| 254 |
+
"""``tt.model.device_worker_l1_size`` (the allocatable L1 the device was opened with; recorded in ``meta``)."""
|
| 255 |
+
try:
|
| 256 |
+
from models.experimental.gr00t.tt.model import device_worker_l1_size
|
| 257 |
+
except ImportError:
|
| 258 |
+
return None
|
| 259 |
+
return int(device_worker_l1_size(device))
|
| 260 |
+
|
| 261 |
+
|
| 262 |
def denoise_stage_name(stage_names: Sequence[str]) -> str:
|
| 263 |
"""The trace that holds the 4-step denoise: the last one of either layout (``denoise`` / ``B``)."""
|
| 264 |
if not stage_names:
|
|
|
|
| 430 |
embed_gather: str = "device",
|
| 431 |
cq1_uploads: Optional[bool] = None,
|
| 432 |
profile_encode: bool = False,
|
| 433 |
+
dit_backend: str = DEFAULT_DIT_BACKEND,
|
| 434 |
+
mk_arena_dtype: str = "auto",
|
| 435 |
+
load_ttnn_dit_weights: Optional[bool] = None,
|
| 436 |
) -> Dict[str, Any]:
|
| 437 |
"""Measure the traced end-to-end latency and its split on an open device; returns the perf payload (not written).
|
| 438 |
Main pass = production path (``tracer.call``), split pass = serialised writes + per-trace sync (module docstring).
|
|
|
|
| 445 |
if warmup < 0 or reps <= 0:
|
| 446 |
raise ValueError("warmup must be >= 0 and reps > 0")
|
| 447 |
setup = resolve_setup(
|
| 448 |
+
version,
|
| 449 |
+
policy=policy,
|
| 450 |
+
trace_layout=trace_layout,
|
| 451 |
+
layout_name=layout_name,
|
| 452 |
+
embodiment=embodiment,
|
| 453 |
+
dit_backend=dit_backend,
|
| 454 |
+
mk_arena_dtype=mk_arena_dtype,
|
| 455 |
)
|
| 456 |
ok, reason = golden_available(version, sample)
|
| 457 |
if not ok:
|
|
|
|
| 460 |
obs, mi = inputs["obs"], inputs["mi"]
|
| 461 |
noise = make_noise(setup["layout"], noise_seed)
|
| 462 |
t_load = time.perf_counter()
|
| 463 |
+
model = load_model(
|
| 464 |
+
setup,
|
| 465 |
+
device,
|
| 466 |
+
cache_root,
|
| 467 |
+
embed_gather=embed_gather,
|
| 468 |
+
cq1_uploads=cq1_uploads,
|
| 469 |
+
load_ttnn_dit_weights=load_ttnn_dit_weights,
|
| 470 |
+
)
|
| 471 |
load_s = time.perf_counter() - t_load
|
| 472 |
tracer = resolve_stage_replayer(model, setup["stage_names"])
|
| 473 |
missing = [m for m in PRODUCTION_MEMBERS if not hasattr(tracer, m)]
|
|
|
|
| 550 |
per_call_inputs=[s.describe() for s in executor.plan().inputs],
|
| 551 |
per_call_bytes=ex_desc.get("plan", {}).get("per_call_bytes"),
|
| 552 |
vit_fork=getattr(getattr(model.backbone, "tower", None), "fork", None),
|
| 553 |
+
dit_backend=model.policy.dit_backend,
|
| 554 |
+
mk_arena_dtype=model.recipe.mk_arena_dtype,
|
| 555 |
+
megakernel=model.head.describe().get("megakernel"),
|
| 556 |
+
worker_l1_size=_worker_l1_size(device),
|
| 557 |
+
head_timing={k: v for k, v in model.head.timing.items()},
|
| 558 |
+
model_timing={k: v for k, v in model.timing.items() if k.endswith("_s")},
|
| 559 |
+
weights={
|
| 560 |
+
"n_tensors": len(model.W),
|
| 561 |
+
"device_bytes": float(model.W.device_bytes),
|
| 562 |
+
"load_path": model.W.load_stats.get("path"),
|
| 563 |
+
"load_s": model.timing.get("weights_load_s"),
|
| 564 |
+
"skipped": dict(getattr(model, "weights_skipped", {})),
|
| 565 |
+
},
|
| 566 |
)
|
| 567 |
if profile_encode:
|
| 568 |
meta["encode_profile"] = encode_profile(version, inputs, setup["layout"])
|
|
|
|
| 583 |
"noise_seed": noise_seed,
|
| 584 |
"embed_gather": embed_gather,
|
| 585 |
"cq1_uploads": cq1_uploads,
|
| 586 |
+
"dit_backend": dit_backend,
|
| 587 |
+
"mk_arena_dtype": mk_arena_dtype,
|
| 588 |
+
"load_ttnn_dit_weights": load_ttnn_dit_weights,
|
| 589 |
},
|
| 590 |
method={
|
| 591 |
"main_pass": f"{warmup} warm-up calls, then {reps} calls: encode -> tracer.call (TraceExecutor.replay: "
|
|
|
|
| 642 |
host_reps: int = DEFAULT_HOST_REPS,
|
| 643 |
verbose: bool = True,
|
| 644 |
profile_encode: bool = False,
|
| 645 |
+
dit_backend: str = DEFAULT_DIT_BACKEND,
|
| 646 |
+
mk_arena_dtype: str = "auto",
|
| 647 |
) -> Dict[str, Any]:
|
| 648 |
"""Validate the setup and time the host side (encode / decode) when the golden data is present; no device."""
|
| 649 |
setup = resolve_setup(
|
| 650 |
+
version,
|
| 651 |
+
policy=policy,
|
| 652 |
+
trace_layout=trace_layout,
|
| 653 |
+
layout_name=layout_name,
|
| 654 |
+
embodiment=embodiment,
|
| 655 |
+
dit_backend=dit_backend,
|
| 656 |
+
mk_arena_dtype=mk_arena_dtype,
|
| 657 |
)
|
| 658 |
stages: Dict[str, Dict[str, Any]] = {}
|
| 659 |
meta = dict(setup["meta"])
|
|
|
|
| 685 |
"embodiment": setup["embodiment_tag"],
|
| 686 |
"layout": setup["layout"].name,
|
| 687 |
"host_reps": host_reps,
|
| 688 |
+
"dit_backend": dit_backend,
|
| 689 |
+
"mk_arena_dtype": mk_arena_dtype,
|
| 690 |
},
|
| 691 |
method={"host": f"{host_reps} host calls after 1 warm-up (encode / decode), median / p90 ms; no device"},
|
| 692 |
stages=stages,
|
|
|
|
| 721 |
)
|
| 722 |
m = payload["meta"]
|
| 723 |
print(
|
| 724 |
+
" layout {l} | embodiment {e} | S_pad {s} | shape_key {sk} | policy {p} | dit_backend {b} (arena {a}) | "
|
| 725 |
+
"stages {st} | embed_gather {g} | cq1_uploads {c} ({n} CQ){skip}".format(
|
| 726 |
l=m.get("layout_name"),
|
| 727 |
e=m.get("embodiment_tag"),
|
| 728 |
s=m.get("s_pad_llm"),
|
| 729 |
sk=m.get("shape_key", "-"),
|
| 730 |
p=(m.get("policy") or {}).get("dtype_policy"),
|
| 731 |
+
b=m.get("dit_backend", (m.get("policy") or {}).get("dit_backend")),
|
| 732 |
+
a=m.get("mk_arena_dtype", "-"),
|
| 733 |
st=m.get("stage_names", m.get("expected_stage_names")),
|
| 734 |
g=m.get("embed_gather", "-"),
|
| 735 |
c=m.get("cq1_uploads", "-"),
|
|
|
|
| 748 |
)
|
| 749 |
ap.add_argument("--version", required=True, choices=VERSIONS)
|
| 750 |
ap.add_argument("--policy", default="mixed_dit", choices=POLICIES, help="TTPolicy.dtype_policy")
|
| 751 |
+
ap.add_argument(
|
| 752 |
+
"--dit-backend",
|
| 753 |
+
default=DEFAULT_DIT_BACKEND,
|
| 754 |
+
choices=DIT_BACKENDS,
|
| 755 |
+
help=f"TTPolicy.dit_backend: Stage-1 ttnn ops | megakernel (default: the policy default, {DEFAULT_DIT_BACKEND})",
|
| 756 |
+
)
|
| 757 |
+
ap.add_argument(
|
| 758 |
+
"--load-ttnn-dit-weights",
|
| 759 |
+
action="store_true",
|
| 760 |
+
help="keep the Stage-1 DiT / precompute / encoder upload under the megakernel backend (startup A/B; "
|
| 761 |
+
"default: the megakernel skips them, tt.model.megakernel_weight_names)",
|
| 762 |
+
)
|
| 763 |
+
ap.add_argument(
|
| 764 |
+
"--mk-arena-dtype",
|
| 765 |
+
default="auto",
|
| 766 |
+
choices=MK_ARENA_DTYPES,
|
| 767 |
+
help="TTPolicy.mk_arena_dtype: the megakernel's weight arena (auto = the per-version default)",
|
| 768 |
+
)
|
| 769 |
ap.add_argument(
|
| 770 |
"--trace-layout", default="per_stage", choices=TRACE_LAYOUTS, help="TTPolicy.trace_layout (plan §4)"
|
| 771 |
)
|
|
|
|
| 819 |
layout_name=args.layout,
|
| 820 |
host_reps=args.host_reps,
|
| 821 |
profile_encode=args.profile_encode,
|
| 822 |
+
dit_backend=args.dit_backend,
|
| 823 |
+
mk_arena_dtype=args.mk_arena_dtype,
|
| 824 |
)
|
| 825 |
SO.write_perf_json(payload, Path(args.out))
|
| 826 |
return 0
|
| 827 |
if import_gr00t_tt() is None:
|
| 828 |
raise ImportError("tt.model.Gr00tTT is not importable (WP-14 pending) -- run with --dry-run")
|
| 829 |
SO.ensure_tt_metal_cache()
|
| 830 |
+
from models.experimental.gr00t.tt import model as M
|
| 831 |
+
from models.experimental.gr00t.tt.policy import TTPolicy
|
| 832 |
+
|
| 833 |
+
device = M.open_model_device(
|
| 834 |
+
TTPolicy(
|
| 835 |
+
dtype_policy=args.policy,
|
| 836 |
+
trace_layout=args.trace_layout,
|
| 837 |
+
dit_backend=args.dit_backend,
|
| 838 |
+
mk_arena_dtype=args.mk_arena_dtype,
|
| 839 |
+
),
|
| 840 |
trace_region_size=args.trace_region_mb << 20,
|
| 841 |
l1_small_size=D.DEFAULT_L1_SMALL_SIZE,
|
| 842 |
num_command_queues=args.num_cqs,
|
|
|
|
| 858 |
embed_gather=args.embed_gather,
|
| 859 |
cq1_uploads=resolve_cq1(args.cq1_uploads),
|
| 860 |
profile_encode=args.profile_encode,
|
| 861 |
+
dit_backend=args.dit_backend,
|
| 862 |
+
mk_arena_dtype=args.mk_arena_dtype,
|
| 863 |
+
load_ttnn_dit_weights=True if args.load_ttnn_dit_weights else None,
|
| 864 |
)
|
| 865 |
finally:
|
| 866 |
D.close_gr00t_device(device)
|
|
|
|
| 872 |
__all__ = [
|
| 873 |
"BENCH_NAME",
|
| 874 |
"CQ1_MODES",
|
| 875 |
+
"DIT_BACKENDS",
|
| 876 |
+
"MK_ARENA_DTYPES",
|
| 877 |
"EMBED_GATHERS",
|
| 878 |
"POLICIES",
|
| 879 |
"PRODUCTION_MEMBERS",
|
code/models/experimental/gr00t/benchmarks/bench_load.py
ADDED
|
@@ -0,0 +1,236 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
"""Model start-up cost of ``Gr00tTT.from_pretrained`` per DiT backend (megakernel-default pass; WP-K5 review NB 7 /
|
| 5 |
+
§8 pre-condition 5): wall time of the weight upload and the module builds, the tensors and bytes ``TTWeights`` puts
|
| 6 |
+
on the device, and what the megakernel backend skips (``tt.model.megakernel_weight_names``).
|
| 7 |
+
|
| 8 |
+
Per invocation the device is opened once through ``tt.model.open_model_device(policy)`` and the model is built
|
| 9 |
+
``--rounds`` times for each configuration -- ``skip`` (the default: the megakernel packs the DiT / precompute / encoder
|
| 10 |
+
tensors into its arena and ``TTWeights`` leaves them out) and ``keep`` (``load_ttnn_dit_weights=True``: the full
|
| 11 |
+
Stage-1 upload next to the arena, what WP-K5 measured) -- alternating ``keep, skip, keep, skip, ...`` so both see the
|
| 12 |
+
same OS page-cache state of the ``.tensorbin`` tier; every model is released before the next build. With
|
| 13 |
+
``--dit-backend ttnn`` only the ``keep`` configuration exists (the ttnn head reads those weights).
|
| 14 |
+
|
| 15 |
+
Output ``benchmarks/results/bench_load_<version>_<ts>.json`` (schema ``gr00t-tt-perf/1``): stages
|
| 16 |
+
``from_pretrained/<cfg>`` (median ms over the rounds), ``weights_load/<cfg>``, ``head_build/<cfg>`` (the megakernel
|
| 17 |
+
arena plan + pack + upload sit inside it); ``meta.configs`` holds the per-round numbers, the tensor counts, the bytes
|
| 18 |
+
on the device and the skipped set, ``meta.device_mb`` the MB of plan tensors per configuration and ``meta.delta`` the
|
| 19 |
+
keep - skip savings.
|
| 20 |
+
|
| 21 |
+
Run::
|
| 22 |
+
|
| 23 |
+
DEVICE_LOCK_TIMEOUT=14400 /home/deepgadget/experiments/gr00t/bin/with-device.sh \\
|
| 24 |
+
python -m models.experimental.gr00t.benchmarks.bench_load --version n16 [--rounds 2] [--dit-backend ttnn]
|
| 25 |
+
"""
|
| 26 |
+
|
| 27 |
+
from __future__ import annotations
|
| 28 |
+
|
| 29 |
+
import argparse
|
| 30 |
+
import statistics
|
| 31 |
+
import time
|
| 32 |
+
from pathlib import Path
|
| 33 |
+
from typing import Any, Dict, List, Optional, Sequence
|
| 34 |
+
|
| 35 |
+
from models.experimental.gr00t.benchmarks import bench_e2e as E2E
|
| 36 |
+
from models.experimental.gr00t.benchmarks import bench_stage_ops as SO
|
| 37 |
+
from models.experimental.gr00t.tt import device as D
|
| 38 |
+
from models.experimental.gr00t.tt.policy import DEFAULT_DIT_BACKEND, DIT_BACKENDS, MK_ARENA_DTYPES
|
| 39 |
+
|
| 40 |
+
BENCH_NAME = "bench_load"
|
| 41 |
+
CONFIGS: Dict[str, Optional[bool]] = {"keep": True, "skip": None} # -> from_pretrained(load_ttnn_dit_weights=)
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def configs_for(dit_backend: str) -> Dict[str, Optional[bool]]:
|
| 45 |
+
"""The configurations that exist for a backend: both for the megakernel, ``keep`` only for ttnn."""
|
| 46 |
+
if dit_backend == "megakernel":
|
| 47 |
+
return dict(CONFIGS)
|
| 48 |
+
return {"keep": None}
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def one_build(
|
| 52 |
+
setup: Dict[str, Any], device: Any, cache_root: Optional[Path], load_ttnn: Optional[bool]
|
| 53 |
+
) -> Dict[str, Any]:
|
| 54 |
+
"""Build the model once, record the timings / weight facts, release it."""
|
| 55 |
+
t0 = time.perf_counter()
|
| 56 |
+
model = E2E.load_model(setup, device, cache_root, load_ttnn_dit_weights=load_ttnn)
|
| 57 |
+
total = time.perf_counter() - t0
|
| 58 |
+
try:
|
| 59 |
+
timing = {k: float(v) for k, v in model.timing.items() if k.endswith("_s")}
|
| 60 |
+
head = {k: float(v) for k, v in model.head.timing.items()}
|
| 61 |
+
rec = {
|
| 62 |
+
"from_pretrained_s": total,
|
| 63 |
+
"weights_load_s": timing.get("weights_load_s"),
|
| 64 |
+
"backbone_build_s": timing.get("backbone_build_s"),
|
| 65 |
+
"adapter_build_s": timing.get("adapter_build_s"),
|
| 66 |
+
"embed_gather_build_s": timing.get("embed_gather_build_s"),
|
| 67 |
+
"executor_and_head_build_s": timing.get("executor_and_head_build_s"),
|
| 68 |
+
"head_timing": head,
|
| 69 |
+
"n_tensors": len(model.W),
|
| 70 |
+
"device_bytes": float(model.W.device_bytes),
|
| 71 |
+
"load_path": model.W.load_stats.get("path"),
|
| 72 |
+
"n_warm": model.W.load_stats.get("n_warm"),
|
| 73 |
+
"n_cold": model.W.load_stats.get("n_cold"),
|
| 74 |
+
"skipped": dict(model.weights_skipped),
|
| 75 |
+
"dit_backend": model.policy.dit_backend,
|
| 76 |
+
"mk_arena_dtype": model.recipe.mk_arena_dtype,
|
| 77 |
+
"adapter_mem": model.adapter_mem,
|
| 78 |
+
}
|
| 79 |
+
finally:
|
| 80 |
+
model.release()
|
| 81 |
+
import ttnn
|
| 82 |
+
|
| 83 |
+
ttnn.synchronize_device(device)
|
| 84 |
+
return rec
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
def run(
|
| 88 |
+
device: Any,
|
| 89 |
+
*,
|
| 90 |
+
version: str,
|
| 91 |
+
policy: str = "mixed_dit",
|
| 92 |
+
trace_layout: str = "per_stage",
|
| 93 |
+
dit_backend: str = DEFAULT_DIT_BACKEND,
|
| 94 |
+
mk_arena_dtype: str = "auto",
|
| 95 |
+
rounds: int = 2,
|
| 96 |
+
cache_root: Optional[Path] = None,
|
| 97 |
+
verbose: bool = True,
|
| 98 |
+
) -> Dict[str, Any]:
|
| 99 |
+
if rounds <= 0:
|
| 100 |
+
raise ValueError("rounds must be positive")
|
| 101 |
+
setup = E2E.resolve_setup(
|
| 102 |
+
version, policy=policy, trace_layout=trace_layout, dit_backend=dit_backend, mk_arena_dtype=mk_arena_dtype
|
| 103 |
+
)
|
| 104 |
+
cfgs = configs_for(dit_backend)
|
| 105 |
+
per_cfg: Dict[str, List[Dict[str, Any]]] = {name: [] for name in cfgs}
|
| 106 |
+
for r in range(rounds):
|
| 107 |
+
for name, load_ttnn in cfgs.items():
|
| 108 |
+
rec = one_build(setup, device, cache_root, load_ttnn)
|
| 109 |
+
rec["round"] = r
|
| 110 |
+
per_cfg[name].append(rec)
|
| 111 |
+
if verbose:
|
| 112 |
+
print(
|
| 113 |
+
f"[bench_load] {version}/{dit_backend}/{name} round {r}: from_pretrained {rec['from_pretrained_s']:.2f} s "
|
| 114 |
+
f"(weights {rec['weights_load_s']:.2f} s, {rec['n_tensors']} tensors, {rec['device_bytes'] / 1e6:.0f} MB, "
|
| 115 |
+
f"{rec['load_path']}; head {rec['executor_and_head_build_s']:.2f} s"
|
| 116 |
+
f"{'; skipped ' + str(rec['skipped'].get('n_tensors')) + ' tensors / ' + format(rec['skipped'].get('device_bytes', 0) / 1e6, '.0f') + ' MB' if rec['skipped'] else ''})"
|
| 117 |
+
)
|
| 118 |
+
stages: Dict[str, Dict[str, Any]] = {}
|
| 119 |
+
for name, recs in per_cfg.items():
|
| 120 |
+
|
| 121 |
+
def med(key: str) -> float:
|
| 122 |
+
vals = [float(x[key]) for x in recs if x.get(key) is not None]
|
| 123 |
+
return statistics.median(vals) if vals else float("nan")
|
| 124 |
+
|
| 125 |
+
stages[f"from_pretrained/{name}"] = {
|
| 126 |
+
"metric": "ms",
|
| 127 |
+
"value": med("from_pretrained_s") * 1e3,
|
| 128 |
+
"min": min(float(x["from_pretrained_s"]) for x in recs) * 1e3,
|
| 129 |
+
"per_round": [round(float(x["from_pretrained_s"]) * 1e3, 1) for x in recs],
|
| 130 |
+
}
|
| 131 |
+
stages[f"weights_load/{name}"] = {"metric": "ms", "value": med("weights_load_s") * 1e3}
|
| 132 |
+
stages[f"head_build/{name}"] = {
|
| 133 |
+
"metric": "ms",
|
| 134 |
+
"value": med("executor_and_head_build_s") * 1e3,
|
| 135 |
+
"note": "executor + ActionHeadTT build (megakernel: arena plan + pack + upload + io inside)",
|
| 136 |
+
}
|
| 137 |
+
facts = D.device_facts(device)
|
| 138 |
+
facts["worker_l1_size"] = E2E._worker_l1_size(device)
|
| 139 |
+
meta = dict(
|
| 140 |
+
setup["meta"],
|
| 141 |
+
configs={name: recs for name, recs in per_cfg.items()},
|
| 142 |
+
device_mb={
|
| 143 |
+
name: {"mb": recs[0]["device_bytes"] / 1e6, "n_tensors": recs[0]["n_tensors"]}
|
| 144 |
+
for name, recs in per_cfg.items()
|
| 145 |
+
},
|
| 146 |
+
rounds=rounds,
|
| 147 |
+
worker_l1_size=facts["worker_l1_size"],
|
| 148 |
+
)
|
| 149 |
+
if "skip" in per_cfg and "keep" in per_cfg:
|
| 150 |
+
k, s = per_cfg["keep"][0], per_cfg["skip"][0]
|
| 151 |
+
meta["delta"] = {
|
| 152 |
+
"device_mb_saved": (k["device_bytes"] - s["device_bytes"]) / 1e6,
|
| 153 |
+
"n_tensors_saved": k["n_tensors"] - s["n_tensors"],
|
| 154 |
+
"from_pretrained_ms_saved_median": stages["from_pretrained/keep"]["value"]
|
| 155 |
+
- stages["from_pretrained/skip"]["value"],
|
| 156 |
+
"weights_load_ms_saved_median": stages["weights_load/keep"]["value"] - stages["weights_load/skip"]["value"],
|
| 157 |
+
}
|
| 158 |
+
payload = SO.perf_payload(
|
| 159 |
+
BENCH_NAME,
|
| 160 |
+
version,
|
| 161 |
+
dry_run=False,
|
| 162 |
+
args={
|
| 163 |
+
"version": version,
|
| 164 |
+
"policy": policy,
|
| 165 |
+
"trace_layout": trace_layout,
|
| 166 |
+
"dit_backend": dit_backend,
|
| 167 |
+
"mk_arena_dtype": mk_arena_dtype,
|
| 168 |
+
"rounds": rounds,
|
| 169 |
+
},
|
| 170 |
+
method={
|
| 171 |
+
"builds": f"{rounds} round(s) of alternating keep / skip Gr00tTT.from_pretrained on one open device, "
|
| 172 |
+
"each model released before the next build; medians over the rounds",
|
| 173 |
+
"clock": "host wall clock (time.perf_counter)",
|
| 174 |
+
},
|
| 175 |
+
stages=stages,
|
| 176 |
+
rows=[],
|
| 177 |
+
device=facts,
|
| 178 |
+
meta=meta,
|
| 179 |
+
)
|
| 180 |
+
return payload
|
| 181 |
+
|
| 182 |
+
|
| 183 |
+
def build_parser() -> argparse.ArgumentParser:
|
| 184 |
+
ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
|
| 185 |
+
ap.add_argument("--version", required=True, choices=E2E.VERSIONS)
|
| 186 |
+
ap.add_argument("--policy", default="mixed_dit", choices=E2E.POLICIES)
|
| 187 |
+
ap.add_argument("--trace-layout", default="per_stage", choices=E2E.TRACE_LAYOUTS)
|
| 188 |
+
ap.add_argument("--dit-backend", default=DEFAULT_DIT_BACKEND, choices=DIT_BACKENDS)
|
| 189 |
+
ap.add_argument("--mk-arena-dtype", default="auto", choices=MK_ARENA_DTYPES)
|
| 190 |
+
ap.add_argument("--rounds", type=int, default=2)
|
| 191 |
+
ap.add_argument("--cache-root", default=None)
|
| 192 |
+
ap.add_argument("--out", default=str(SO.RESULTS_DIR))
|
| 193 |
+
ap.add_argument("--trace-region-mb", type=int, default=64)
|
| 194 |
+
ap.add_argument("--num-cqs", type=int, default=2, choices=(1, 2))
|
| 195 |
+
return ap
|
| 196 |
+
|
| 197 |
+
|
| 198 |
+
def main(argv: Optional[Sequence[str]] = None) -> int:
|
| 199 |
+
args = build_parser().parse_args(argv)
|
| 200 |
+
if E2E.import_gr00t_tt() is None:
|
| 201 |
+
raise ImportError("tt.model.Gr00tTT is not importable")
|
| 202 |
+
SO.ensure_tt_metal_cache()
|
| 203 |
+
from models.experimental.gr00t.tt import model as M
|
| 204 |
+
from models.experimental.gr00t.tt.policy import TTPolicy
|
| 205 |
+
|
| 206 |
+
device = M.open_model_device(
|
| 207 |
+
TTPolicy(
|
| 208 |
+
dtype_policy=args.policy,
|
| 209 |
+
trace_layout=args.trace_layout,
|
| 210 |
+
dit_backend=args.dit_backend,
|
| 211 |
+
mk_arena_dtype=args.mk_arena_dtype,
|
| 212 |
+
),
|
| 213 |
+
trace_region_size=args.trace_region_mb << 20,
|
| 214 |
+
l1_small_size=D.DEFAULT_L1_SMALL_SIZE,
|
| 215 |
+
num_command_queues=args.num_cqs,
|
| 216 |
+
)
|
| 217 |
+
try:
|
| 218 |
+
payload = run(
|
| 219 |
+
device,
|
| 220 |
+
version=args.version,
|
| 221 |
+
policy=args.policy,
|
| 222 |
+
trace_layout=args.trace_layout,
|
| 223 |
+
dit_backend=args.dit_backend,
|
| 224 |
+
mk_arena_dtype=args.mk_arena_dtype,
|
| 225 |
+
rounds=args.rounds,
|
| 226 |
+
cache_root=Path(args.cache_root) if args.cache_root else None,
|
| 227 |
+
)
|
| 228 |
+
finally:
|
| 229 |
+
D.close_gr00t_device(device)
|
| 230 |
+
path = SO.write_perf_json(payload, Path(args.out))
|
| 231 |
+
print(f"[bench_load] wrote {path}")
|
| 232 |
+
return 0
|
| 233 |
+
|
| 234 |
+
|
| 235 |
+
if __name__ == "__main__":
|
| 236 |
+
raise SystemExit(main())
|
code/models/experimental/gr00t/benchmarks/bench_mk_step.py
ADDED
|
@@ -0,0 +1,291 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
"""WP-K4 timing / zone bench of the DiT megakernel (IMPLEMENTATION_PLAN.md §5.7 gate G2, §5.4 timing chain).
|
| 5 |
+
|
| 6 |
+
Modes::
|
| 7 |
+
|
| 8 |
+
# G2: the steps4 program captured in a trace, N replays -> ms per step vs 1.3 x arena bytes / measured GB/s
|
| 9 |
+
bench_mk_step.py --version n16 --dtype bf16 --replays 20
|
| 10 |
+
# per-phase DeviceZoneScopedN zones of the block variant over a few blocks (TT_METAL_DEVICE_PROFILER=1 is set
|
| 11 |
+
# before ttnn is imported; the profiler cannot be combined with the watcher on Blackhole)
|
| 12 |
+
bench_mk_step.py --version n16 --dtype bf16 --profile --blocks 4
|
| 13 |
+
|
| 14 |
+
Every run writes ``benchmarks/results/bench_mk_step_<version>_<stamp>.json`` (schema ``gr00t-mk-k4/1``). The
|
| 15 |
+
arena is packed from the real plan tensors (``dit_program.real_arena_plan``), the K/V are hoisted on device from the
|
| 16 |
+
canonical golden sample and the masks are the call's (``tests/tt/test_mk_dit_block.K4Session``). Run through
|
| 17 |
+
``bin/with-device.sh`` from the tt-metal root; ``ttnn`` is imported lazily in :func:`main`.
|
| 18 |
+
"""
|
| 19 |
+
|
| 20 |
+
from __future__ import annotations
|
| 21 |
+
|
| 22 |
+
import argparse
|
| 23 |
+
import csv
|
| 24 |
+
import json
|
| 25 |
+
import os
|
| 26 |
+
import statistics
|
| 27 |
+
import sys
|
| 28 |
+
import time
|
| 29 |
+
from collections import defaultdict
|
| 30 |
+
from datetime import datetime
|
| 31 |
+
from pathlib import Path
|
| 32 |
+
from typing import Any, Dict, List, Optional, Sequence, Tuple
|
| 33 |
+
|
| 34 |
+
import torch
|
| 35 |
+
|
| 36 |
+
TT_METAL_HOME = Path(os.environ.get("TT_METAL_HOME", str(Path(__file__).resolve().parents[4])))
|
| 37 |
+
PROFILE_CSV = TT_METAL_HOME / "generated" / "profiler" / ".logs" / "profile_log_device.csv"
|
| 38 |
+
RESULTS_DIR = Path(__file__).resolve().parent / "results"
|
| 39 |
+
SCHEMA = "gr00t-mk-k4/1"
|
| 40 |
+
MEASURED_GBPS: Dict[str, float] = {"bf16": 464.0, "bfp8_b": 414.0} # mk_k1_summary.md §3 (direct mode)
|
| 41 |
+
G2_FACTOR = 1.3
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def parse_args(argv: Optional[Sequence[str]] = None) -> argparse.Namespace:
|
| 45 |
+
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
| 46 |
+
ap.add_argument("--version", default="n16", choices=("n15", "n16", "n17"))
|
| 47 |
+
ap.add_argument("--dtype", default="bf16", choices=("bf16", "bfp8_b"), help="arena stream dtype")
|
| 48 |
+
ap.add_argument("--replays", type=int, default=20, help="traced steps4 replays for the G2 timing")
|
| 49 |
+
ap.add_argument("--profile", action="store_true", help="collect the per-phase zones of a short block launch")
|
| 50 |
+
ap.add_argument(
|
| 51 |
+
"--profile-variant",
|
| 52 |
+
default="block",
|
| 53 |
+
choices=("block", "step"),
|
| 54 |
+
help="profiled launch: `block` = blocks [0, --blocks) of the block variant; `step` = one whole denoising step "
|
| 55 |
+
"(all blocks + output head; set GR00T_MK_ZONE_ROUNDS=0 so the hub BRISC's 250 profiler markers hold every block)",
|
| 56 |
+
)
|
| 57 |
+
ap.add_argument("--blocks", type=int, default=4, help="blocks of the profiled block launch")
|
| 58 |
+
ap.add_argument(
|
| 59 |
+
"--first-block",
|
| 60 |
+
type=int,
|
| 61 |
+
default=0,
|
| 62 |
+
help="first block of the profiled block launch (the launch's first block streams its weights cold: profile "
|
| 63 |
+
"from block 1 or 2 to see the steady-state prefetch)",
|
| 64 |
+
)
|
| 65 |
+
ap.add_argument("--step", type=int, default=0, help="denoising step of the profiled block launch")
|
| 66 |
+
ap.add_argument(
|
| 67 |
+
"--cb-w-slots", type=int, default=None, help="weight-ring depth in K-block slots (K4Budget default)"
|
| 68 |
+
)
|
| 69 |
+
ap.add_argument(
|
| 70 |
+
"--k5-device",
|
| 71 |
+
action="store_true",
|
| 72 |
+
help="open the device as the K5 model does (tt.model.worker_l1_size_for(TTPolicy(dit_backend='megakernel')), the "
|
| 73 |
+
"64 KiB worker-L1 cut) instead of the K4 test device (96 KiB): the compute-core binaries must fit K5's "
|
| 74 |
+
"kernel-config ring buffer (mk_k4b_summary.md section 5) -- a program too large fails at the first launch",
|
| 75 |
+
)
|
| 76 |
+
ap.add_argument("--results-dir", default=str(RESULTS_DIR))
|
| 77 |
+
return ap.parse_args(argv)
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def read_zones(path: Path = PROFILE_CSV) -> Dict[str, Any]:
|
| 81 |
+
"""``{"freq_mhz", "zones": {"<RISC>:<zone>": {count, median_us, mean_us, min_us, max_us, cores}}}`` over the
|
| 82 |
+
``mk_*`` zones of every launch in the profiler CSV (nested zones handled with a per-(run, core, RISC) stack)."""
|
| 83 |
+
if not path.is_file():
|
| 84 |
+
raise FileNotFoundError(f"{path} not found (set TT_METAL_DEVICE_PROFILER=1 before importing ttnn)")
|
| 85 |
+
with open(path) as fh:
|
| 86 |
+
header = fh.readline()
|
| 87 |
+
freq_mhz = float(header.split("CHIP_FREQ[MHz]:")[1].split(",")[0])
|
| 88 |
+
reader = csv.reader(fh)
|
| 89 |
+
cols = [c.strip() for c in next(reader)]
|
| 90 |
+
idx = {c: i for i, c in enumerate(cols)}
|
| 91 |
+
stacks: Dict[tuple, List[Tuple[str, int]]] = defaultdict(list)
|
| 92 |
+
durations: Dict[str, List[int]] = defaultdict(list)
|
| 93 |
+
cores: Dict[str, set] = defaultdict(set)
|
| 94 |
+
for row in reader:
|
| 95 |
+
if len(row) < len(cols) - 1:
|
| 96 |
+
continue
|
| 97 |
+
zone = row[idx["zone name"]].strip()
|
| 98 |
+
if not zone.startswith("mk_"):
|
| 99 |
+
continue
|
| 100 |
+
core = (
|
| 101 |
+
row[idx["run host ID"]].strip(),
|
| 102 |
+
row[idx["core_x"]],
|
| 103 |
+
row[idx["core_y"]],
|
| 104 |
+
row[idx["RISC processor type"]].strip(),
|
| 105 |
+
)
|
| 106 |
+
t = int(row[idx["time[cycles since reset]"]])
|
| 107 |
+
typ = row[idx["type"]].strip()
|
| 108 |
+
if typ == "ZONE_START":
|
| 109 |
+
stacks[core].append((zone, t))
|
| 110 |
+
elif typ == "ZONE_END":
|
| 111 |
+
st = stacks[core]
|
| 112 |
+
for j in range(len(st) - 1, -1, -1):
|
| 113 |
+
if st[j][0] == zone:
|
| 114 |
+
name = f"{core[3]}:{zone}"
|
| 115 |
+
durations[name].append(t - st[j][1])
|
| 116 |
+
cores[name].add(core[1:3])
|
| 117 |
+
del st[j]
|
| 118 |
+
break
|
| 119 |
+
zones = {}
|
| 120 |
+
for name, ds in durations.items():
|
| 121 |
+
us = [d / freq_mhz for d in ds]
|
| 122 |
+
zones[name] = {
|
| 123 |
+
"count": len(us),
|
| 124 |
+
"cores": len(cores[name]),
|
| 125 |
+
"median_us": statistics.median(us),
|
| 126 |
+
"mean_us": sum(us) / len(us),
|
| 127 |
+
"min_us": min(us),
|
| 128 |
+
"max_us": max(us),
|
| 129 |
+
}
|
| 130 |
+
return {"freq_mhz": freq_mhz, "zones": zones}
|
| 131 |
+
|
| 132 |
+
|
| 133 |
+
def g2_record(per_step_ms: float, bytes_per_step: int, dtype: str) -> Dict[str, Any]:
|
| 134 |
+
floor_ms = bytes_per_step / MEASURED_GBPS[dtype] / 1e6
|
| 135 |
+
return {
|
| 136 |
+
"per_step_ms": per_step_ms,
|
| 137 |
+
"bytes_per_step": bytes_per_step,
|
| 138 |
+
"measured_gbps": MEASURED_GBPS[dtype],
|
| 139 |
+
"floor_ms_per_step": floor_ms,
|
| 140 |
+
"gate_ms_per_step": G2_FACTOR * floor_ms,
|
| 141 |
+
"ratio_to_floor": per_step_ms / floor_ms,
|
| 142 |
+
"effective_gbps": bytes_per_step / per_step_ms / 1e6,
|
| 143 |
+
"pass": per_step_ms <= G2_FACTOR * floor_ms,
|
| 144 |
+
}
|
| 145 |
+
|
| 146 |
+
|
| 147 |
+
def main(argv: Optional[Sequence[str]] = None) -> int:
|
| 148 |
+
args = parse_args(argv)
|
| 149 |
+
if args.profile:
|
| 150 |
+
os.environ["TT_METAL_DEVICE_PROFILER"] = "1"
|
| 151 |
+
# the device CSV of the last run and the profiler's accumulated zone-source-location logs: stale entries of
|
| 152 |
+
# earlier builds of the same zone names at other line numbers can collide in its 16-bit hash (profiler.cpp
|
| 153 |
+
# `populateZoneSrcLocations` throws at device close) -- they are rebuilt from this run's compile logs
|
| 154 |
+
for stale in (PROFILE_CSV, *PROFILE_CSV.parent.glob("*zone_src_locations.log")):
|
| 155 |
+
if stale.exists():
|
| 156 |
+
stale.unlink()
|
| 157 |
+
os.environ.setdefault("TT_METAL_CACHE", str(Path.home() / ".cache" / "tt-metal-cache-gr00t"))
|
| 158 |
+
torch.set_grad_enabled(False)
|
| 159 |
+
import ttnn # noqa: F401 (after the profiler env var)
|
| 160 |
+
|
| 161 |
+
from models.experimental.gr00t.common.preprocessing import encode
|
| 162 |
+
from models.experimental.gr00t.reference.model import load_golden_sample
|
| 163 |
+
from models.experimental.gr00t.tests.tt import test_mk_dit_block as TB
|
| 164 |
+
from models.experimental.gr00t.tests.tt.harness import device_facts
|
| 165 |
+
from models.experimental.gr00t.tt.device import close_gr00t_device
|
| 166 |
+
from models.experimental.gr00t.tt.megakernel import descriptors as D
|
| 167 |
+
from models.experimental.gr00t.tt.megakernel import dit_program as P
|
| 168 |
+
|
| 169 |
+
TB.STREAM_DTYPE = args.dtype # the session packs the arena in this dtype and picks the matching oracle policy
|
| 170 |
+
if args.cb_w_slots is not None:
|
| 171 |
+
os.environ["GR00T_MK_CB_W_SLOTS"] = str(args.cb_w_slots)
|
| 172 |
+
gs, obs = load_golden_sample(args.version, "canonical")
|
| 173 |
+
mi = encode(args.version, obs)
|
| 174 |
+
out: Dict[str, Any] = {
|
| 175 |
+
"schema": SCHEMA,
|
| 176 |
+
"created": datetime.now().astimezone().isoformat(timespec="seconds"),
|
| 177 |
+
"version": args.version,
|
| 178 |
+
"dtype": args.dtype,
|
| 179 |
+
"cb_w_slots": args.cb_w_slots,
|
| 180 |
+
"profile": bool(args.profile),
|
| 181 |
+
"profile_variant": args.profile_variant if args.profile else None,
|
| 182 |
+
"env": {k: v for k, v in os.environ.items() if k.startswith("GR00T_MK_")},
|
| 183 |
+
"argv": list(argv) if argv is not None else sys.argv[1:],
|
| 184 |
+
}
|
| 185 |
+
worker_l1 = None
|
| 186 |
+
if args.k5_device:
|
| 187 |
+
from models.experimental.gr00t.tt.model import worker_l1_size_for
|
| 188 |
+
from models.experimental.gr00t.tt.policy import TTPolicy
|
| 189 |
+
|
| 190 |
+
worker_l1 = int(worker_l1_size_for(TTPolicy(dit_backend="megakernel")))
|
| 191 |
+
out["k5_device"] = bool(args.k5_device)
|
| 192 |
+
out["worker_l1_size"] = worker_l1 if worker_l1 is not None else P.K4_WORKER_L1_SIZE
|
| 193 |
+
device = P.open_k4_device(worker_l1_size=out["worker_l1_size"])
|
| 194 |
+
try:
|
| 195 |
+
out["device"] = device_facts(device)
|
| 196 |
+
s = TB.K4Session(args.version, (gs, obs, mi), device, variant="block")
|
| 197 |
+
out["arena"] = {
|
| 198 |
+
"n_blocks": s.arena.n_blocks,
|
| 199 |
+
"bytes_per_step": s.mk.arena_bytes_per_step(),
|
| 200 |
+
"stream_bytes_per_step": s.arena.bytes_per_bank("stream") * 8,
|
| 201 |
+
"small_bytes_per_step": s.arena.bytes_per_bank("small") * 8,
|
| 202 |
+
"weights_load_s": s.weights_s,
|
| 203 |
+
"arena_pack_upload_s": s.arena_s,
|
| 204 |
+
}
|
| 205 |
+
if args.profile:
|
| 206 |
+
# one short launch: the zones of every RISC of every core land in the profiler CSV
|
| 207 |
+
tap = "sa_embs" if args.first_block == 0 else f"dit_block{args.first_block - 1}_out"
|
| 208 |
+
x = s.golden_rows(tap, args.step)
|
| 209 |
+
xr = s.mk.make_x_rows(x)
|
| 210 |
+
if args.profile_variant == "step":
|
| 211 |
+
s.mk.set_variant("step")
|
| 212 |
+
t0 = time.perf_counter()
|
| 213 |
+
s.mk.run_step(xr, args.step)
|
| 214 |
+
wall = time.perf_counter() - t0
|
| 215 |
+
n_blk = s.arena.n_blocks
|
| 216 |
+
else:
|
| 217 |
+
t0 = time.perf_counter()
|
| 218 |
+
s.mk.run_block(xr, args.first_block, matmul_only=False, n_blk=args.blocks, step=args.step)
|
| 219 |
+
wall = time.perf_counter() - t0
|
| 220 |
+
n_blk = args.blocks
|
| 221 |
+
s.mk._ttnn.deallocate(xr)
|
| 222 |
+
out["profile_launch"] = {
|
| 223 |
+
"variant": args.profile_variant,
|
| 224 |
+
"blocks": n_blk,
|
| 225 |
+
"first_block": args.first_block if args.profile_variant == "block" else 0,
|
| 226 |
+
"step": args.step,
|
| 227 |
+
"wall_s": wall,
|
| 228 |
+
"zone_rounds": P.K4_ZONE_ROUNDS,
|
| 229 |
+
}
|
| 230 |
+
else:
|
| 231 |
+
s.mk.set_variant("steps4")
|
| 232 |
+
noise = gs.load("initial_noise")
|
| 233 |
+
state = gs.load("state_features").reshape(-1)
|
| 234 |
+
t0 = time.perf_counter()
|
| 235 |
+
ref = s.mk.run_steps4(noise, state)
|
| 236 |
+
out["steps4_untraced_wall_s"] = time.perf_counter() - t0
|
| 237 |
+
s.mk.set_steps4_inputs(noise, state)
|
| 238 |
+
tid = s.mk.capture_steps4()
|
| 239 |
+
try:
|
| 240 |
+
times = [s.mk.replay(tid) for _ in range(args.replays)]
|
| 241 |
+
traced = s.mk.read_steps4_outputs()
|
| 242 |
+
finally:
|
| 243 |
+
s.mk.release_trace(tid)
|
| 244 |
+
ms = sorted(t * 1e3 for t in times)
|
| 245 |
+
out["steps4_trace"] = {
|
| 246 |
+
"replays": args.replays,
|
| 247 |
+
"ms_min": ms[0],
|
| 248 |
+
"ms_median": statistics.median(ms),
|
| 249 |
+
"ms_max": ms[-1],
|
| 250 |
+
"traced_eq_untraced": bool(torch.equal(traced["action_pred"], ref["action_pred"])),
|
| 251 |
+
}
|
| 252 |
+
out["g2"] = g2_record(statistics.median(ms) / P.STEPS4_N_STEPS, s.mk.arena_bytes_per_step(), args.dtype)
|
| 253 |
+
out["g2_min"] = g2_record(ms[0] / P.STEPS4_N_STEPS, s.mk.arena_bytes_per_step(), args.dtype)
|
| 254 |
+
out["stats"] = s.mk.stats().__dict__
|
| 255 |
+
out["k4b"] = { # the K4b schedule knobs of this run (dit_program / descriptors)
|
| 256 |
+
"mcast_split": P.K4_MCAST_SPLIT,
|
| 257 |
+
"zone_rounds": P.K4_ZONE_ROUNDS,
|
| 258 |
+
"land_pages": D.K4_LAND_PAGES,
|
| 259 |
+
"hub": list(s.cm.hub),
|
| 260 |
+
"hub2": list(s.cm.hub2),
|
| 261 |
+
"split_blocks": [b for b in range(s.arena.n_blocks) if P.is_split_block(s.arena, b)],
|
| 262 |
+
"cb_w_slots": s.mk.k4.cb_w_slots[args.dtype],
|
| 263 |
+
}
|
| 264 |
+
s.release()
|
| 265 |
+
finally:
|
| 266 |
+
close_gr00t_device(device)
|
| 267 |
+
if args.profile:
|
| 268 |
+
zones = read_zones()
|
| 269 |
+
out["zones"] = zones
|
| 270 |
+
print(f"zones ({zones['freq_mhz']} MHz), {args.blocks} blocks:")
|
| 271 |
+
for name, z in sorted(zones["zones"].items()):
|
| 272 |
+
print(
|
| 273 |
+
f" {name:28s} n={z['count']:5d} cores={z['cores']:3d} median {z['median_us']:8.2f} us min {z['min_us']:8.2f} max {z['max_us']:8.2f}"
|
| 274 |
+
)
|
| 275 |
+
else:
|
| 276 |
+
g = out["g2"]
|
| 277 |
+
print(
|
| 278 |
+
f"{args.version} {args.dtype}: steps4 trace median {out['steps4_trace']['ms_median']:.3f} ms "
|
| 279 |
+
f"(min {out['steps4_trace']['ms_min']:.3f}) -> {g['per_step_ms']:.3f} ms/step; floor {g['floor_ms_per_step']:.3f}, "
|
| 280 |
+
f"gate {g['gate_ms_per_step']:.3f}, ratio {g['ratio_to_floor']:.2f}, {g['effective_gbps']:.0f} GB/s -> G2 {'PASS' if g['pass'] else 'FAIL'}"
|
| 281 |
+
)
|
| 282 |
+
Path(args.results_dir).mkdir(parents=True, exist_ok=True)
|
| 283 |
+
path = Path(args.results_dir) / f"bench_mk_step_{args.version}_{datetime.now().strftime('%Y%m%d-%H%M%S')}.json"
|
| 284 |
+
with open(path, "w") as fh:
|
| 285 |
+
json.dump(out, fh, indent=1, default=str)
|
| 286 |
+
print("wrote", path)
|
| 287 |
+
return 0
|
| 288 |
+
|
| 289 |
+
|
| 290 |
+
if __name__ == "__main__":
|
| 291 |
+
raise SystemExit(main())
|
code/models/experimental/gr00t/benchmarks/sweep_backbone.py
CHANGED
|
@@ -1394,7 +1394,8 @@ def sweep_cq(
|
|
| 1394 |
row: Dict[str, Any] = {"cq1_uploads": flag}
|
| 1395 |
model = None
|
| 1396 |
try:
|
| 1397 |
-
model
|
|
|
|
| 1398 |
model.executor.set_cq1_uploads(flag)
|
| 1399 |
model.warm_and_capture(mi, noise)
|
| 1400 |
times = []
|
|
|
|
| 1394 |
row: Dict[str, Any] = {"cq1_uploads": flag}
|
| 1395 |
model = None
|
| 1396 |
try:
|
| 1397 |
+
# the Stage-1 model on the default-L1 device (the megakernel default needs open_model_device)
|
| 1398 |
+
model = Gr00tTT.from_pretrained(version, policy=TTPolicy(dit_backend="ttnn"), device=device)
|
| 1399 |
model.executor.set_cq1_uploads(flag)
|
| 1400 |
model.warm_and_capture(mi, noise)
|
| 1401 |
times = []
|
code/models/experimental/gr00t/tests/tt/test_mk_2cq.py
ADDED
|
@@ -0,0 +1,384 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
"""Command-queue bit-equality of the traced ``Gr00tTT`` under the megakernel DiT backend (megakernel-default pass;
|
| 5 |
+
``docs/plan/reviews/WP-K5-review-round1.md`` §2 "Command-queue paths", NB 3 and §8 pre-condition 3).
|
| 6 |
+
|
| 7 |
+
The production path opens the device with **two** command queues (``Gr00tTT.from_pretrained`` default
|
| 8 |
+
``num_command_queues=2``, the policy server's ``GR00T_NUM_CQS`` default, ``bench_e2e --num-cqs 2``): the per-call input
|
| 9 |
+
writes of call ``i + 1`` go to CQ 1 while the traces of call ``i`` run on CQ 0, fenced by ``record_event`` /
|
| 10 |
+
``wait_for_event`` (``tt.tracer`` ``cq1_uploads``). Every WP-K5 device test ran on one CQ and only the bench exercised
|
| 11 |
+
two, checking finiteness and shape. This module runs the same version / sample on a device opened **twice in one
|
| 12 |
+
process** through ``tt.model.open_model_device(policy, num_command_queues=...)`` -- first with 2 CQs (CQ-1 writes on),
|
| 13 |
+
then with 1 CQ -- and asserts, per selected sample:
|
| 14 |
+
|
| 15 |
+
* on each device: traced == untraced bit for bit on the raw ``action_pred`` buffer, the kernel's per-step
|
| 16 |
+
``velocity_out`` taps, ``backbone_features`` and ``K_0`` / ``V_0``; ``N_REPLAYS`` traced replays bit-identical;
|
| 17 |
+
* across the devices: the untraced raw tensors, the traced raw tensors, the host ``action_pred_normalized`` and the
|
| 18 |
+
per-step velocities of the 2-CQ run equal the 1-CQ run's bit for bit (``torch.equal``; PCC 1.0, max |d| 0 recorded);
|
| 19 |
+
* the facts: ``cq1_uploads`` on and ``num_command_queues == 2`` on the first device, off / 1 on the second, the
|
| 20 |
+
policy (backend, arena) and the ``worker_l1_size`` the device was opened with.
|
| 21 |
+
|
| 22 |
+
The backend follows the ``TTPolicy`` default (the megakernel since 2026-09-18) and ``GR00T_DIT_BACKEND`` /
|
| 23 |
+
``GR00T_MK_ARENA_DTYPE`` like the other whole-model files, so the same test doubles as the ttnn backend's 2-CQ check
|
| 24 |
+
when asked (``test_tt_uploads.py`` covers the ttnn path's CQ-1 writes vs the host gather independently).
|
| 25 |
+
|
| 26 |
+
Run (one pytest process per version, through the device lock; alloc tracking is asserted because traced == untraced
|
| 27 |
+
is)::
|
| 28 |
+
|
| 29 |
+
TT_METAL_TRACE_ALLOC_TRACKING=1 DEVICE_LOCK_TIMEOUT=14400 /home/deepgadget/experiments/gr00t/bin/with-device.sh \\
|
| 30 |
+
python -m pytest -q -rA --timeout=3600 -p no:cacheprovider \\
|
| 31 |
+
models/experimental/gr00t/tests/tt/test_mk_2cq.py --gr00t-version n16 [--sample multi]
|
| 32 |
+
|
| 33 |
+
CPU tier (``-k cpu``): the comparison helper.
|
| 34 |
+
"""
|
| 35 |
+
|
| 36 |
+
from __future__ import annotations
|
| 37 |
+
|
| 38 |
+
import os
|
| 39 |
+
import time
|
| 40 |
+
from typing import Any, Dict, Iterator, List, Mapping, Sequence, Tuple
|
| 41 |
+
|
| 42 |
+
import pytest
|
| 43 |
+
import torch
|
| 44 |
+
|
| 45 |
+
from models.experimental.gr00t.common.golden import GoldenSet
|
| 46 |
+
from models.experimental.gr00t.common.pcc import max_abs, pcc
|
| 47 |
+
from models.experimental.gr00t.tests.tt import harness
|
| 48 |
+
from models.experimental.gr00t.tests.tt.conftest import (
|
| 49 |
+
DEFAULT_L1_SMALL_SIZE,
|
| 50 |
+
DEFAULT_TRACE_REGION_SIZE,
|
| 51 |
+
DEFAULT_TT_METAL_CACHE,
|
| 52 |
+
OPT_SAMPLE,
|
| 53 |
+
SAMPLE_SETS,
|
| 54 |
+
samples_for,
|
| 55 |
+
skip_if_missing,
|
| 56 |
+
)
|
| 57 |
+
from models.experimental.gr00t.tt import model as M
|
| 58 |
+
from models.experimental.gr00t.tt import tracer as T
|
| 59 |
+
from models.experimental.gr00t.tt.layout import StaticShapeError
|
| 60 |
+
from models.experimental.gr00t.tt.policy import TTPolicy
|
| 61 |
+
|
| 62 |
+
torch.set_grad_enabled(False)
|
| 63 |
+
|
| 64 |
+
CQ_CONFIGS: Tuple[int, ...] = (2, 1) # the production path first, then the reference
|
| 65 |
+
N_REPLAYS = 5
|
| 66 |
+
RAW_KEYS: Tuple[str, ...] = ("action_pred", "backbone_features", "K_0", "V_0")
|
| 67 |
+
_RESULTS: Dict[str, Dict[int, Dict[str, Any]]] = {} # version -> num_cqs -> {sample: record, "_facts": {...}}
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 71 |
+
# helpers (CPU-testable)
|
| 72 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 73 |
+
def equal_tensors(a: Mapping[str, torch.Tensor], b: Mapping[str, torch.Tensor], keys: Sequence[str]) -> Dict[str, Any]:
|
| 74 |
+
"""``{key: {"ok", "pcc", "max_abs", "shape"}}`` -- ``ok`` is ``torch.equal`` (same shape, dtype and bits); the
|
| 75 |
+
PCC / max |d| are recorded for the diagnosis when it is not."""
|
| 76 |
+
out: Dict[str, Any] = {}
|
| 77 |
+
for k in keys:
|
| 78 |
+
x, y = a[k], b[k]
|
| 79 |
+
same_shape = tuple(x.shape) == tuple(y.shape)
|
| 80 |
+
ok = bool(same_shape and x.dtype == y.dtype and torch.equal(x, y))
|
| 81 |
+
out[k] = {
|
| 82 |
+
"ok": ok,
|
| 83 |
+
"pcc": pcc(x.float(), y.float()) if same_shape else None,
|
| 84 |
+
"max_abs": max_abs(x.float(), y.float()) if same_shape else None,
|
| 85 |
+
"shape": list(x.shape),
|
| 86 |
+
"dtype": str(x.dtype),
|
| 87 |
+
}
|
| 88 |
+
return out
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
def raw_from_untraced(taps: Mapping[str, torch.Tensor]) -> Dict[str, torch.Tensor]:
|
| 92 |
+
out = {k[len(M.RAW_PREFIX) :]: v for k, v in taps.items() if k.startswith(M.RAW_PREFIX)}
|
| 93 |
+
missing = [k for k in RAW_KEYS if k not in out]
|
| 94 |
+
if missing:
|
| 95 |
+
raise KeyError(f"run_untraced(raw=True) lacks {missing}")
|
| 96 |
+
return out
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
def test_cpu_equal_tensors_helper() -> None:
|
| 100 |
+
a = {k: torch.arange(6, dtype=torch.bfloat16).reshape(2, 3) for k in RAW_KEYS}
|
| 101 |
+
res = equal_tensors(a, {k: t.clone() for k, t in a.items()}, RAW_KEYS)
|
| 102 |
+
assert all(v["ok"] and v["pcc"] == 1.0 and v["max_abs"] == 0.0 for v in res.values())
|
| 103 |
+
b = dict(a)
|
| 104 |
+
b["K_0"] = a["K_0"].clone()
|
| 105 |
+
b["K_0"][0, 0] += 1
|
| 106 |
+
res = equal_tensors(a, b, RAW_KEYS)
|
| 107 |
+
assert not res["K_0"]["ok"] and res["K_0"]["max_abs"] == 1.0 and res["action_pred"]["ok"]
|
| 108 |
+
c = dict(a)
|
| 109 |
+
c["V_0"] = a["V_0"].float() # a dtype change is not bit-equality
|
| 110 |
+
assert not equal_tensors(a, c, RAW_KEYS)["V_0"]["ok"]
|
| 111 |
+
d = dict(a)
|
| 112 |
+
d["action_pred"] = a["action_pred"].reshape(3, 2)
|
| 113 |
+
assert not equal_tensors(a, d, RAW_KEYS)["action_pred"]["ok"]
|
| 114 |
+
with pytest.raises(KeyError):
|
| 115 |
+
raw_from_untraced({M.RAW_PREFIX + "action_pred": a["action_pred"]})
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 119 |
+
# fixtures
|
| 120 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 121 |
+
@pytest.fixture(scope="module")
|
| 122 |
+
def policy(policy_name: str, trace_layout: str) -> TTPolicy:
|
| 123 |
+
"""conftest's ``TTPolicy(dtype_policy=--policy, trace_layout=--trace-layout)`` plus the environment overrides
|
| 124 |
+
(``GR00T_DIT_BACKEND`` / ``GR00T_MK_ARENA_DTYPE``); the code default is the megakernel."""
|
| 125 |
+
return TTPolicy(dtype_policy=policy_name, trace_layout=trace_layout).with_env_overrides()
|
| 126 |
+
|
| 127 |
+
|
| 128 |
+
def _selected_samples(request: pytest.FixtureRequest, version: str) -> Tuple[str, ...]:
|
| 129 |
+
opt = str(request.config.getoption(OPT_SAMPLE))
|
| 130 |
+
return tuple(samples_for(version, opt) if opt in SAMPLE_SETS else (opt,))
|
| 131 |
+
|
| 132 |
+
|
| 133 |
+
@pytest.fixture(scope="module")
|
| 134 |
+
def cq_results(request: pytest.FixtureRequest, version: str, policy: TTPolicy) -> Dict[int, Dict[str, Any]]:
|
| 135 |
+
"""``num_cqs -> {sample: record, "_facts": {...}}`` of every selected sample: one device open per CQ
|
| 136 |
+
configuration (sequentially, module docstring), the model built once per configuration."""
|
| 137 |
+
if version not in _RESULTS:
|
| 138 |
+
if not T.alloc_tracking_requested():
|
| 139 |
+
pytest.fail(
|
| 140 |
+
f"{T.ALLOC_TRACKING_ENV}=1 must be set in the shell before pytest starts (plan §4 / static-shape §10 "
|
| 141 |
+
"item 9): this test asserts traced == untraced"
|
| 142 |
+
)
|
| 143 |
+
samples = _selected_samples(request, version)
|
| 144 |
+
for s in samples:
|
| 145 |
+
skip_if_missing(version, s)
|
| 146 |
+
os.environ.setdefault("TT_METAL_CACHE", str(DEFAULT_TT_METAL_CACHE))
|
| 147 |
+
_RESULTS[version] = {n: _collect(version, samples, policy, n) for n in CQ_CONFIGS}
|
| 148 |
+
request.config._gr00t_device_facts = _RESULTS[version][CQ_CONFIGS[0]]["_facts"]["device"] # type: ignore[attr-defined]
|
| 149 |
+
return _RESULTS[version]
|
| 150 |
+
|
| 151 |
+
|
| 152 |
+
def _collect(version: str, samples: Tuple[str, ...], policy: TTPolicy, num_cqs: int) -> Dict[str, Any]:
|
| 153 |
+
"""Open the device with ``num_cqs`` command queues, build the model, run every sample untraced + traced, release."""
|
| 154 |
+
from models.experimental.gr00t.reference.model import load_golden_sample
|
| 155 |
+
from models.experimental.gr00t.tt.device import close_gr00t_device
|
| 156 |
+
|
| 157 |
+
device = M.open_model_device(
|
| 158 |
+
policy,
|
| 159 |
+
trace_region_size=DEFAULT_TRACE_REGION_SIZE,
|
| 160 |
+
l1_small_size=DEFAULT_L1_SMALL_SIZE,
|
| 161 |
+
num_command_queues=num_cqs,
|
| 162 |
+
)
|
| 163 |
+
out: Dict[str, Any] = {}
|
| 164 |
+
try:
|
| 165 |
+
facts = harness.device_facts(device)
|
| 166 |
+
facts["worker_l1_size"] = M.device_worker_l1_size(device)
|
| 167 |
+
facts["dit_backend"] = policy.dit_backend
|
| 168 |
+
facts["mk_arena_dtype"] = policy.mk_arena_dtype_for(version) if policy.dit_backend == "megakernel" else None
|
| 169 |
+
facts["num_command_queues_requested"] = num_cqs
|
| 170 |
+
t0 = time.perf_counter()
|
| 171 |
+
model = M.Gr00tTT.from_pretrained(version, policy=policy, device=device)
|
| 172 |
+
build_s = time.perf_counter() - t0
|
| 173 |
+
try:
|
| 174 |
+
n_cqs = int(model.executor.backend.num_command_queues(device))
|
| 175 |
+
ex_desc = model.executor.describe()
|
| 176 |
+
print(
|
| 177 |
+
f"\n[test_mk_2cq] {version}/{num_cqs}cq: from_pretrained {build_s:.1f} s (weights "
|
| 178 |
+
f"{model.timing['weights_load_s']:.1f} s, {len(model.W)} tensors, {model.W.device_bytes / 1e6:.0f} MB; "
|
| 179 |
+
f"backend {policy.dit_backend}, arena {model.recipe.mk_arena_dtype}); device reports {n_cqs} CQ(s), "
|
| 180 |
+
f"cq1_uploads {ex_desc.get('cq1_uploads')}"
|
| 181 |
+
)
|
| 182 |
+
out["_facts"] = {
|
| 183 |
+
"device": facts,
|
| 184 |
+
"num_command_queues": n_cqs,
|
| 185 |
+
"cq1_uploads": bool(ex_desc.get("cq1_uploads")),
|
| 186 |
+
"build_s": build_s,
|
| 187 |
+
"weights": {"n_tensors": len(model.W), "device_bytes": float(model.W.device_bytes)},
|
| 188 |
+
"recipe": model.recipe.describe(),
|
| 189 |
+
"head_timing": dict(model.head.timing),
|
| 190 |
+
}
|
| 191 |
+
for sample in samples:
|
| 192 |
+
gs, obs = load_golden_sample(version, sample)
|
| 193 |
+
try:
|
| 194 |
+
mi = model.encode(obs)
|
| 195 |
+
except (ValueError, StaticShapeError) as exc: # another embodiment / layout (the N1.7 extras)
|
| 196 |
+
out[sample] = {"skip": f"{type(exc).__name__}: {exc}"}
|
| 197 |
+
continue
|
| 198 |
+
if mi.shape_key != model.dlayout.shape_key:
|
| 199 |
+
out[sample] = {"skip": f"shape key {mi.shape_key} != model's"}
|
| 200 |
+
continue
|
| 201 |
+
noise = gs.load("initial_noise")
|
| 202 |
+
rec: Dict[str, Any] = {"n_text": int(mi.n_text)}
|
| 203 |
+
pred, taps = model.run_untraced(mi, noise, return_taps=True, raw=True)
|
| 204 |
+
rec["untraced_pred"] = pred
|
| 205 |
+
rec["untraced_raw"] = raw_from_untraced(taps)
|
| 206 |
+
rec["untraced_velocity"] = _velocity_from_taps(model, taps)
|
| 207 |
+
if not model.captured:
|
| 208 |
+
t0 = time.perf_counter()
|
| 209 |
+
model.warm_and_capture(mi, noise)
|
| 210 |
+
rec["warm_and_capture_s"] = time.perf_counter() - t0
|
| 211 |
+
preds: List[torch.Tensor] = []
|
| 212 |
+
states: List[Dict[str, torch.Tensor]] = []
|
| 213 |
+
replay_s: List[float] = []
|
| 214 |
+
for _ in range(N_REPLAYS):
|
| 215 |
+
t0 = time.perf_counter()
|
| 216 |
+
preds.append(model.predict_normalized(mi, noise))
|
| 217 |
+
replay_s.append(time.perf_counter() - t0)
|
| 218 |
+
states.append(model.traced_state())
|
| 219 |
+
rec["traced_preds"], rec["traced_states"], rec["replay_s"] = preds, states, replay_s
|
| 220 |
+
rec["traced_velocity"] = _velocity_traced(model)
|
| 221 |
+
rec["golden_pred"] = gs.load("action_pred_normalized").float()
|
| 222 |
+
out[sample] = rec
|
| 223 |
+
finally:
|
| 224 |
+
model.release()
|
| 225 |
+
finally:
|
| 226 |
+
close_gr00t_device(device)
|
| 227 |
+
return out
|
| 228 |
+
|
| 229 |
+
|
| 230 |
+
def _velocity_from_taps(model: M.Gr00tTT, taps: Mapping[str, torch.Tensor]) -> torch.Tensor:
|
| 231 |
+
"""The untraced per-step decoder output ``[n_steps, M, A_max]`` from the head's golden-shaped taps (the megakernel's
|
| 232 |
+
``velocity_out`` host copies; the ttnn backend's decoder taps under the same key)."""
|
| 233 |
+
keys = [f"action_decoder_out[k={k}]" for k in range(int(model.recipe.n_steps))]
|
| 234 |
+
missing = [k for k in keys if k not in taps]
|
| 235 |
+
if missing:
|
| 236 |
+
raise KeyError(f"run_untraced lacks {missing}")
|
| 237 |
+
return torch.stack([taps[k][0] for k in keys]).contiguous()
|
| 238 |
+
|
| 239 |
+
|
| 240 |
+
def _velocity_traced(model: M.Gr00tTT) -> torch.Tensor:
|
| 241 |
+
"""The per-step decoder output of the last traced replay on the same ``[n_steps, M, A_max]`` slice: the kernel's
|
| 242 |
+
``velocity_out`` tap tensor (megakernel) or ``None`` for the ttnn backend (its decoder outputs are not persistent).
|
| 243 |
+
"""
|
| 244 |
+
if model.policy.dit_backend != "megakernel":
|
| 245 |
+
return torch.empty(0)
|
| 246 |
+
vel = model.head.velocity_taps_host() # [n_steps, 64, A_pad]
|
| 247 |
+
m, a_max = int(model.dlayout.m_logical), int(model.cfg.encoders.max_action_dim)
|
| 248 |
+
return vel[:, :m, :a_max].contiguous()
|
| 249 |
+
|
| 250 |
+
|
| 251 |
+
def _row(results: Any, tap: str, key: str, ok: bool, **extra: Any) -> None:
|
| 252 |
+
results.add_row({"tap": tap, "key": key, "rule": "exact", "ok": bool(ok), **extra})
|
| 253 |
+
|
| 254 |
+
|
| 255 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 256 |
+
# device tier
|
| 257 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 258 |
+
def test_mk_2cq_bit_equality(
|
| 259 |
+
version: str,
|
| 260 |
+
sample: str,
|
| 261 |
+
golden: Tuple[GoldenSet, Any, Any],
|
| 262 |
+
policy: TTPolicy,
|
| 263 |
+
cq_results: Dict[int, Dict[str, Any]],
|
| 264 |
+
results: harness.ResultsWriter,
|
| 265 |
+
) -> None:
|
| 266 |
+
"""2-CQ (CQ-1 input writes) vs 1-CQ: traced == untraced on each device, replays bit-identical, and the two
|
| 267 |
+
devices' results bit-equal (module docstring)."""
|
| 268 |
+
for n in CQ_CONFIGS:
|
| 269 |
+
rec = cq_results[n].get(sample)
|
| 270 |
+
if rec is None or "skip" in rec:
|
| 271 |
+
pytest.skip(f"{version}/{sample}/{n}cq: {rec['skip'] if rec else 'not collected'}")
|
| 272 |
+
facts = {n: cq_results[n]["_facts"] for n in CQ_CONFIGS}
|
| 273 |
+
results.set_meta(
|
| 274 |
+
sample=sample,
|
| 275 |
+
policy=policy.to_dict(),
|
| 276 |
+
alloc_tracking_env=os.environ.get(T.ALLOC_TRACKING_ENV),
|
| 277 |
+
n_replays=N_REPLAYS,
|
| 278 |
+
cq_facts={str(n): {k: v for k, v in f.items() if k != "device"} for n, f in facts.items()},
|
| 279 |
+
device_2cq=facts[2]["device"],
|
| 280 |
+
device_1cq=facts[1]["device"],
|
| 281 |
+
n_text=cq_results[2][sample]["n_text"],
|
| 282 |
+
)
|
| 283 |
+
ok_all = True
|
| 284 |
+
# the CQ facts themselves: the production configuration really ran with two queues and the CQ-1 writes on
|
| 285 |
+
fact_ok = (
|
| 286 |
+
facts[2]["num_command_queues"] == 2
|
| 287 |
+
and facts[2]["cq1_uploads"] is True
|
| 288 |
+
and facts[1]["num_command_queues"] == 1
|
| 289 |
+
and facts[1]["cq1_uploads"] is False
|
| 290 |
+
)
|
| 291 |
+
_row(results, "device", "cq_facts", fact_ok, **{f"{n}cq": facts[n]["cq1_uploads"] for n in CQ_CONFIGS})
|
| 292 |
+
ok_all = ok_all and fact_ok
|
| 293 |
+
for n in CQ_CONFIGS:
|
| 294 |
+
rec = cq_results[n][sample]
|
| 295 |
+
st0 = rec["traced_states"][0]
|
| 296 |
+
eq_raw = equal_tensors(st0, rec["untraced_raw"], RAW_KEYS)
|
| 297 |
+
for k, r in eq_raw.items():
|
| 298 |
+
_row(results, k, f"{n}cq-traced_equals_untraced", r["ok"], pcc=r["pcc"], max_abs=r["max_abs"])
|
| 299 |
+
eq_host = bool(torch.equal(rec["traced_preds"][0], rec["untraced_pred"]))
|
| 300 |
+
_row(results, "action_pred_normalized", f"{n}cq-traced_equals_untraced", eq_host)
|
| 301 |
+
eq_replays = all(
|
| 302 |
+
all(equal_tensors(st0, s, RAW_KEYS)[k]["ok"] for k in RAW_KEYS) for s in rec["traced_states"][1:]
|
| 303 |
+
) and all(torch.equal(p, rec["traced_preds"][0]) for p in rec["traced_preds"][1:])
|
| 304 |
+
_row(
|
| 305 |
+
results,
|
| 306 |
+
"action_pred",
|
| 307 |
+
f"{n}cq-{N_REPLAYS}_replays_bit_identical",
|
| 308 |
+
eq_replays,
|
| 309 |
+
replay_ms=[round(s * 1e3, 3) for s in rec["replay_s"]],
|
| 310 |
+
)
|
| 311 |
+
ok_all = ok_all and all(r["ok"] for r in eq_raw.values()) and eq_host and eq_replays
|
| 312 |
+
if rec["traced_velocity"].numel():
|
| 313 |
+
eq_vel = bool(torch.equal(rec["traced_velocity"], rec["untraced_velocity"]))
|
| 314 |
+
_row(
|
| 315 |
+
results,
|
| 316 |
+
"velocity_out",
|
| 317 |
+
f"{n}cq-traced_equals_untraced",
|
| 318 |
+
eq_vel,
|
| 319 |
+
shape=list(rec["traced_velocity"].shape),
|
| 320 |
+
)
|
| 321 |
+
ok_all = ok_all and eq_vel
|
| 322 |
+
print(
|
| 323 |
+
f"[{version}/{sample}] {n}cq: traced==untraced {all(r['ok'] for r in eq_raw.values()) and eq_host}, "
|
| 324 |
+
f"{N_REPLAYS} replays identical {eq_replays}, replay+D2H ms {[round(s * 1e3, 2) for s in rec['replay_s']]}"
|
| 325 |
+
)
|
| 326 |
+
# across the two devices: untraced, traced, host prediction and the per-step velocities
|
| 327 |
+
two, one = cq_results[2][sample], cq_results[1][sample]
|
| 328 |
+
cross_untraced = equal_tensors(two["untraced_raw"], one["untraced_raw"], RAW_KEYS)
|
| 329 |
+
cross_traced = equal_tensors(two["traced_states"][0], one["traced_states"][0], RAW_KEYS)
|
| 330 |
+
for k in RAW_KEYS:
|
| 331 |
+
_row(
|
| 332 |
+
results,
|
| 333 |
+
k,
|
| 334 |
+
"2cq_equals_1cq-untraced",
|
| 335 |
+
cross_untraced[k]["ok"],
|
| 336 |
+
**{x: cross_untraced[k][x] for x in ("pcc", "max_abs")},
|
| 337 |
+
)
|
| 338 |
+
_row(
|
| 339 |
+
results,
|
| 340 |
+
k,
|
| 341 |
+
"2cq_equals_1cq-traced",
|
| 342 |
+
cross_traced[k]["ok"],
|
| 343 |
+
**{x: cross_traced[k][x] for x in ("pcc", "max_abs")},
|
| 344 |
+
)
|
| 345 |
+
eq_pred = bool(torch.equal(two["traced_preds"][0], one["traced_preds"][0]))
|
| 346 |
+
_row(
|
| 347 |
+
results,
|
| 348 |
+
"action_pred_normalized",
|
| 349 |
+
"2cq_equals_1cq-host_pred",
|
| 350 |
+
eq_pred,
|
| 351 |
+
pcc=pcc(two["traced_preds"][0], one["traced_preds"][0]),
|
| 352 |
+
max_abs=max_abs(two["traced_preds"][0], one["traced_preds"][0]),
|
| 353 |
+
)
|
| 354 |
+
eq_vel_cross = True
|
| 355 |
+
if two["traced_velocity"].numel() and one["traced_velocity"].numel():
|
| 356 |
+
eq_vel_cross = bool(torch.equal(two["traced_velocity"], one["traced_velocity"]))
|
| 357 |
+
_row(results, "velocity_out", "2cq_equals_1cq-traced", eq_vel_cross)
|
| 358 |
+
ok_cross = all(r["ok"] for r in cross_untraced.values()) and all(r["ok"] for r in cross_traced.values())
|
| 359 |
+
ok_all = ok_all and ok_cross and eq_pred and eq_vel_cross
|
| 360 |
+
# sanity, not the subject: the prediction is a valid answer for this sample (the e2e files gate it)
|
| 361 |
+
p_gold = pcc(two["traced_preds"][0], two["golden_pred"])
|
| 362 |
+
results.add_row(
|
| 363 |
+
{
|
| 364 |
+
"tap": "action_pred_normalized",
|
| 365 |
+
"key": "2cq-vs-golden",
|
| 366 |
+
"rule": "info",
|
| 367 |
+
"ok": bool(p_gold > 0.99),
|
| 368 |
+
"pcc": p_gold,
|
| 369 |
+
}
|
| 370 |
+
)
|
| 371 |
+
results.set_timing(
|
| 372 |
+
replay_median_2cq_s=sorted(two["replay_s"])[len(two["replay_s"]) // 2],
|
| 373 |
+
replay_median_1cq_s=sorted(one["replay_s"])[len(one["replay_s"]) // 2],
|
| 374 |
+
build_2cq_s=facts[2]["build_s"],
|
| 375 |
+
build_1cq_s=facts[1]["build_s"],
|
| 376 |
+
)
|
| 377 |
+
print(
|
| 378 |
+
f"[{version}/{sample}] 2cq == 1cq: untraced {all(r['ok'] for r in cross_untraced.values())}, traced "
|
| 379 |
+
f"{all(r['ok'] for r in cross_traced.values())}, host pred {eq_pred}, velocity {eq_vel_cross}; "
|
| 380 |
+
f"pcc vs golden {p_gold:.6f}"
|
| 381 |
+
)
|
| 382 |
+
assert fact_ok, {n: facts[n] for n in CQ_CONFIGS}
|
| 383 |
+
assert ok_all, "see the rows: a CQ configuration is not bit-equal to its untraced run or to the other one"
|
| 384 |
+
assert p_gold > 0.99
|
code/models/experimental/gr00t/tests/tt/test_mk_cpu.py
CHANGED
|
@@ -174,7 +174,8 @@ def test_core_map_vc_and_per_core_args(core_map: CM.CoreMap):
|
|
| 174 |
for c, core in enumerate(core_map.compute):
|
| 175 |
g = CM.group_of(c)
|
| 176 |
assert args[core] == (c, g, core_map.vc[g], c % 32, g)
|
| 177 |
-
assert args[core_map.hub] == (0xFFFF,) *
|
|
|
|
| 178 |
|
| 179 |
|
| 180 |
def test_core_map_other_worker_lists_and_errors():
|
|
@@ -612,7 +613,10 @@ def test_descriptor_spec_contents(core_map):
|
|
| 612 |
assert ctas[name] == cb_id
|
| 613 |
for name, sem_id in D.SEM_IDS.items():
|
| 614 |
assert ctas[name] == sem_id
|
| 615 |
-
|
|
|
|
|
|
|
|
|
|
| 616 |
# common RT args: addresses (0 in the dry run), step index, offsets table
|
| 617 |
io = spec.io_order()
|
| 618 |
assert io[:3] == ["x_rep", "a_rep", "arena_stream"] and io[-3:] == ["x_t", "state_features", "action_pred"]
|
|
|
|
| 174 |
for c, core in enumerate(core_map.compute):
|
| 175 |
g = CM.group_of(c)
|
| 176 |
assert args[core] == (c, g, core_map.vc[g], c % 32, g)
|
| 177 |
+
assert args[core_map.hub] == (0xFFFF,) * 4 + (0,) # K4b: the group word of a hub is its hub index
|
| 178 |
+
assert args[core_map.hub2] == (0xFFFF,) * 4 + (1,)
|
| 179 |
|
| 180 |
|
| 181 |
def test_core_map_other_worker_lists_and_errors():
|
|
|
|
| 613 |
assert ctas[name] == cb_id
|
| 614 |
for name, sem_id in D.SEM_IDS.items():
|
| 615 |
assert ctas[name] == sem_id
|
| 616 |
+
# 9 K1 semaphores + the 4 K4b ones (hub done / odd ready counter / hub-1 flag / hub-1 landing base)
|
| 617 |
+
n_sems = len(D.SEM_IDS)
|
| 618 |
+
assert n_sems == 13 and len(spec.semaphores) == n_sems
|
| 619 |
+
assert sorted(s.sem_id for s in spec.semaphores) == list(range(n_sems))
|
| 620 |
# common RT args: addresses (0 in the dry run), step index, offsets table
|
| 621 |
io = spec.io_order()
|
| 622 |
assert io[:3] == ["x_rep", "a_rep", "arena_stream"] and io[-3:] == ["x_t", "state_features", "action_pred"]
|
code/models/experimental/gr00t/tests/tt/test_mk_dit_block.py
ADDED
|
@@ -0,0 +1,398 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
"""WP-K4 device tests of the megakernel's full DiT blocks (``block`` variant; IMPLEMENTATION_PLAN.md §5.4 P1-P12,
|
| 5 |
+
§5.7 K4 acceptance rows ``dit_block0_out`` / ``dit_block1_out``).
|
| 6 |
+
|
| 7 |
+
* ``test_kernel_cta_names_cpu`` -- every ``get_named_compile_time_arg_val`` name of ``dit_kernel.cpp`` / ``ops/*.hpp``
|
| 8 |
+
is emitted by the K4 program spec (or by ``materialise``), so a missing CTA fails on the CPU, not as a JIT error;
|
| 9 |
+
* ``test_k4_spec_cpu`` -- the K4 spec of every version / stream dtype keeps the descriptor limits, the tail geometry
|
| 10 |
+
mirror and the tile-row ``x_rep`` shard bookkeeping;
|
| 11 |
+
* ``test_block_vs_oracle`` -- one self block (block 1 from the golden ``dit_block0_out[k]``) and one cross block
|
| 12 |
+
(block 0 from the golden ``sa_embs[k]``) with the real bf16 weights, device-hoisted K/V and the call's masks, steps 0
|
| 13 |
+
and 3: PCC >= 0.9999 against the Stage-1 ``DiTSelfBlockTT`` / ``DiTCrossBlockTT`` on device (all 64 rows), the
|
| 14 |
+
golden ``dit_block{0,1}_out[k]`` gates of ``gates_multi.json`` (rows < M), every core's resident tile row
|
| 15 |
+
bit-identical, and a second launch bit-identical;
|
| 16 |
+
* ``test_blocks_chain_vs_golden`` -- blocks [0, 2) in one launch from ``sa_embs[k]`` vs ``dit_block1_out[k]`` and
|
| 17 |
+
vs the chained oracle.
|
| 18 |
+
|
| 19 |
+
Run (one pytest process per file, through the device lock)::
|
| 20 |
+
|
| 21 |
+
DEVICE_LOCK_TIMEOUT=14400 bin/with-device.sh python -m pytest -q -rA --timeout=1800 -p no:cacheprovider \\
|
| 22 |
+
models/experimental/gr00t/tests/tt/test_mk_dit_block.py --gr00t-version n16
|
| 23 |
+
"""
|
| 24 |
+
|
| 25 |
+
from __future__ import annotations
|
| 26 |
+
|
| 27 |
+
import os
|
| 28 |
+
import re
|
| 29 |
+
import time
|
| 30 |
+
from pathlib import Path
|
| 31 |
+
from typing import Any, Dict, Iterator, List, Optional, Sequence, Tuple
|
| 32 |
+
|
| 33 |
+
import pytest
|
| 34 |
+
import torch
|
| 35 |
+
|
| 36 |
+
from models.experimental.gr00t.common import pcc as pcc_mod
|
| 37 |
+
from models.experimental.gr00t.common.configs import DIT_QUERY_ROWS_PAD, GR00TConfig, get_config
|
| 38 |
+
from models.experimental.gr00t.common.golden import GoldenSet
|
| 39 |
+
from models.experimental.gr00t.tests.tt.harness import check_taps
|
| 40 |
+
from models.experimental.gr00t.tt.megakernel import arena as A
|
| 41 |
+
from models.experimental.gr00t.tt.megakernel import descriptors as D
|
| 42 |
+
from models.experimental.gr00t.tt.megakernel import dit_program as P
|
| 43 |
+
from models.experimental.gr00t.tt.megakernel.core_map import CoreMap, default_core_map
|
| 44 |
+
|
| 45 |
+
torch.set_grad_enabled(False)
|
| 46 |
+
|
| 47 |
+
KERNEL_DIR = Path(P.__file__).resolve().parent / "kernels"
|
| 48 |
+
GATE_ORACLE = 0.9999 # plan §5.7 K4: PCC vs the S1 device oracle
|
| 49 |
+
STEPS_TESTED: Tuple[int, ...] = (0, 3)
|
| 50 |
+
SELF_BLK, CROSS_BLK = 1, 0
|
| 51 |
+
#: arena stream dtype of the session (bring-up order: bf16 first, bfp8 after G2); ``GR00T_MK_STREAM_DTYPE=bfp8_b``
|
| 52 |
+
#: runs the same tests on the bfp8 arena with the ``mixed_dit`` (bfp8 DiT weights) Stage-1 oracle
|
| 53 |
+
STREAM_DTYPE = os.environ.get("GR00T_MK_STREAM_DTYPE", "bf16")
|
| 54 |
+
if STREAM_DTYPE not in ("bf16", "bfp8_b"):
|
| 55 |
+
raise ValueError(f"GR00T_MK_STREAM_DTYPE must be bf16 or bfp8_b, got {STREAM_DTYPE!r}")
|
| 56 |
+
HEAD_CATEGORIES: Tuple[str, ...] = ("dit", "dit_precompute", "encoders")
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 60 |
+
# host helpers
|
| 61 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 62 |
+
def pad_rows(t: torch.Tensor) -> torch.Tensor:
|
| 63 |
+
"""Golden ``[1, M, C]`` fp32 -> ``[64, C]`` bf16-representable fp32 with exactly-zero pad rows (plan D5)."""
|
| 64 |
+
if t.ndim != 3 or t.shape[0] != 1 or t.shape[1] > DIT_QUERY_ROWS_PAD:
|
| 65 |
+
raise ValueError(f"expected [1, M <= 64, C], got {tuple(t.shape)}")
|
| 66 |
+
out = torch.zeros(DIT_QUERY_ROWS_PAD, t.shape[2], dtype=torch.float32)
|
| 67 |
+
out[: t.shape[1]] = t[0].to(torch.bfloat16).float()
|
| 68 |
+
return out
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def kernel_cta_names(kernel_dir: Path = KERNEL_DIR) -> List[str]:
|
| 72 |
+
"""Every distinct ``get_named_compile_time_arg_val("<name>")`` of the megakernel sources (kernel + ops headers)."""
|
| 73 |
+
pat = re.compile(r'get_named_compile_time_arg_val\("([A-Za-z0-9_]+)"\)')
|
| 74 |
+
names: List[str] = []
|
| 75 |
+
for f in sorted(list(kernel_dir.glob("*.cpp")) + list((kernel_dir / "ops").glob("*.hpp"))):
|
| 76 |
+
for m in pat.finditer(f.read_text()):
|
| 77 |
+
if m.group(1) not in names:
|
| 78 |
+
names.append(m.group(1))
|
| 79 |
+
return names
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def oracle_policy() -> Any:
|
| 83 |
+
from models.experimental.gr00t.tt.policy import TTPolicy
|
| 84 |
+
|
| 85 |
+
return TTPolicy(dtype_policy="bf16" if STREAM_DTYPE == "bf16" else "mixed_dit")
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 89 |
+
# CPU
|
| 90 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 91 |
+
def test_kernel_cta_names_cpu() -> None:
|
| 92 |
+
"""The K4 spec (+ ``fill_missing_ctas`` + ``materialise``'s NoC names) covers every CTA the sources read."""
|
| 93 |
+
cm = default_core_map()
|
| 94 |
+
L = A.ArenaLayout.plan("n16", dtype_stream="bf16", reader_mode="direct", tile_order="row_major")
|
| 95 |
+
spec = D.build_program_spec_k4(cm, L, "steps4", {}, dry_run=True)
|
| 96 |
+
spec.named_ctas += [(f"mk_io_{n}", 0) for n in P.MK_IO_CTAS] + [
|
| 97 |
+
("mk_k1", 0),
|
| 98 |
+
("mk_ring_depth", 4),
|
| 99 |
+
("mk_relay_window", 6),
|
| 100 |
+
("mk_rt_base", 0),
|
| 101 |
+
("mk_io_kv_base", 0),
|
| 102 |
+
]
|
| 103 |
+
P.fill_missing_ctas(spec)
|
| 104 |
+
emitted = {n for n, _ in spec.named_ctas} | {flag for flag, _ in spec.role_flags}
|
| 105 |
+
emitted |= set(P.MATERIALISE_CTAS)
|
| 106 |
+
names = kernel_cta_names()
|
| 107 |
+
missing = [n for n in names if n not in emitted]
|
| 108 |
+
assert len(names) > 60, names
|
| 109 |
+
assert not missing, f"CTAs read by the kernel sources but never emitted: {missing}"
|
| 110 |
+
# and the union list itself only names things the sources read or the K1 spec emits
|
| 111 |
+
k1 = D.build_program_spec(cm, L, "block", {}, dry_run=True)
|
| 112 |
+
k1_emitted = {n for n, _ in k1.named_ctas}
|
| 113 |
+
stale = [n for n in P.KERNEL_CTA_UNION if n not in names and n not in k1_emitted]
|
| 114 |
+
assert not stale, f"KERNEL_CTA_UNION entries no source reads: {stale}"
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
@pytest.mark.parametrize("dtype", ("bf16", "bfp8_b"))
|
| 118 |
+
def test_k4_spec_cpu(version: str, dtype: str) -> None:
|
| 119 |
+
"""Descriptor limits, tail geometry mirror and the x_rep tile-row bookkeeping of the K4 spec."""
|
| 120 |
+
cm = default_core_map()
|
| 121 |
+
L = A.ArenaLayout.plan(version, dtype_stream=dtype, reader_mode="direct", tile_order="row_major")
|
| 122 |
+
assert P.check_kernel_geometry(L) == 4 * L.n_blocks
|
| 123 |
+
assert P.check_tail_geometry(L) == len(P.TAIL_OP_NAMES)
|
| 124 |
+
for variant in D.VARIANTS:
|
| 125 |
+
spec = D.build_program_spec_k4(cm, L, variant, {}, dry_run=True)
|
| 126 |
+
st = spec.stats()
|
| 127 |
+
assert st.n_common_rt_args + len(P.MK4_RT) + cm.n_compute <= D.MAX_COMMON_RT_ARGS, st
|
| 128 |
+
assert st.n_kernel_groups <= D.MAX_KERNEL_GROUPS and st.n_cbs <= D.MAX_CBS
|
| 129 |
+
assert st.l1_bytes_compute <= D.K4Budget().l1_budget_compute
|
| 130 |
+
cb = {c.name: c for c in spec.cbs}
|
| 131 |
+
assert cb["cb_x"].total_bytes == 48 * 2048 and cb["cb_x"].tensor_backed == "x_rep"
|
| 132 |
+
assert cb["cb_a"].total_bytes == 128 * 2048
|
| 133 |
+
assert cb["cb_stat32"].cb_id == D.K4_HUB_CB_IDS["cb_act"]
|
| 134 |
+
# the tile row of every compute core (dist_layernorm.hpp: r = c // 48, t = c % 48)
|
| 135 |
+
order = P.shard_order(cm)
|
| 136 |
+
assert sorted(order) == list(range(cm.n_compute))
|
| 137 |
+
assert all((c // 48) in (0, 1) for c in order) and sum(1 for c in order if c // 48 == 0) == 48
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 141 |
+
# device session (one arena + weights + hoisted K/V per version)
|
| 142 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 143 |
+
class K4Session:
|
| 144 |
+
"""Real bf16 weights on device (``TTWeights``, the Stage-1 oracle's), the same plan tensors packed into the
|
| 145 |
+
megakernel arena, the call's masks and the device-hoisted K/V of the golden sample, one ``DiTMegakernel``."""
|
| 146 |
+
|
| 147 |
+
def __init__(self, version: str, golden: Tuple[GoldenSet, Any, Any], device: Any, variant: str = "block") -> None:
|
| 148 |
+
import ttnn
|
| 149 |
+
from models.experimental.gr00t.tests.tt.test_tt_dit_ops import check_timesteps, gathered_vl_sets, hoist_kv
|
| 150 |
+
from models.experimental.gr00t.tt import layers as TL
|
| 151 |
+
from models.experimental.gr00t.tt import layout as LAY
|
| 152 |
+
from models.experimental.gr00t.tt import weights as TW
|
| 153 |
+
from models.experimental.gr00t.tt.dit_ops import dit_masks
|
| 154 |
+
|
| 155 |
+
gs, obs, mi = golden
|
| 156 |
+
self.gs, self.mi = gs, mi
|
| 157 |
+
self.cfg: GR00TConfig = get_config(version)
|
| 158 |
+
self.dl = LAY.DeviceLayout.for_inputs(self.cfg, mi)
|
| 159 |
+
self.dl.assert_call(mi)
|
| 160 |
+
check_timesteps(gs, self.cfg)
|
| 161 |
+
if self.dl.layout_name != self.cfg.canonical_layout.name:
|
| 162 |
+
raise RuntimeError(f"{version}: golden sample layout {self.dl.layout_name} is not the canonical layout")
|
| 163 |
+
self.device = device
|
| 164 |
+
self.policy = oracle_policy()
|
| 165 |
+
t0 = time.perf_counter()
|
| 166 |
+
entries = TW.upload_plan(version, None, self.policy)
|
| 167 |
+
names = sorted(n for n, e in entries.items() if e.category in HEAD_CATEGORIES)
|
| 168 |
+
self.W = TW.TTWeights.load(version, None, self.policy, device, names=names)
|
| 169 |
+
self.weights_s = time.perf_counter() - t0
|
| 170 |
+
# masks (materialise dictionary) and hoisted K/V from the golden VL rows (the WP-10 hoist recipe on device)
|
| 171 |
+
self.masks = LAY.materialise(device, self.dl, mi.slots)
|
| 172 |
+
self.kv = hoist_kv(
|
| 173 |
+
self.cfg, self.dl, self.W, gathered_vl_sets(gs, mi, self.dl), lambda t: TL.to_device(t, device, mem="DRAM")
|
| 174 |
+
)
|
| 175 |
+
# the megakernel arena from the same host plan tensors
|
| 176 |
+
t0 = time.perf_counter()
|
| 177 |
+
self.arena = A.ArenaLayout.plan(
|
| 178 |
+
version, dtype_stream=STREAM_DTYPE, reader_mode="direct", tile_order="row_major"
|
| 179 |
+
)
|
| 180 |
+
plan = P.real_arena_plan(version, policy=self.policy)
|
| 181 |
+
self.pack = self.arena.pack(plan)
|
| 182 |
+
self.cm = CoreMap.from_device(device)
|
| 183 |
+
# GR00T_MK_CB_W_SLOTS overrides the weight-ring depth (K-block slots) for the G2 sweeps (K4Budget default 13 bf16 / 15 bfp8)
|
| 184 |
+
slots = os.environ.get("GR00T_MK_CB_W_SLOTS")
|
| 185 |
+
k4_budget = D.K4Budget(cb_w_slots={STREAM_DTYPE: int(slots)}) if slots else None
|
| 186 |
+
self.mk = P.DiTMegakernel(
|
| 187 |
+
self.cfg, self.dl, None, self.arena, self.cm, device, variant=variant, k1_extras=False, k4_budget=k4_budget
|
| 188 |
+
)
|
| 189 |
+
self.mk.upload(self.pack)
|
| 190 |
+
self.mk.bind_kv(self.kv)
|
| 191 |
+
self.mk.bind_masks(dit_masks(self.masks))
|
| 192 |
+
ttnn.synchronize_device(device)
|
| 193 |
+
self.arena_s = time.perf_counter() - t0
|
| 194 |
+
self.dit_masks = dit_masks(self.masks)
|
| 195 |
+
self._blocks: Dict[int, Any] = {}
|
| 196 |
+
|
| 197 |
+
def oracle_block(self, i: int) -> Any:
|
| 198 |
+
from models.experimental.gr00t.tt.dit_ops import DiTCrossBlockTT, DiTSelfBlockTT
|
| 199 |
+
|
| 200 |
+
if i not in self._blocks:
|
| 201 |
+
if self.arena.is_cross(i):
|
| 202 |
+
self._blocks[i] = DiTCrossBlockTT(self.cfg, i, self.W, self.dl, self.masks, self.arena.key_sets[i])
|
| 203 |
+
else:
|
| 204 |
+
self._blocks[i] = DiTSelfBlockTT(self.cfg, i, self.W, self.dl, self.masks)
|
| 205 |
+
return self._blocks[i]
|
| 206 |
+
|
| 207 |
+
def oracle_blocks(self, x: torch.Tensor, blocks: Sequence[int], step: int) -> torch.Tensor:
|
| 208 |
+
"""Stage-1 device path: blocks ``blocks`` in order from the host ``x [64, 1536]`` -> host fp32 ``[64, 1536]``."""
|
| 209 |
+
import ttnn
|
| 210 |
+
from models.experimental.gr00t.tt import layers as TL
|
| 211 |
+
|
| 212 |
+
cur = TL.to_device(x.reshape(1, 1, DIT_QUERY_ROWS_PAD, -1).to(torch.bfloat16), self.device, mem="L1")
|
| 213 |
+
for i in blocks:
|
| 214 |
+
blk = self.oracle_block(i)
|
| 215 |
+
nxt = blk.forward(cur, step, self.kv[i] if self.arena.is_cross(i) else None)
|
| 216 |
+
ttnn.deallocate(cur)
|
| 217 |
+
cur = nxt
|
| 218 |
+
out = TL.to_torch(cur).reshape(DIT_QUERY_ROWS_PAD, -1).float()
|
| 219 |
+
ttnn.deallocate(cur)
|
| 220 |
+
return out
|
| 221 |
+
|
| 222 |
+
def golden_rows(self, tap: str, step: int) -> torch.Tensor:
|
| 223 |
+
return pad_rows(self.gs.load(tap, step=step))
|
| 224 |
+
|
| 225 |
+
def release(self) -> None:
|
| 226 |
+
import ttnn
|
| 227 |
+
|
| 228 |
+
self.mk.deallocate()
|
| 229 |
+
for t in self.masks.values():
|
| 230 |
+
if t is not None:
|
| 231 |
+
ttnn.deallocate(t)
|
| 232 |
+
for k, v in self.kv.values():
|
| 233 |
+
ttnn.deallocate(k)
|
| 234 |
+
ttnn.deallocate(v)
|
| 235 |
+
self.W.release()
|
| 236 |
+
|
| 237 |
+
|
| 238 |
+
_SESSIONS: Dict[str, K4Session] = {}
|
| 239 |
+
|
| 240 |
+
|
| 241 |
+
@pytest.fixture(scope="module")
|
| 242 |
+
def k4_device() -> Iterator[Any]:
|
| 243 |
+
"""The p150a opened with the smaller ``worker_l1_size`` the K4 binaries need (``dit_program.open_k4_device``: the
|
| 244 |
+
compute-core BRISC + NCRISC + 3 TRISC binaries are ~91 KB, the default kernel-config ring buffer 70,656 B). The
|
| 245 |
+
session-scoped device fixture of conftest must not be open in the same process (one device at a time)."""
|
| 246 |
+
import os
|
| 247 |
+
|
| 248 |
+
from models.experimental.gr00t.tests.tt.conftest import DEFAULT_TT_METAL_CACHE
|
| 249 |
+
from models.experimental.gr00t.tt.device import close_gr00t_device
|
| 250 |
+
|
| 251 |
+
os.environ.setdefault("TT_METAL_CACHE", str(DEFAULT_TT_METAL_CACHE))
|
| 252 |
+
device = P.open_k4_device()
|
| 253 |
+
try:
|
| 254 |
+
yield device
|
| 255 |
+
finally:
|
| 256 |
+
for v, s in list(_SESSIONS.items()):
|
| 257 |
+
s.release()
|
| 258 |
+
del _SESSIONS[v]
|
| 259 |
+
close_gr00t_device(device)
|
| 260 |
+
|
| 261 |
+
|
| 262 |
+
@pytest.fixture(scope="module")
|
| 263 |
+
def k4(version: str, golden: Tuple[GoldenSet, Any, Any], k4_device: Any) -> Iterator[K4Session]:
|
| 264 |
+
for v, s in list(_SESSIONS.items()):
|
| 265 |
+
if v != version:
|
| 266 |
+
s.release()
|
| 267 |
+
del _SESSIONS[v]
|
| 268 |
+
if version not in _SESSIONS:
|
| 269 |
+
_SESSIONS[version] = K4Session(version, golden, k4_device)
|
| 270 |
+
yield _SESSIONS[version]
|
| 271 |
+
|
| 272 |
+
|
| 273 |
+
def _row(results: Any, key: str, got: torch.Tensor, ref: torch.Tensor, gate: float, **extra: Any) -> Dict[str, Any]:
|
| 274 |
+
m = pcc_mod.metrics(got.float(), ref.float())
|
| 275 |
+
row = {
|
| 276 |
+
"tap": "mk_k4_block",
|
| 277 |
+
"key": key,
|
| 278 |
+
"rule": "pcc",
|
| 279 |
+
"pcc": m["pcc"],
|
| 280 |
+
"gate_pcc": gate,
|
| 281 |
+
"max_abs": m["max_abs"],
|
| 282 |
+
"rel_l2": m["rel_l2"],
|
| 283 |
+
"finite": m["finite"],
|
| 284 |
+
"ok": bool(m["finite"] and m["pcc"] >= gate),
|
| 285 |
+
}
|
| 286 |
+
row.update(extra)
|
| 287 |
+
results.add_row(row)
|
| 288 |
+
return row
|
| 289 |
+
|
| 290 |
+
|
| 291 |
+
def _session_meta(results: Any, s: K4Session) -> None:
|
| 292 |
+
from models.experimental.gr00t.tests.tt.harness import device_facts
|
| 293 |
+
|
| 294 |
+
results.set_device_facts(device_facts(s.device))
|
| 295 |
+
results.set_meta(
|
| 296 |
+
worker_l1_size=P.K4_WORKER_L1_SIZE,
|
| 297 |
+
layout=s.dl.layout_name,
|
| 298 |
+
embodiment_id=s.dl.embodiment_id,
|
| 299 |
+
stream_dtype=STREAM_DTYPE,
|
| 300 |
+
weights_load_s=s.weights_s,
|
| 301 |
+
arena_pack_upload_s=s.arena_s,
|
| 302 |
+
key_sets=list(s.arena.key_sets[:4]),
|
| 303 |
+
stats=s.mk.stats().__dict__ if s.mk.last_spec is not None else None,
|
| 304 |
+
compile_s={str(k): v for k, v in s.mk.compile_seconds.items()},
|
| 305 |
+
)
|
| 306 |
+
|
| 307 |
+
|
| 308 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 309 |
+
# device
|
| 310 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 311 |
+
@pytest.mark.parametrize("step", STEPS_TESTED)
|
| 312 |
+
@pytest.mark.parametrize("kind", ("self", "cross"))
|
| 313 |
+
def test_block_vs_oracle(results: Any, version: str, k4: K4Session, kind: str, step: int) -> None:
|
| 314 |
+
"""One full block (P1-P12) from a golden input vs the S1 block on device (all 64 rows, PCC >= 0.9999) and vs the
|
| 315 |
+
golden block output (rows < M, ``gates_multi.json``); every core's tile row bit-identical; deterministic."""
|
| 316 |
+
s = k4
|
| 317 |
+
blk = SELF_BLK if kind == "self" else CROSS_BLK
|
| 318 |
+
assert s.arena.is_cross(blk) == (kind == "cross")
|
| 319 |
+
x = s.golden_rows("sa_embs" if blk == 0 else f"dit_block{blk - 1}_out", step)
|
| 320 |
+
x_rows = s.mk.make_x_rows(x)
|
| 321 |
+
try:
|
| 322 |
+
t0 = time.perf_counter()
|
| 323 |
+
got = s.mk.run_block(x_rows, blk, matmul_only=False, step=step)
|
| 324 |
+
dt = time.perf_counter() - t0
|
| 325 |
+
all_rows = s.mk.read_x_rows_all(x_rows)
|
| 326 |
+
# second launch from the same input: bit-identical
|
| 327 |
+
x_rows2 = s.mk.make_x_rows(x)
|
| 328 |
+
try:
|
| 329 |
+
got2 = s.mk.run_block(x_rows2, blk, matmul_only=False, step=step)
|
| 330 |
+
finally:
|
| 331 |
+
s.mk._ttnn.deallocate(x_rows2)
|
| 332 |
+
finally:
|
| 333 |
+
s.mk._ttnn.deallocate(x_rows)
|
| 334 |
+
ref = s.oracle_blocks(x, [blk], step)
|
| 335 |
+
m = s.dl.m_logical
|
| 336 |
+
row = _row(
|
| 337 |
+
results,
|
| 338 |
+
f"{kind}-blk{blk}-k{step}-vs-oracle",
|
| 339 |
+
got,
|
| 340 |
+
ref,
|
| 341 |
+
GATE_ORACLE,
|
| 342 |
+
blk=blk,
|
| 343 |
+
step=step,
|
| 344 |
+
kind=kind,
|
| 345 |
+
pcc_valid_rows=pcc_mod.pcc(got[:m], ref[:m]),
|
| 346 |
+
wall_s=dt,
|
| 347 |
+
)
|
| 348 |
+
rows_ok = all(torch.equal(all_rows[c], all_rows[48 * (c // 48)]) for c in range(s.cm.n_compute))
|
| 349 |
+
results.add_row(
|
| 350 |
+
{
|
| 351 |
+
"tap": "mk_k4_block",
|
| 352 |
+
"key": f"{kind}-blk{blk}-k{step}-consistency",
|
| 353 |
+
"rule": "exact",
|
| 354 |
+
"rows_bit_identical": rows_ok,
|
| 355 |
+
"deterministic": bool(torch.equal(got, got2)),
|
| 356 |
+
"ok": bool(rows_ok and torch.equal(got, got2)),
|
| 357 |
+
}
|
| 358 |
+
)
|
| 359 |
+
taps = {s.gs.threshold_key(f"dit_block{blk}_out", step=step): got[:m].reshape(1, m, -1)}
|
| 360 |
+
report = check_taps(version, taps, s.gs, results=results)
|
| 361 |
+
_session_meta(results, s)
|
| 362 |
+
assert row["finite"], "non-finite block output"
|
| 363 |
+
assert rows_ok, "the compute cores of one tile row hold different residual rows"
|
| 364 |
+
assert torch.equal(got, got2), "two launches of the same block differ"
|
| 365 |
+
assert row["pcc"] >= GATE_ORACLE, f"{kind} block {blk} k={step}: PCC vs S1 {row['pcc']:.6f} < {GATE_ORACLE}"
|
| 366 |
+
assert report.ok, report.summary(failures_only=True)
|
| 367 |
+
|
| 368 |
+
|
| 369 |
+
@pytest.mark.parametrize("step", STEPS_TESTED)
|
| 370 |
+
def test_blocks_chain_vs_golden(results: Any, version: str, k4: K4Session, step: int) -> None:
|
| 371 |
+
"""Blocks [0, 2) in one launch from ``sa_embs[k]``: the pending-residual hand-over between blocks (P12 -> P1) vs
|
| 372 |
+
the golden ``dit_block1_out[k]`` and the chained S1 oracle."""
|
| 373 |
+
s = k4
|
| 374 |
+
x = s.golden_rows("sa_embs", step)
|
| 375 |
+
x_rows = s.mk.make_x_rows(x)
|
| 376 |
+
try:
|
| 377 |
+
t0 = time.perf_counter()
|
| 378 |
+
got = s.mk.run_block(x_rows, 0, matmul_only=False, n_blk=2, step=step)
|
| 379 |
+
dt = time.perf_counter() - t0
|
| 380 |
+
finally:
|
| 381 |
+
s.mk._ttnn.deallocate(x_rows)
|
| 382 |
+
ref = s.oracle_blocks(x, [0, 1], step)
|
| 383 |
+
m = s.dl.m_logical
|
| 384 |
+
row = _row(
|
| 385 |
+
results,
|
| 386 |
+
f"chain-blk0-1-k{step}-vs-oracle",
|
| 387 |
+
got,
|
| 388 |
+
ref,
|
| 389 |
+
GATE_ORACLE,
|
| 390 |
+
step=step,
|
| 391 |
+
pcc_valid_rows=pcc_mod.pcc(got[:m], ref[:m]),
|
| 392 |
+
wall_s=dt,
|
| 393 |
+
)
|
| 394 |
+
taps = {s.gs.threshold_key("dit_block1_out", step=step): got[:m].reshape(1, m, -1)}
|
| 395 |
+
report = check_taps(version, taps, s.gs, results=results)
|
| 396 |
+
_session_meta(results, s)
|
| 397 |
+
assert row["finite"] and row["pcc"] >= GATE_ORACLE, f"chain k={step}: PCC vs S1 {row['pcc']:.6f}"
|
| 398 |
+
assert report.ok, report.summary(failures_only=True)
|
code/models/experimental/gr00t/tests/tt/test_mk_dit_step.py
ADDED
|
@@ -0,0 +1,335 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
|
| 2 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 3 |
+
|
| 4 |
+
"""WP-K4 device tests of the megakernel's ``step`` and ``steps4`` variants (IMPLEMENTATION_PLAN.md §5.4 "Per step",
|
| 5 |
+
§5.7 K4 acceptance rows ``dit_block_last_out`` / ``dit_out`` / ``pred_velocity`` / ``action_pred_*``, gate G2).
|
| 6 |
+
|
| 7 |
+
* ``test_step_vs_golden[k]`` -- all ``n_blocks`` blocks + output AdaLN + ``proj_out_2`` of denoising step ``k`` from
|
| 8 |
+
the golden ``sa_embs[k]``: ``dit_block_last_out[k]`` / ``dit_out[k]`` against the golden gates and against the
|
| 9 |
+
Stage-1 ``DiTStepTT`` on device (``dit_block_last_out`` PCC >= 0.9999; ``dit_out`` >= 0.9998 **and** not further from
|
| 10 |
+
the fp32 golden than Stage 1 is, see ``GATE_ORACLE_HEAD`` / ``GOLDEN_SLACK``); a second launch bit-identical;
|
| 11 |
+
* ``test_steps4_vs_golden`` -- the whole denoise in one launch (action encoder -> sa_embs -> blocks -> head ->
|
| 12 |
+
decoder -> Euler, 4 steps) from the golden ``initial_noise`` / ``state_features``: ``pred_velocity[k]``,
|
| 13 |
+
``action_decoder_out[k]``, ``action_pred_normalized``, ``action_pred_valid`` gates + PCC vs the Stage-1
|
| 14 |
+
``ActionHeadTT.denoise`` on device;
|
| 15 |
+
* ``test_steps4_determinism_and_trace`` -- 10 untraced launches bit-identical, the traced replay bit-identical to the
|
| 16 |
+
untraced result (``action_pred``, and ``velocity`` when the taps are allocated), and the G2 timing: 20 traced replays
|
| 17 |
+
(min / median ms), per-step ms vs ``1.3 x bytes / GB/s``. G2 is **recorded, not asserted**: the ``G2-step-time`` row
|
| 18 |
+
carries ``ok = pass`` (a miss shows as ``ok: False`` in the results JSON while the test stays green -- plan section
|
| 19 |
+
5.7 fallback, ratified 2026-09-18 in docs/plan/reviews/WP-K4-review-round1.md).
|
| 20 |
+
|
| 21 |
+
Run (one pytest process per file, through the device lock)::
|
| 22 |
+
|
| 23 |
+
DEVICE_LOCK_TIMEOUT=14400 bin/with-device.sh python -m pytest -q -rA --timeout=1800 -p no:cacheprovider \\
|
| 24 |
+
models/experimental/gr00t/tests/tt/test_mk_dit_step.py --gr00t-version n16
|
| 25 |
+
"""
|
| 26 |
+
|
| 27 |
+
from __future__ import annotations
|
| 28 |
+
|
| 29 |
+
import statistics
|
| 30 |
+
import time
|
| 31 |
+
from typing import Any, Dict, Iterator, List, Tuple
|
| 32 |
+
|
| 33 |
+
import pytest
|
| 34 |
+
import torch
|
| 35 |
+
|
| 36 |
+
from models.experimental.gr00t.common import pcc as pcc_mod
|
| 37 |
+
from models.experimental.gr00t.common.checkpoint import LazyCheckpoint
|
| 38 |
+
from models.experimental.gr00t.common.configs import DIT_QUERY_ROWS_PAD
|
| 39 |
+
from models.experimental.gr00t.common.golden import GoldenSet
|
| 40 |
+
from models.experimental.gr00t.tests.tt.harness import check_taps, device_facts
|
| 41 |
+
from models.experimental.gr00t.tests.tt.test_mk_dit_block import K4Session, STREAM_DTYPE, pad_rows
|
| 42 |
+
from models.experimental.gr00t.tt import encoders as E
|
| 43 |
+
from models.experimental.gr00t.tt.megakernel import dit_program as P
|
| 44 |
+
|
| 45 |
+
torch.set_grad_enabled(False)
|
| 46 |
+
|
| 47 |
+
GATE_ORACLE = 0.9999 # plan §5.7 K4: PCC vs the S1 device oracle (blocks, residual stream, action taps)
|
| 48 |
+
#: the output head after 32 blocks: the two bf16 device pipelines (different op / rounding order per block) drift to
|
| 49 |
+
#: 1 - PCC ~ 1e-4 on dit_out (measured 0.99990-0.99997 on N1.6, mk_k4_summary.md §4); gated at 0.9998 against the
|
| 50 |
+
#: oracle AND "not further from the fp32 golden than Stage 1 is" (GOLDEN_SLACK). Post-hoc change from the plan's
|
| 51 |
+
#: 0.9999: N1.6 k=3 measured 0.9998980 / 0.9998996 under 0.9999 (test_step_vs_golden_3_n16_20260917-233300 / -233501,
|
| 52 |
+
#: failed); ratified by the plan owner 2026-09-18 (WP-K4-review-round1, "Plan-owner decisions" 1; summary §9). The
|
| 53 |
+
#: golden gates (gates_multi.json) are untouched.
|
| 54 |
+
GATE_ORACLE_HEAD = 0.9998
|
| 55 |
+
GOLDEN_SLACK = 2e-5
|
| 56 |
+
STEPS: Tuple[int, ...] = (0, 1, 2, 3)
|
| 57 |
+
N_DETERMINISM = 10 # plan D12: megakernel 10x bit-identical
|
| 58 |
+
N_REPLAYS = 20 # G2 timing (work_packages.json WP-K4)
|
| 59 |
+
#: measured direct-mode arena rates (mk_k1_summary.md §3: 464 GB/s bf16, 414 GB/s bfp8) and the G2 factor (plan §5.7)
|
| 60 |
+
MEASURED_GBPS: Dict[str, float] = {"bf16": 464.0, "bfp8_b": 414.0}
|
| 61 |
+
G2_FACTOR = 1.3
|
| 62 |
+
|
| 63 |
+
_SESSIONS: Dict[str, K4Session] = {}
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
@pytest.fixture(scope="module")
|
| 67 |
+
def k4_device() -> Iterator[Any]:
|
| 68 |
+
"""The p150a opened with the smaller ``worker_l1_size`` the K4 binaries need (see test_mk_dit_block)."""
|
| 69 |
+
import os
|
| 70 |
+
|
| 71 |
+
from models.experimental.gr00t.tests.tt.conftest import DEFAULT_TT_METAL_CACHE
|
| 72 |
+
from models.experimental.gr00t.tt.device import close_gr00t_device
|
| 73 |
+
|
| 74 |
+
os.environ.setdefault("TT_METAL_CACHE", str(DEFAULT_TT_METAL_CACHE))
|
| 75 |
+
device = P.open_k4_device()
|
| 76 |
+
try:
|
| 77 |
+
yield device
|
| 78 |
+
finally:
|
| 79 |
+
for v, s in list(_SESSIONS.items()):
|
| 80 |
+
s.release()
|
| 81 |
+
del _SESSIONS[v]
|
| 82 |
+
close_gr00t_device(device)
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
@pytest.fixture(scope="module")
|
| 86 |
+
def k4(version: str, golden: Tuple[GoldenSet, Any, Any], k4_device: Any) -> Iterator[K4Session]:
|
| 87 |
+
for v, s in list(_SESSIONS.items()):
|
| 88 |
+
if v != version:
|
| 89 |
+
s.release()
|
| 90 |
+
del _SESSIONS[v]
|
| 91 |
+
if version not in _SESSIONS:
|
| 92 |
+
_SESSIONS[version] = K4Session(version, golden, k4_device, variant="step")
|
| 93 |
+
yield _SESSIONS[version]
|
| 94 |
+
|
| 95 |
+
|
| 96 |
+
def _row(results: Any, key: str, got: torch.Tensor, ref: torch.Tensor, gate: float, **extra: Any) -> Dict[str, Any]:
|
| 97 |
+
m = pcc_mod.metrics(got.float(), ref.float())
|
| 98 |
+
row = {
|
| 99 |
+
"tap": "mk_k4_step",
|
| 100 |
+
"key": key,
|
| 101 |
+
"rule": "pcc",
|
| 102 |
+
"pcc": m["pcc"],
|
| 103 |
+
"gate_pcc": gate,
|
| 104 |
+
"max_abs": m["max_abs"],
|
| 105 |
+
"rel_l2": m["rel_l2"],
|
| 106 |
+
"finite": m["finite"],
|
| 107 |
+
"ok": bool(m["finite"] and m["pcc"] >= gate),
|
| 108 |
+
}
|
| 109 |
+
row.update(extra)
|
| 110 |
+
results.add_row(row)
|
| 111 |
+
return row
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
def _meta(results: Any, s: K4Session) -> None:
|
| 115 |
+
results.set_device_facts(device_facts(s.device))
|
| 116 |
+
results.set_meta(
|
| 117 |
+
layout=s.dl.layout_name,
|
| 118 |
+
embodiment_id=s.dl.embodiment_id,
|
| 119 |
+
stream_dtype=STREAM_DTYPE,
|
| 120 |
+
n_blocks=s.arena.n_blocks,
|
| 121 |
+
key_sets=sorted(set(s.arena.key_sets)),
|
| 122 |
+
worker_l1_size=P.K4_WORKER_L1_SIZE,
|
| 123 |
+
stats=s.mk.stats().__dict__ if s.mk.last_spec is not None else None,
|
| 124 |
+
compile_s={str(k): v for k, v in s.mk.compile_seconds.items()},
|
| 125 |
+
)
|
| 126 |
+
|
| 127 |
+
|
| 128 |
+
def _b3_bias(s: K4Session) -> torch.Tensor:
|
| 129 |
+
"""``action_encoder.W3.b[e]`` (folded into ``b3_pos`` on device; the golden ``action_encoder_out`` includes it)."""
|
| 130 |
+
ck = LazyCheckpoint(s.cfg.version)
|
| 131 |
+
try:
|
| 132 |
+
return ck.slot(
|
| 133 |
+
s.cfg.checkpoint.action_head_prefix + "action_encoder.W3.b", int(s.dl.embodiment_id), torch.float32
|
| 134 |
+
)
|
| 135 |
+
finally:
|
| 136 |
+
ck.close()
|
| 137 |
+
|
| 138 |
+
|
| 139 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 140 |
+
# step variant
|
| 141 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 142 |
+
@pytest.mark.parametrize("step", STEPS)
|
| 143 |
+
def test_step_vs_golden(results: Any, version: str, k4: K4Session, step: int) -> None:
|
| 144 |
+
"""One denoising step (all blocks + head) from the golden ``sa_embs[k]`` vs the golden taps and the S1 step."""
|
| 145 |
+
import ttnn
|
| 146 |
+
from models.experimental.gr00t.tt import layers as TL
|
| 147 |
+
from models.experimental.gr00t.tt.dit_ops import DiTStepTT
|
| 148 |
+
|
| 149 |
+
s = k4
|
| 150 |
+
m = s.dl.m_logical
|
| 151 |
+
x = s.golden_rows("sa_embs", step)
|
| 152 |
+
# S1 oracle first (blocks + head, every block output kept as a tap), with the megakernel's L1 residents released:
|
| 153 |
+
# under the K4 worker_l1_size the oracle's layer_norm CBs do not fit next to a_rep on 32-block versions
|
| 154 |
+
s.mk.release_l1_residents()
|
| 155 |
+
oracle = DiTStepTT(s.cfg, s.W, s.dl, s.masks, s.kv, collect_taps=True)
|
| 156 |
+
xd = TL.to_device(x.reshape(1, 1, DIT_QUERY_ROWS_PAD, -1).to(torch.bfloat16), s.device, mem="L1")
|
| 157 |
+
y_dev = oracle.forward(xd, step)
|
| 158 |
+
ref_y = TL.to_torch(y_dev).float().reshape(DIT_QUERY_ROWS_PAD, -1)
|
| 159 |
+
ref_x = TL.to_torch(oracle.debug_taps["dit_block_last_out"]).float().reshape(DIT_QUERY_ROWS_PAD, -1)
|
| 160 |
+
for t in list(oracle.debug_taps.values()) + [xd]:
|
| 161 |
+
ttnn.deallocate(t)
|
| 162 |
+
s.mk.set_variant("step")
|
| 163 |
+
t0 = time.perf_counter()
|
| 164 |
+
y, x_last = s.mk.run_step(x, step)
|
| 165 |
+
dt = time.perf_counter() - t0
|
| 166 |
+
y2, x_last2 = s.mk.run_step(x, step)
|
| 167 |
+
row_x = _row(results, f"dit_block_last_out-k{step}-vs-oracle", x_last, ref_x, GATE_ORACLE, step=step, wall_s=dt)
|
| 168 |
+
row_y = _row(results, f"dit_out-k{step}-vs-oracle", y, ref_y, GATE_ORACLE_HEAD, step=step)
|
| 169 |
+
det = bool(torch.equal(y, y2) and torch.equal(x_last, x_last2))
|
| 170 |
+
results.add_row({"tap": "mk_k4_step", "key": f"step-k{step}-deterministic", "rule": "exact", "ok": det})
|
| 171 |
+
taps = {
|
| 172 |
+
s.gs.threshold_key("dit_block_last_out", step=step): x_last[:m].reshape(1, m, -1),
|
| 173 |
+
s.gs.threshold_key("dit_out", step=step): y[:m].reshape(1, m, -1),
|
| 174 |
+
}
|
| 175 |
+
report = check_taps(version, taps, s.gs, results=results)
|
| 176 |
+
# the megakernel must be at least as close to the fp32 golden as the Stage-1 path is (both bf16 pipelines)
|
| 177 |
+
gold_x = s.golden_rows("dit_block_last_out", step)
|
| 178 |
+
gold_y = pad_rows(s.gs.load("dit_out", step=step))
|
| 179 |
+
closeness = {
|
| 180 |
+
"dit_block_last_out": (pcc_mod.pcc(x_last[:m], gold_x[:m]), pcc_mod.pcc(ref_x[:m], gold_x[:m])),
|
| 181 |
+
"dit_out": (pcc_mod.pcc(y[:m], gold_y[:m]), pcc_mod.pcc(ref_y[:m], gold_y[:m])),
|
| 182 |
+
}
|
| 183 |
+
for tap, (mk_pcc, s1_pcc) in closeness.items():
|
| 184 |
+
results.add_row(
|
| 185 |
+
{
|
| 186 |
+
"tap": "mk_k4_step",
|
| 187 |
+
"key": f"{tap}-k{step}-mk-vs-s1-closeness-to-golden",
|
| 188 |
+
"rule": "pcc",
|
| 189 |
+
"pcc": mk_pcc,
|
| 190 |
+
"s1_pcc_vs_golden": s1_pcc,
|
| 191 |
+
"gate_pcc": s1_pcc - GOLDEN_SLACK,
|
| 192 |
+
"ok": bool(mk_pcc >= s1_pcc - GOLDEN_SLACK),
|
| 193 |
+
}
|
| 194 |
+
)
|
| 195 |
+
_meta(results, s)
|
| 196 |
+
assert row_x["finite"] and row_y["finite"], "non-finite step output"
|
| 197 |
+
assert det, "two launches of the same step differ"
|
| 198 |
+
assert report.ok, report.summary(failures_only=True)
|
| 199 |
+
assert row_x["pcc"] >= GATE_ORACLE, f"dit_block_last_out k={step}: PCC vs S1 {row_x['pcc']:.6f} < {GATE_ORACLE}"
|
| 200 |
+
assert row_y["pcc"] >= GATE_ORACLE_HEAD, f"dit_out k={step}: PCC vs S1 {row_y['pcc']:.6f} < {GATE_ORACLE_HEAD}"
|
| 201 |
+
for tap, (mk_pcc, s1_pcc) in closeness.items():
|
| 202 |
+
assert mk_pcc >= s1_pcc - GOLDEN_SLACK, f"{tap} k={step}: mk vs golden {mk_pcc:.6f} < S1 vs golden {s1_pcc:.6f}"
|
| 203 |
+
|
| 204 |
+
|
| 205 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 206 |
+
# steps4 variant
|
| 207 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 208 |
+
def _steps4_inputs(s: K4Session) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 209 |
+
noise = s.gs.load("initial_noise")
|
| 210 |
+
state = s.gs.load("state_features").reshape(-1)
|
| 211 |
+
return noise, state
|
| 212 |
+
|
| 213 |
+
|
| 214 |
+
def _steps4_taps(s: K4Session, out: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]:
|
| 215 |
+
"""Megakernel outputs -> golden-keyed host tensors (``taps_to_golden`` shapes)."""
|
| 216 |
+
lo, hi = E.action_rows(s.cfg, s.dl)
|
| 217 |
+
m = s.dl.m_logical
|
| 218 |
+
a_max = int(s.cfg.encoders.max_action_dim)
|
| 219 |
+
pred = E.xt_to_action_pred(out["action_pred"], s.cfg, s.dl)
|
| 220 |
+
taps = {"action_pred_normalized": pred, "action_pred_valid": pred}
|
| 221 |
+
v = out["velocity"] # [4, 64, A_pad]
|
| 222 |
+
for k in STEPS:
|
| 223 |
+
taps[s.gs.threshold_key("pred_velocity", step=k)] = v[k, lo:hi, :a_max].reshape(1, hi - lo, a_max)
|
| 224 |
+
taps[s.gs.threshold_key("action_decoder_out", step=k)] = v[k, :m, :a_max].reshape(1, m, a_max)
|
| 225 |
+
return taps
|
| 226 |
+
|
| 227 |
+
|
| 228 |
+
def test_steps4_vs_golden(results: Any, version: str, k4: K4Session) -> None:
|
| 229 |
+
"""The 4-step denoise in one launch vs the golden action taps and the S1 ``ActionHeadTT`` on device."""
|
| 230 |
+
import ttnn
|
| 231 |
+
from models.experimental.gr00t.tt import action_head as AH
|
| 232 |
+
from models.experimental.gr00t.tt import layers as TL
|
| 233 |
+
from models.experimental.gr00t.tt.policy import TTPolicy
|
| 234 |
+
|
| 235 |
+
s = k4
|
| 236 |
+
noise, state = _steps4_inputs(s)
|
| 237 |
+
# S1 oracle first, with the megakernel's L1 residents released (a_rep 256 KB on all cores, the steps4 rows): the
|
| 238 |
+
# same inputs through ActionHeadTT (ttnn ops, same weights / K/V / masks, every intermediate kept as a tap). Under
|
| 239 |
+
# the K4 worker_l1_size its layer_norm CBs (~980 KB on the op's 2 cores) do not fit next to a_rep.
|
| 240 |
+
s.mk.release_l1_residents()
|
| 241 |
+
head = AH.ActionHeadTT(
|
| 242 |
+
s.cfg,
|
| 243 |
+
s.dl,
|
| 244 |
+
s.W,
|
| 245 |
+
TTPolicy(dtype_policy=s.policy.dtype_policy),
|
| 246 |
+
s.device,
|
| 247 |
+
masks=s.masks,
|
| 248 |
+
kv=s.kv,
|
| 249 |
+
collect_taps=True,
|
| 250 |
+
)
|
| 251 |
+
try:
|
| 252 |
+
head.upload_noise(noise)
|
| 253 |
+
state_dev = TL.to_device(state.reshape(1, 1, 1, -1).to(torch.bfloat16), s.device, mem="DRAM")
|
| 254 |
+
head.denoise(head.x_t, state_dev)
|
| 255 |
+
ref = head.golden_taps(b3=_b3_bias(s), block_taps=False)
|
| 256 |
+
ttnn.deallocate(state_dev)
|
| 257 |
+
head.free_taps()
|
| 258 |
+
finally:
|
| 259 |
+
head.release()
|
| 260 |
+
s.mk.set_variant("steps4")
|
| 261 |
+
t0 = time.perf_counter()
|
| 262 |
+
out = s.mk.run_steps4(noise, state)
|
| 263 |
+
dt = time.perf_counter() - t0
|
| 264 |
+
taps = _steps4_taps(s, out)
|
| 265 |
+
rows = {}
|
| 266 |
+
for key in ("action_pred_normalized",) + tuple(s.gs.threshold_key("pred_velocity", step=k) for k in STEPS):
|
| 267 |
+
rows[key] = _row(results, f"{key}-vs-oracle", taps[key], ref[key], GATE_ORACLE, wall_s=dt)
|
| 268 |
+
report = check_taps(version, taps, s.gs, results=results)
|
| 269 |
+
_meta(results, s)
|
| 270 |
+
results.set_meta(steps4_wall_s=dt, oracle_min_pcc=min(r["pcc"] for r in rows.values()))
|
| 271 |
+
assert all(r["finite"] for r in rows.values()), "non-finite steps4 output"
|
| 272 |
+
assert report.ok, report.summary(failures_only=True)
|
| 273 |
+
bad = {k: r["pcc"] for k, r in rows.items() if r["pcc"] < GATE_ORACLE}
|
| 274 |
+
assert not bad, f"steps4 vs S1 below {GATE_ORACLE}: {bad}"
|
| 275 |
+
|
| 276 |
+
|
| 277 |
+
def test_steps4_determinism_and_trace(results: Any, version: str, k4: K4Session) -> None:
|
| 278 |
+
"""10 untraced launches bit-identical; the trace replay bit-identical to the untraced result; G2 timing."""
|
| 279 |
+
s = k4
|
| 280 |
+
s.mk.set_variant("steps4")
|
| 281 |
+
noise, state = _steps4_inputs(s)
|
| 282 |
+
outs: List[Dict[str, torch.Tensor]] = [s.mk.run_steps4(noise, state) for _ in range(N_DETERMINISM)]
|
| 283 |
+
identical = all(
|
| 284 |
+
torch.equal(o["action_pred"], outs[0]["action_pred"]) and torch.equal(o["velocity"], outs[0]["velocity"])
|
| 285 |
+
for o in outs[1:]
|
| 286 |
+
)
|
| 287 |
+
results.add_row({"tap": "mk_k4_step", "key": "steps4-10x-bit-identical", "rule": "exact", "ok": identical})
|
| 288 |
+
# trace: capture once, replay N_REPLAYS times (the inputs are the persistent tensors, refreshed before capture)
|
| 289 |
+
s.mk.set_steps4_inputs(noise, state)
|
| 290 |
+
tid = s.mk.capture_steps4()
|
| 291 |
+
try:
|
| 292 |
+
times = [s.mk.replay(tid) for _ in range(N_REPLAYS)]
|
| 293 |
+
traced = s.mk.read_steps4_outputs()
|
| 294 |
+
finally:
|
| 295 |
+
s.mk.release_trace(tid)
|
| 296 |
+
traced_ok = bool(torch.equal(traced["action_pred"], outs[0]["action_pred"]))
|
| 297 |
+
cmp_velocity = "velocity" in traced and "velocity" in outs[0]
|
| 298 |
+
if cmp_velocity:
|
| 299 |
+
traced_ok = traced_ok and bool(torch.equal(traced["velocity"], outs[0]["velocity"]))
|
| 300 |
+
results.add_row(
|
| 301 |
+
{
|
| 302 |
+
"tap": "mk_k4_step",
|
| 303 |
+
"key": "steps4-traced-eq-untraced",
|
| 304 |
+
"rule": "exact",
|
| 305 |
+
"ok": traced_ok,
|
| 306 |
+
"velocity_compared": cmp_velocity,
|
| 307 |
+
}
|
| 308 |
+
)
|
| 309 |
+
ms = sorted(t * 1e3 for t in times)
|
| 310 |
+
per_step_ms = statistics.median(ms) / P.STEPS4_N_STEPS
|
| 311 |
+
bytes_step = s.mk.arena_bytes_per_step()
|
| 312 |
+
floor_ms = bytes_step / MEASURED_GBPS[STREAM_DTYPE] / 1e6
|
| 313 |
+
gate_ms = G2_FACTOR * floor_ms
|
| 314 |
+
g2 = {
|
| 315 |
+
"replays": N_REPLAYS,
|
| 316 |
+
"trace_ms_min": ms[0],
|
| 317 |
+
"trace_ms_median": statistics.median(ms),
|
| 318 |
+
"trace_ms_max": ms[-1],
|
| 319 |
+
"per_step_ms_median": per_step_ms,
|
| 320 |
+
"per_step_ms_min": ms[0] / P.STEPS4_N_STEPS,
|
| 321 |
+
"bytes_per_step": bytes_step,
|
| 322 |
+
"measured_gbps": MEASURED_GBPS[STREAM_DTYPE],
|
| 323 |
+
"floor_ms_per_step": floor_ms,
|
| 324 |
+
"gate_ms_per_step": gate_ms,
|
| 325 |
+
"ratio_to_floor": per_step_ms / floor_ms,
|
| 326 |
+
"pass": per_step_ms <= gate_ms,
|
| 327 |
+
}
|
| 328 |
+
# recorded, not asserted (plan §5.7 fallback ratified 2026-09-18): the row's ok follows the gate so a miss is
|
| 329 |
+
# visible as ok: False / n_fail in the JSON while the determinism assertions below decide the test outcome
|
| 330 |
+
results.add_row({"tap": "mk_k4_step", "key": "G2-step-time", "rule": "timing", "ok": bool(g2["pass"]), **g2})
|
| 331 |
+
results.set_timing(steps4_trace_median_s=statistics.median(times), steps4_trace_min_s=min(times))
|
| 332 |
+
_meta(results, s)
|
| 333 |
+
results.set_meta(g2=g2)
|
| 334 |
+
assert identical, "10 steps4 launches are not bit-identical"
|
| 335 |
+
assert traced_ok, "the traced steps4 replay differs from the untraced launch"
|
code/models/experimental/gr00t/tt/action_head.py
CHANGED
|
@@ -24,17 +24,45 @@ and refreshed per call with :meth:`ActionHeadTT.update_masks`) and the hoisted K
|
|
| 24 |
(``HoistOut.kv``). Weights and AdaLN vector tables live in ``TTWeights`` (DRAM).
|
| 25 |
|
| 26 |
``backend`` (``TTPolicy.dit_backend``): ``"ttnn"`` runs the Stage-1 op sequence above; ``"megakernel"`` (Stage 2,
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 34 |
"""
|
| 35 |
|
| 36 |
from __future__ import annotations
|
| 37 |
|
|
|
|
| 38 |
from dataclasses import dataclass
|
| 39 |
from typing import Any, Dict, List, Mapping, Optional, Tuple
|
| 40 |
|
|
@@ -59,6 +87,40 @@ STEP_TAPS: Tuple[str, ...] = (
|
|
| 59 |
"actions_out",
|
| 60 |
)
|
| 61 |
FINAL_TAPS: Tuple[str, ...] = ("action_pred_normalized",)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 62 |
|
| 63 |
|
| 64 |
def _shape(t: Any) -> Tuple[int, ...]:
|
|
@@ -236,15 +298,17 @@ class ActionHeadTT:
|
|
| 236 |
Args:
|
| 237 |
cfg, dlayout: version config and the baked static layout of the call shape.
|
| 238 |
W: ``tt.weights.TTWeights`` (categories ``dit``, ``dit_precompute``, ``encoders`` resident; ``constants``).
|
| 239 |
-
policy: ``TTPolicy`` -- ``dit_backend`` selects the Stage-1 op path or the megakernel (WP-K5)
|
| 240 |
-
|
| 241 |
-
|
|
|
|
| 242 |
masks: the ``tt.layout.materialise`` dictionary shared with the model (default: own buffers from
|
| 243 |
:func:`allocate_mask_buffers`, refreshed with :meth:`update_masks`).
|
| 244 |
kv: block index -> persistent ``(K_i, V_i)`` (``VLAdapterTT`` ``HoistOut.kv``); or :meth:`bind_kv` later.
|
| 245 |
-
weights: name -> tensor overrides passed to the DiT modules (negative controls, tests).
|
| 246 |
persist_sa_embs: keep a persistent DRAM ``sa_embs`` buffer updated every step (one extra copy per step).
|
| 247 |
collect_taps: keep the golden-named device taps in ``taps`` (never inside a trace capture).
|
|
|
|
| 248 |
"""
|
| 249 |
|
| 250 |
def __init__(
|
|
@@ -267,6 +331,7 @@ class ActionHeadTT:
|
|
| 267 |
fuse_swish: bool = False,
|
| 268 |
persist_sa_embs: bool = False,
|
| 269 |
collect_taps: bool = False,
|
|
|
|
| 270 |
) -> None:
|
| 271 |
if not isinstance(cfg, GR00TConfig) or not isinstance(dlayout, DeviceLayout):
|
| 272 |
raise TypeError("ActionHeadTT(cfg: GR00TConfig, dlayout: DeviceLayout, W, policy: TTPolicy, device)")
|
|
@@ -290,6 +355,7 @@ class ActionHeadTT:
|
|
| 290 |
self.width = self.plan.width
|
| 291 |
self.rows = E.action_rows(cfg, dlayout)
|
| 292 |
self.buckets = tuple(int(b) for b in cfg.sampler.timestep_buckets)
|
|
|
|
| 293 |
|
| 294 |
# masks: shared dictionary or own buffers (materialise-compatible, refreshed in place)
|
| 295 |
if masks is None:
|
|
@@ -309,18 +375,34 @@ class ActionHeadTT:
|
|
| 309 |
self.owns_masks = False
|
| 310 |
self.dit_masks = dit_masks(self.masks)
|
| 311 |
|
| 312 |
-
# modules
|
| 313 |
-
self.encoders =
|
| 314 |
-
|
| 315 |
-
)
|
| 316 |
-
self._check_encoders(self.encoders)
|
| 317 |
-
self.step = DiTStepTT(
|
| 318 |
-
cfg, W, dlayout, self.masks, kv, weights=weights, mem=mem, ckc=ckc, head_ckc=head_ckc, sdpa_ckc=sdpa_ckc
|
| 319 |
-
)
|
| 320 |
-
self.step.set_collect_taps(self.collect_taps)
|
| 321 |
-
self.encoders.set_collect_taps(self.collect_taps)
|
| 322 |
self._tap_pool: List[Any] = [] # encoder intermediates kept while taps are on (freed by free_taps)
|
| 323 |
-
self.megakernel: Any = None #
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 324 |
|
| 325 |
# persistent buffers
|
| 326 |
zeros_x = torch.zeros(1, 1, DIT_QUERY_ROWS_PAD, self.a_pad, dtype=E.TORCH_DTYPES[self.xt_dtype])
|
|
@@ -331,6 +413,8 @@ class ActionHeadTT:
|
|
| 331 |
zeros_sa = torch.zeros(1, 1, DIT_QUERY_ROWS_PAD, self.width, dtype=torch.bfloat16)
|
| 332 |
self.sa_embs = L.to_device(zeros_sa, device, dtype="bf16", mem=buffer_mem)
|
| 333 |
self.n_calls = 0
|
|
|
|
|
|
|
| 334 |
|
| 335 |
# ------------------------------------------------------------------ checks
|
| 336 |
def _check_encoders(self, enc: HeadEncoders) -> None:
|
|
@@ -351,14 +435,183 @@ class ActionHeadTT:
|
|
| 351 |
f"action encoder has {enc.action_encoder.n_steps} step biases, sampler has {self.n_steps} steps"
|
| 352 |
)
|
| 353 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 354 |
# ------------------------------------------------------------------ per-call state
|
| 355 |
@property
|
| 356 |
def kv(self) -> Optional[Dict[int, Tuple[Any, Any]]]:
|
| 357 |
-
return self.step.kv
|
| 358 |
|
| 359 |
def bind_kv(self, kv: Mapping[int, Any]) -> Dict[int, Tuple[Any, Any]]:
|
| 360 |
-
"""Bind the persistent hoisted K/V (``VLAdapterTT`` ``HoistOut.kv``) once; the objects are what trace B captures.
|
| 361 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 362 |
|
| 363 |
def update_masks(self, slots: SlotLayout) -> Dict[str, Any]:
|
| 364 |
"""Write this call's per-call masks (``cross_text`` / ``cross_all`` / ``llm`` / ``vl``) into the persistent
|
|
@@ -378,8 +631,10 @@ class ActionHeadTT:
|
|
| 378 |
def set_collect_taps(self, flag: bool) -> None:
|
| 379 |
"""Toggle tap collection on the head and its DiT step (off before a trace capture)."""
|
| 380 |
self.collect_taps = bool(flag)
|
| 381 |
-
self.step
|
| 382 |
-
|
|
|
|
|
|
|
| 383 |
if not flag:
|
| 384 |
self.taps = {}
|
| 385 |
self._tap_pool = []
|
|
@@ -403,16 +658,7 @@ class ActionHeadTT:
|
|
| 403 |
for m in modules:
|
| 404 |
self._tap_pool.extend(m.debug_taps.values())
|
| 405 |
|
| 406 |
-
def
|
| 407 |
-
"""Run the 4 unrolled steps from ``x_t`` (``[1,1,64,A_pad]`` in the version dtype; the persistent buffer or a
|
| 408 |
-
tensor that is copied into it) with ``state_features [1,1,1,1536]`` bf16 (``VLAdapterTT`` ``HoistOut``) and
|
| 409 |
-
the hoisted ``kv`` (default: the bound handles) -> the persistent ``action_pred`` buffer."""
|
| 410 |
-
import ttnn
|
| 411 |
-
|
| 412 |
-
if self.backend != "ttnn":
|
| 413 |
-
raise NotImplementedError(
|
| 414 |
-
f"ActionHeadTT backend {self.backend!r}: the megakernel denoise (DiTMegakernel.run) is wired by WP-K5"
|
| 415 |
-
)
|
| 416 |
want = (1, 1, DIT_QUERY_ROWS_PAD, self.a_pad)
|
| 417 |
if _shape(x_t) != want:
|
| 418 |
raise ValueError(f"x_t must be {want}, got {_shape(x_t)}")
|
|
@@ -420,15 +666,31 @@ class ActionHeadTT:
|
|
| 420 |
raise TypeError(f"x_t must be {self.xt_dtype} (cfg.sampler.euler_dtype), got {L.dtype_name(x_t.dtype)}")
|
| 421 |
if _shape(state_features) != (1, 1, 1, self.width):
|
| 422 |
raise ValueError(f"state_features must be (1, 1, 1, {self.width}), got {_shape(state_features)}")
|
| 423 |
-
|
| 424 |
-
|
| 425 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 426 |
raise ValueError("no hoisted K/V bound (VLAdapterTT HoistOut.kv): pass kv= or call bind_kv first")
|
| 427 |
if x_t is not self.x_t:
|
| 428 |
ttnn.copy(x_t, self.x_t)
|
| 429 |
self.taps = {}
|
| 430 |
self._tap_pool = []
|
| 431 |
-
enc = self.encoders
|
| 432 |
x = self.x_t
|
| 433 |
for k in range(self.n_steps):
|
| 434 |
a = enc.action_encoder.forward(x, k)
|
|
@@ -446,8 +708,8 @@ class ActionHeadTT:
|
|
| 446 |
ttnn.deallocate(sa)
|
| 447 |
sa = self.sa_embs
|
| 448 |
self._tap(f"sa_embs[k={k}]", sa)
|
| 449 |
-
y =
|
| 450 |
-
for name, t in
|
| 451 |
self._tap(f"{name}[k={k}]", t)
|
| 452 |
if sa is not self.sa_embs:
|
| 453 |
self._free(sa)
|
|
@@ -468,6 +730,47 @@ class ActionHeadTT:
|
|
| 468 |
self.n_calls += 1
|
| 469 |
return self.action_pred
|
| 470 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 471 |
__call__ = denoise
|
| 472 |
|
| 473 |
# ------------------------------------------------------------------ host side
|
|
@@ -479,26 +782,44 @@ class ActionHeadTT:
|
|
| 479 |
""":func:`taps_to_golden` of the last ``denoise`` call's taps."""
|
| 480 |
return taps_to_golden(self.taps, self.cfg, self.dlayout, b3=b3, block_taps=block_taps)
|
| 481 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 482 |
def free_taps(self) -> None:
|
| 483 |
-
"""Deallocate every collected tap that is not a persistent buffer (each object once)."""
|
| 484 |
import ttnn
|
| 485 |
|
| 486 |
keep = {id(self.x_t), id(self.action_pred), id(self.sa_embs)}
|
| 487 |
seen: set = set()
|
| 488 |
-
|
| 489 |
-
|
|
|
|
| 490 |
continue
|
| 491 |
seen.add(id(t))
|
| 492 |
ttnn.deallocate(t)
|
| 493 |
self.taps = {}
|
| 494 |
self._tap_pool = []
|
| 495 |
-
self.step
|
| 496 |
-
|
|
|
|
|
|
|
| 497 |
|
| 498 |
def release(self) -> None:
|
| 499 |
-
"""Deallocate the persistent buffers this head owns (masks only when it allocated them; never the K/V)
|
|
|
|
| 500 |
import ttnn
|
| 501 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 502 |
for t in (self.x_t, self.action_pred, self.sa_embs):
|
| 503 |
if t is not None:
|
| 504 |
ttnn.deallocate(t)
|
|
@@ -509,7 +830,7 @@ class ActionHeadTT:
|
|
| 509 |
self.masks = {}
|
| 510 |
|
| 511 |
def describe(self) -> Dict[str, Any]:
|
| 512 |
-
d = self.step.describe()
|
| 513 |
d.update(
|
| 514 |
{
|
| 515 |
"backend": self.backend,
|
|
@@ -518,19 +839,39 @@ class ActionHeadTT:
|
|
| 518 |
"buckets": list(self.buckets),
|
| 519 |
"persist_sa_embs": self.sa_embs is not None,
|
| 520 |
"owns_masks": self.owns_masks,
|
| 521 |
-
"fuse_swish": self.encoders.action_encoder.fuse_swish,
|
| 522 |
"buffer_mem": getattr(self.buffer_mem, "kind", str(self.buffer_mem)),
|
|
|
|
| 523 |
}
|
| 524 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 525 |
return d
|
| 526 |
|
| 527 |
|
| 528 |
__all__ = [
|
|
|
|
| 529 |
"FINAL_TAPS",
|
|
|
|
|
|
|
| 530 |
"STEP_TAPS",
|
| 531 |
"ActionHeadTT",
|
| 532 |
"HeadEncoders",
|
| 533 |
"allocate_mask_buffers",
|
|
|
|
| 534 |
"placeholder_masks",
|
| 535 |
"taps_to_golden",
|
| 536 |
]
|
|
|
|
| 24 |
(``HoistOut.kv``). Weights and AdaLN vector tables live in ``TTWeights`` (DRAM).
|
| 25 |
|
| 26 |
``backend`` (``TTPolicy.dit_backend``): ``"ttnn"`` runs the Stage-1 op sequence above; ``"megakernel"`` (Stage 2,
|
| 27 |
+
plan §5, WP-K5) runs the whole denoise as **one** persistent ``ttnn.generic_op`` -- the ``steps4`` variant of
|
| 28 |
+
``tt.megakernel.dit_program.DiTMegakernel`` (action encoder -> ``sa_embs`` -> blocks -> head -> decoder -> Euler,
|
| 29 |
+
four steps in one launch; ``tests/tt/results/mk_k4_summary.md`` §3, §8). Under that backend the constructor packs the
|
| 30 |
+
DiT / precompute / encoder host plan tensors into the DRAM weight arena (``TTPolicy.mk_arena_dtype_for(version)``:
|
| 31 |
+
bf16 or bfp8_b stream), uploads it once, binds the hoisted K/V handles and the mask buffers (the same persistent
|
| 32 |
+
objects trace A / the executor write), allocates the kernel's io tensors once (rank-4 shapes so that ``ttnn.copy`` can
|
| 33 |
+
feed them from the model's buffers) and builds the cached program descriptor. Per call (inside trace B)::
|
| 34 |
+
|
| 35 |
+
ttnn.copy(x_t -> mk x_t fp32 [1,1,64,A_pad]) the version dtype -> the kernel's fp32 Euler state (TILE typecast)
|
| 36 |
+
ttnn.to_layout(state_features, ROW_MAJOR) -> copy bf16 [1,1,1,1536] row the hub reads
|
| 37 |
+
ttnn.generic_op(tensors, program) encoder / 16-32 blocks / head / decoder / Euler x 4
|
| 38 |
+
ttnn.copy(mk action_pred fp32 -> action_pred) back to the version dtype: the same buffer contract as "ttnn"
|
| 39 |
+
|
| 40 |
+
so ``Gr00tTT``'s denoise stage, decode and the executor are unchanged. The kernel's two L1 tensors -- ``a_rep``
|
| 41 |
+
(the multicast target, 256 KB on all 110 cores) and the ``x_rep`` tile row (96 KB per compute core) -- are **not**
|
| 42 |
+
kept resident: next to them the Stage-1 backbone programs (matmuls with ~1 MB of static CBs) do not fit under the K4
|
| 43 |
+
``worker_l1_size``, so :meth:`ActionHeadTT._denoise_megakernel` allocates them (``ttnn.allocate_tensor_on_device``,
|
| 44 |
+
no host write) right before the generic_op, rebuilds the program descriptor on their addresses (the addresses are
|
| 45 |
+
common runtime args, ``descriptors._address_of``) and frees them after -- exactly like any traced ttnn op's
|
| 46 |
+
intermediates, so trace B replays them at the captured addresses. The device must be opened with a ``worker_l1_size``
|
| 47 |
+
reduced by ``tt.model.MK_L1_CUT_BYTES_DEFAULT`` (**64 KiB**: the compute-core K4 binaries -- 108,912 B on N1.5 -- need a
|
| 48 |
+
kernel-config ring buffer larger than the p150a default of 70,656 B, ``mk_k4_summary.md`` §7.2 fact 2;
|
| 49 |
+
``tt.model.open_model_device`` does it, and ``Gr00tTT`` moves the VL adapter's intermediates to DRAM because no cut
|
| 50 |
+
that fits the binaries leaves the N1.5 / N1.7 adapter programs room for their L1 activations; the K4 unit tests'
|
| 51 |
+
own 96 KiB cut, ``dit_program.K4_WORKER_L1_SIZE``, is not what the model uses). The megakernel is the
|
| 52 |
+
``TTPolicy.dit_backend`` **default** since 2026-09-18 (``tt.policy.DEFAULT_DIT_BACKEND``).
|
| 53 |
+
|
| 54 |
+
Taps (``collect_taps=True``): the ttnn backend keeps device tensors under the golden keys (``sa_embs[k=0]``,
|
| 55 |
+
``dit_block{i}_out[k=0]``, ``dit_block_last_out[k=0]``, ``dit_out[k=0]``, ``action_decoder_out[k=0]``,
|
| 56 |
+
``actions_out[k=0]``, ``action_pred_normalized``; plus ``action_encoder_pre_pos[k=0]``); the megakernel exposes the
|
| 57 |
+
per-step decoder output (``action_decoder_out[k]`` from its ``velocity_out`` tap tensor, host copies) and
|
| 58 |
+
``action_pred_normalized`` -- the residual-stream intermediates live inside the kernel (:func:`head_tap_names`).
|
| 59 |
+
:func:`taps_to_golden` slices them to the goldens' shapes for ``tests.tt.harness.check_taps``. ``ttnn`` is imported
|
| 60 |
+
lazily (constructor and forwards).
|
| 61 |
"""
|
| 62 |
|
| 63 |
from __future__ import annotations
|
| 64 |
|
| 65 |
+
import time
|
| 66 |
from dataclasses import dataclass
|
| 67 |
from typing import Any, Dict, List, Mapping, Optional, Tuple
|
| 68 |
|
|
|
|
| 87 |
"actions_out",
|
| 88 |
)
|
| 89 |
FINAL_TAPS: Tuple[str, ...] = ("action_pred_normalized",)
|
| 90 |
+
#: per-step taps the megakernel backend produces (``velocity_out`` = the decoder output of every step; module docstring)
|
| 91 |
+
MK_STEP_TAPS: Tuple[str, ...] = ("action_decoder_out",)
|
| 92 |
+
#: golden-keyed taps derived by :func:`taps_to_golden` from a per-step tap (``<tap>[k=i]`` -> ``<derived>[k=i]``)
|
| 93 |
+
DERIVED_STEP_TAPS: Dict[str, Tuple[str, ...]] = {
|
| 94 |
+
"action_encoder_pre_pos": ("action_encoder_out",),
|
| 95 |
+
"action_decoder_out": ("pred_velocity",),
|
| 96 |
+
}
|
| 97 |
+
#: the kernel's persistent io tensors :class:`ActionHeadTT` re-allocates with rank-4 shapes (module docstring)
|
| 98 |
+
MK_IO_RANK4: Tuple[str, ...] = ("x_t", "state_features", "action_pred")
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
def head_tap_names(backend: str, n_steps: int, n_blocks: int) -> Tuple[str, ...]:
|
| 102 |
+
"""The golden-keyed tap names ``ActionHeadTT.golden_taps`` produces for ``backend`` in step order
|
| 103 |
+
(``tt.model.expected_tap_keys`` appends them; ``GoldenSet.threshold_key`` form).
|
| 104 |
+
|
| 105 |
+
* ``"ttnn"``: ``action_encoder_out[k]``, ``sa_embs[k]``, every ``dit_block{i}_out[k]``, ``dit_block_last_out[k]``,
|
| 106 |
+
``dit_out[k]``, ``action_decoder_out[k]``, ``pred_velocity[k]``, ``actions_out[k]`` per step, then
|
| 107 |
+
``action_pred_normalized`` / ``action_pred_valid``;
|
| 108 |
+
* ``"megakernel"``: ``action_decoder_out[k]`` / ``pred_velocity[k]`` per step, then ``action_pred_normalized`` /
|
| 109 |
+
``action_pred_valid`` (the residual stream never leaves the kernel; ``mk_k4_summary.md`` §8 "Not needed").
|
| 110 |
+
"""
|
| 111 |
+
if backend not in DIT_BACKENDS:
|
| 112 |
+
raise ValueError(f"unknown dit_backend {backend!r}; expected one of {DIT_BACKENDS}")
|
| 113 |
+
keys: List[str] = []
|
| 114 |
+
for k in range(int(n_steps)):
|
| 115 |
+
if backend == "ttnn":
|
| 116 |
+
keys += [f"action_encoder_out[k={k}]", f"sa_embs[k={k}]"]
|
| 117 |
+
keys += [f"dit_block{i}_out[k={k}]" for i in range(int(n_blocks))]
|
| 118 |
+
keys += [f"dit_block_last_out[k={k}]", f"dit_out[k={k}]"]
|
| 119 |
+
keys += [f"action_decoder_out[k={k}]", f"pred_velocity[k={k}]"]
|
| 120 |
+
if backend == "ttnn":
|
| 121 |
+
keys.append(f"actions_out[k={k}]")
|
| 122 |
+
keys += ["action_pred_normalized", "action_pred_valid"]
|
| 123 |
+
return tuple(keys)
|
| 124 |
|
| 125 |
|
| 126 |
def _shape(t: Any) -> Tuple[int, ...]:
|
|
|
|
| 298 |
Args:
|
| 299 |
cfg, dlayout: version config and the baked static layout of the call shape.
|
| 300 |
W: ``tt.weights.TTWeights`` (categories ``dit``, ``dit_precompute``, ``encoders`` resident; ``constants``).
|
| 301 |
+
policy: ``TTPolicy`` -- ``dit_backend`` selects the Stage-1 op path or the megakernel (WP-K5);
|
| 302 |
+
``mk_arena_dtype_for(version)`` the megakernel's weight arena.
|
| 303 |
+
device: the open mesh device (buffers are allocated here; the megakernel needs ``tt.model.open_model_device``).
|
| 304 |
+
encoders: a prebuilt :class:`HeadEncoders` (default: built from ``W``; ttnn backend only).
|
| 305 |
masks: the ``tt.layout.materialise`` dictionary shared with the model (default: own buffers from
|
| 306 |
:func:`allocate_mask_buffers`, refreshed with :meth:`update_masks`).
|
| 307 |
kv: block index -> persistent ``(K_i, V_i)`` (``VLAdapterTT`` ``HoistOut.kv``); or :meth:`bind_kv` later.
|
| 308 |
+
weights: name -> tensor overrides passed to the DiT modules (negative controls, tests; ttnn backend only).
|
| 309 |
persist_sa_embs: keep a persistent DRAM ``sa_embs`` buffer updated every step (one extra copy per step).
|
| 310 |
collect_taps: keep the golden-named device taps in ``taps`` (never inside a trace capture).
|
| 311 |
+
mk_k4_budget: a ``descriptors.K4Budget`` for the megakernel's CB / L1 budget (default: the K4 defaults).
|
| 312 |
"""
|
| 313 |
|
| 314 |
def __init__(
|
|
|
|
| 331 |
fuse_swish: bool = False,
|
| 332 |
persist_sa_embs: bool = False,
|
| 333 |
collect_taps: bool = False,
|
| 334 |
+
mk_k4_budget: Any = None,
|
| 335 |
) -> None:
|
| 336 |
if not isinstance(cfg, GR00TConfig) or not isinstance(dlayout, DeviceLayout):
|
| 337 |
raise TypeError("ActionHeadTT(cfg: GR00TConfig, dlayout: DeviceLayout, W, policy: TTPolicy, device)")
|
|
|
|
| 355 |
self.width = self.plan.width
|
| 356 |
self.rows = E.action_rows(cfg, dlayout)
|
| 357 |
self.buckets = tuple(int(b) for b in cfg.sampler.timestep_buckets)
|
| 358 |
+
self.timing: Dict[str, float] = {}
|
| 359 |
|
| 360 |
# masks: shared dictionary or own buffers (materialise-compatible, refreshed in place)
|
| 361 |
if masks is None:
|
|
|
|
| 375 |
self.owns_masks = False
|
| 376 |
self.dit_masks = dit_masks(self.masks)
|
| 377 |
|
| 378 |
+
# modules (ttnn backend) / the megakernel launcher (megakernel backend)
|
| 379 |
+
self.encoders: Optional[HeadEncoders] = None
|
| 380 |
+
self.step: Optional[DiTStepTT] = None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 381 |
self._tap_pool: List[Any] = [] # encoder intermediates kept while taps are on (freed by free_taps)
|
| 382 |
+
self.megakernel: Any = None # tt.megakernel.dit_program.DiTMegakernel (megakernel backend)
|
| 383 |
+
self.arena: Any = None # tt.megakernel.arena.ArenaLayout (megakernel backend)
|
| 384 |
+
self.mk_arena_dtype: Optional[str] = None
|
| 385 |
+
self._kv: Optional[Dict[int, Tuple[Any, Any]]] = None
|
| 386 |
+
if self.backend == "ttnn":
|
| 387 |
+
if mk_k4_budget is not None:
|
| 388 |
+
raise ValueError("mk_k4_budget is a megakernel-backend argument")
|
| 389 |
+
self.encoders = (
|
| 390 |
+
encoders
|
| 391 |
+
if encoders is not None
|
| 392 |
+
else HeadEncoders.build(cfg, W, dlayout, mem=mem, fuse_swish=fuse_swish)
|
| 393 |
+
)
|
| 394 |
+
self._check_encoders(self.encoders)
|
| 395 |
+
self.step = DiTStepTT(
|
| 396 |
+
cfg, W, dlayout, self.masks, kv, weights=weights, mem=mem, ckc=ckc, head_ckc=head_ckc, sdpa_ckc=sdpa_ckc
|
| 397 |
+
)
|
| 398 |
+
self.step.set_collect_taps(self.collect_taps)
|
| 399 |
+
self.encoders.set_collect_taps(self.collect_taps)
|
| 400 |
+
else:
|
| 401 |
+
if encoders is not None or weights is not None or persist_sa_embs:
|
| 402 |
+
raise ValueError(
|
| 403 |
+
"encoders= / weights= / persist_sa_embs are Stage-1 (ttnn backend) arguments: the megakernel "
|
| 404 |
+
"assembles sa_embs and runs the encoders / decoder / Euler in-kernel"
|
| 405 |
+
)
|
| 406 |
|
| 407 |
# persistent buffers
|
| 408 |
zeros_x = torch.zeros(1, 1, DIT_QUERY_ROWS_PAD, self.a_pad, dtype=E.TORCH_DTYPES[self.xt_dtype])
|
|
|
|
| 413 |
zeros_sa = torch.zeros(1, 1, DIT_QUERY_ROWS_PAD, self.width, dtype=torch.bfloat16)
|
| 414 |
self.sa_embs = L.to_device(zeros_sa, device, dtype="bf16", mem=buffer_mem)
|
| 415 |
self.n_calls = 0
|
| 416 |
+
if self.backend == "megakernel":
|
| 417 |
+
self._build_megakernel(kv, mk_k4_budget)
|
| 418 |
|
| 419 |
# ------------------------------------------------------------------ checks
|
| 420 |
def _check_encoders(self, enc: HeadEncoders) -> None:
|
|
|
|
| 435 |
f"action encoder has {enc.action_encoder.n_steps} step biases, sampler has {self.n_steps} steps"
|
| 436 |
)
|
| 437 |
|
| 438 |
+
# ------------------------------------------------------------------ megakernel build (WP-K5; mk_k4_summary.md §8)
|
| 439 |
+
def _build_megakernel(self, kv: Optional[Mapping[int, Any]], k4_budget: Any) -> None:
|
| 440 |
+
"""Arena (packed from the same bf16 host plan ``TTWeights`` uploads) + upload, K/V and mask binding, the
|
| 441 |
+
persistent io tensors and the cached ``steps4`` program (built once K/V are bound)."""
|
| 442 |
+
import ttnn
|
| 443 |
+
|
| 444 |
+
from models.experimental.gr00t.tt.megakernel import arena as A
|
| 445 |
+
from models.experimental.gr00t.tt.megakernel import dit_program as P
|
| 446 |
+
from models.experimental.gr00t.tt.megakernel.core_map import CoreMap
|
| 447 |
+
|
| 448 |
+
cfg, dl = self.cfg, self.dlayout
|
| 449 |
+
self.mk_arena_dtype = self.policy.mk_arena_dtype_for(cfg.version)
|
| 450 |
+
t0 = time.perf_counter()
|
| 451 |
+
self.arena = A.ArenaLayout.plan(
|
| 452 |
+
cfg.version,
|
| 453 |
+
dtype_stream=self.mk_arena_dtype,
|
| 454 |
+
reader_mode="direct", # the K4 steps4 program streams in direct mode (mk_k4_summary.md §8)
|
| 455 |
+
tile_order="row_major",
|
| 456 |
+
layout_name=dl.layout_name,
|
| 457 |
+
)
|
| 458 |
+
for name, got, want in (
|
| 459 |
+
("m_pad", int(self.arena.m_pad), DIT_QUERY_ROWS_PAD),
|
| 460 |
+
("a_pad", int(self.arena.a_pad), self.a_pad),
|
| 461 |
+
("width", int(self.arena.width), self.width),
|
| 462 |
+
("n_blocks", int(self.arena.n_blocks), int(cfg.dit.n_blocks)),
|
| 463 |
+
("action_rows", tuple(self.arena.action_rows), tuple(self.rows)),
|
| 464 |
+
):
|
| 465 |
+
if got != want:
|
| 466 |
+
raise ValueError(f"megakernel arena {name} {got} != head {want}")
|
| 467 |
+
plan = P.real_arena_plan(cfg.version, int(dl.embodiment_id), self.policy)
|
| 468 |
+
pack = self.arena.pack(plan)
|
| 469 |
+
t1 = time.perf_counter()
|
| 470 |
+
mk = P.DiTMegakernel(
|
| 471 |
+
cfg,
|
| 472 |
+
dl,
|
| 473 |
+
None,
|
| 474 |
+
self.arena,
|
| 475 |
+
CoreMap.from_device(self.device),
|
| 476 |
+
self.device,
|
| 477 |
+
variant="steps4",
|
| 478 |
+
k1_extras=False,
|
| 479 |
+
k4_budget=k4_budget,
|
| 480 |
+
)
|
| 481 |
+
mk.upload(pack)
|
| 482 |
+
mk.bind_masks(self.dit_masks)
|
| 483 |
+
t2 = time.perf_counter()
|
| 484 |
+
# the kernel's persistent io: rank-4 shapes (same DRAM pages: fp32 tile 4096 B, one bf16 row 3072 B) so that
|
| 485 |
+
# ttnn.copy (identical logical shape + layout, dtype conversion on TILE) feeds them from the model's buffers
|
| 486 |
+
io = mk.allocate_steps4_io(taps=True)
|
| 487 |
+
shapes = {
|
| 488 |
+
"x_t": ((1, 1, DIT_QUERY_ROWS_PAD, self.a_pad), ttnn.float32, ttnn.TILE_LAYOUT, torch.float32, 4096),
|
| 489 |
+
"state_features": (
|
| 490 |
+
(1, 1, 1, self.width),
|
| 491 |
+
ttnn.bfloat16,
|
| 492 |
+
ttnn.ROW_MAJOR_LAYOUT,
|
| 493 |
+
torch.bfloat16,
|
| 494 |
+
self.width * 2,
|
| 495 |
+
),
|
| 496 |
+
"action_pred": (
|
| 497 |
+
(1, 1, DIT_QUERY_ROWS_PAD, self.a_pad),
|
| 498 |
+
ttnn.float32,
|
| 499 |
+
ttnn.TILE_LAYOUT,
|
| 500 |
+
torch.float32,
|
| 501 |
+
4096,
|
| 502 |
+
),
|
| 503 |
+
}
|
| 504 |
+
for name in MK_IO_RANK4:
|
| 505 |
+
shape, dt, lay, tdt, page = shapes[name]
|
| 506 |
+
ttnn.deallocate(io[name])
|
| 507 |
+
t = ttnn.from_torch(
|
| 508 |
+
torch.zeros(*shape, dtype=tdt),
|
| 509 |
+
dtype=dt,
|
| 510 |
+
layout=lay,
|
| 511 |
+
device=self.device,
|
| 512 |
+
memory_config=ttnn.DRAM_MEMORY_CONFIG,
|
| 513 |
+
)
|
| 514 |
+
if int(t.buffer_page_size()) != page:
|
| 515 |
+
raise RuntimeError(f"megakernel io {name}: DRAM page {int(t.buffer_page_size())} B != {page}")
|
| 516 |
+
io[name] = t
|
| 517 |
+
if "velocity_out" not in io:
|
| 518 |
+
raise RuntimeError("allocate_steps4_io(taps=True) gave no velocity_out tensor")
|
| 519 |
+
# the upload-time a_rep is freed: the L1 residents live only inside the denoise (module docstring)
|
| 520 |
+
ttnn.deallocate(mk.io["a_rep"])
|
| 521 |
+
mk.io["a_rep"] = None
|
| 522 |
+
self.megakernel = mk
|
| 523 |
+
if kv is not None:
|
| 524 |
+
self.bind_kv(kv)
|
| 525 |
+
ttnn.synchronize_device(self.device)
|
| 526 |
+
self.timing.update(mk_arena_plan_pack_s=t1 - t0, mk_upload_s=t2 - t1, mk_io_s=time.perf_counter() - t2)
|
| 527 |
+
|
| 528 |
+
def _alloc_residents(self) -> Tuple[Any, Any]:
|
| 529 |
+
"""Allocate the kernel's L1 tensors (``a_rep`` ``[110*64, 2048]`` HEIGHT_SHARDED on all cores = ``cb_a``; the
|
| 530 |
+
``x_rep`` tile rows ``[96*32, width]`` on the compute cores = ``cb_x``) without a host write and hand them to
|
| 531 |
+
the launcher (``io["a_rep"]`` / ``k4_io["x_rep"]``; ``dit_program.DiTMegakernel._alloc_a_rep`` /
|
| 532 |
+
``make_x_rows`` shapes). Their contents are never read before the kernel writes them (the multicasts and the
|
| 533 |
+
``LoadX`` of ``sa_embs`` fill them)."""
|
| 534 |
+
import ttnn
|
| 535 |
+
|
| 536 |
+
from models.experimental.gr00t.tt.megakernel import dit_program as P
|
| 537 |
+
|
| 538 |
+
mk = self.megakernel
|
| 539 |
+
crs = mk.core_map.to_core_range_sets()
|
| 540 |
+
a_cols = int(mk.budget.cb_a) // (P.M_ROWS * 2)
|
| 541 |
+
|
| 542 |
+
def alloc(rows: int, cols: int, shard_rows: int, grid: Any) -> Any:
|
| 543 |
+
mem = ttnn.MemoryConfig(
|
| 544 |
+
ttnn.TensorMemoryLayout.HEIGHT_SHARDED,
|
| 545 |
+
ttnn.BufferType.L1,
|
| 546 |
+
ttnn.ShardSpec(grid, [shard_rows, cols], ttnn.ShardOrientation.ROW_MAJOR),
|
| 547 |
+
)
|
| 548 |
+
return ttnn.allocate_tensor_on_device(
|
| 549 |
+
ttnn.Shape([rows, cols]), ttnn.bfloat16, ttnn.TILE_LAYOUT, self.device, mem
|
| 550 |
+
)
|
| 551 |
+
|
| 552 |
+
a_rep = alloc(mk.core_map.n_cores * P.M_ROWS, a_cols, P.M_ROWS, crs["all"])
|
| 553 |
+
x_rep = alloc(mk.core_map.n_compute * P.TILE, self.width, P.TILE, crs["compute"])
|
| 554 |
+
mk.io["a_rep"] = a_rep
|
| 555 |
+
mk.k4_io["x_rep"] = x_rep
|
| 556 |
+
mk.k4_program = None # the descriptor carries the addresses: rebuild on these tensors
|
| 557 |
+
mk.k4_tensors = []
|
| 558 |
+
return a_rep, x_rep
|
| 559 |
+
|
| 560 |
+
def _free_residents(self) -> None:
|
| 561 |
+
import ttnn
|
| 562 |
+
|
| 563 |
+
mk = self.megakernel
|
| 564 |
+
for holder, name in ((mk.io, "a_rep"), (mk.k4_io, "x_rep")):
|
| 565 |
+
t = holder.get(name)
|
| 566 |
+
if t is not None:
|
| 567 |
+
ttnn.deallocate(t)
|
| 568 |
+
holder[name] = None
|
| 569 |
+
mk.k4_program = None
|
| 570 |
+
mk.k4_tensors = []
|
| 571 |
+
|
| 572 |
+
@property
|
| 573 |
+
def mk_io(self) -> Dict[str, Any]:
|
| 574 |
+
"""The megakernel's persistent io tensors (``x_t`` fp32, ``state_features`` bf16 row, ``action_pred`` fp32,
|
| 575 |
+
``velocity_out`` bf16 ``[4*64, A_pad]``, ``x_rep``)."""
|
| 576 |
+
if self.megakernel is None:
|
| 577 |
+
raise RuntimeError("no megakernel: this head runs the ttnn backend")
|
| 578 |
+
return self.megakernel.k4_io
|
| 579 |
+
|
| 580 |
+
def mk_program(self) -> Tuple[Any, List[Any]]:
|
| 581 |
+
"""The ``(ProgramDescriptor, io tensors)`` of the ``steps4`` generic_op on the currently allocated L1
|
| 582 |
+
residents (:meth:`_alloc_residents` first; what trace B captures)."""
|
| 583 |
+
if self.megakernel is None:
|
| 584 |
+
raise RuntimeError("no megakernel: this head runs the ttnn backend")
|
| 585 |
+
if self._kv is None:
|
| 586 |
+
raise ValueError("no hoisted K/V bound (VLAdapterTT HoistOut.kv): pass kv= or call bind_kv first")
|
| 587 |
+
mk = self.megakernel
|
| 588 |
+
if mk.io.get("a_rep") is None or mk.k4_io.get("x_rep") is None:
|
| 589 |
+
raise RuntimeError("the megakernel's L1 residents are not allocated (_alloc_residents)")
|
| 590 |
+
t0 = time.perf_counter()
|
| 591 |
+
out = mk.steps4_program()
|
| 592 |
+
self.timing["mk_program_build_s"] = time.perf_counter() - t0
|
| 593 |
+
return out
|
| 594 |
+
|
| 595 |
# ------------------------------------------------------------------ per-call state
|
| 596 |
@property
|
| 597 |
def kv(self) -> Optional[Dict[int, Tuple[Any, Any]]]:
|
| 598 |
+
return self.step.kv if self.step is not None else self._kv
|
| 599 |
|
| 600 |
def bind_kv(self, kv: Mapping[int, Any]) -> Dict[int, Tuple[Any, Any]]:
|
| 601 |
+
"""Bind the persistent hoisted K/V (``VLAdapterTT`` ``HoistOut.kv``) once; the objects are what trace B captures.
|
| 602 |
+
Under the megakernel the program descriptor is (re)built on the bound handles."""
|
| 603 |
+
if self.step is not None:
|
| 604 |
+
return self.step.bind_kv(kv)
|
| 605 |
+
bound = {int(i): (k, v) for i, (k, v) in kv.items()}
|
| 606 |
+
self.megakernel.bind_kv(bound)
|
| 607 |
+
self._kv = bound
|
| 608 |
+
# validate the whole descriptor once (limits, CTAs, runtime args) on transient residents
|
| 609 |
+
self._alloc_residents()
|
| 610 |
+
try:
|
| 611 |
+
self.mk_program()
|
| 612 |
+
finally:
|
| 613 |
+
self._free_residents()
|
| 614 |
+
return bound
|
| 615 |
|
| 616 |
def update_masks(self, slots: SlotLayout) -> Dict[str, Any]:
|
| 617 |
"""Write this call's per-call masks (``cross_text`` / ``cross_all`` / ``llm`` / ``vl``) into the persistent
|
|
|
|
| 631 |
def set_collect_taps(self, flag: bool) -> None:
|
| 632 |
"""Toggle tap collection on the head and its DiT step (off before a trace capture)."""
|
| 633 |
self.collect_taps = bool(flag)
|
| 634 |
+
if self.step is not None:
|
| 635 |
+
self.step.set_collect_taps(flag)
|
| 636 |
+
if self.encoders is not None:
|
| 637 |
+
self.encoders.set_collect_taps(flag)
|
| 638 |
if not flag:
|
| 639 |
self.taps = {}
|
| 640 |
self._tap_pool = []
|
|
|
|
| 658 |
for m in modules:
|
| 659 |
self._tap_pool.extend(m.debug_taps.values())
|
| 660 |
|
| 661 |
+
def _check_inputs(self, x_t: Any, state_features: Any) -> None:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 662 |
want = (1, 1, DIT_QUERY_ROWS_PAD, self.a_pad)
|
| 663 |
if _shape(x_t) != want:
|
| 664 |
raise ValueError(f"x_t must be {want}, got {_shape(x_t)}")
|
|
|
|
| 666 |
raise TypeError(f"x_t must be {self.xt_dtype} (cfg.sampler.euler_dtype), got {L.dtype_name(x_t.dtype)}")
|
| 667 |
if _shape(state_features) != (1, 1, 1, self.width):
|
| 668 |
raise ValueError(f"state_features must be (1, 1, 1, {self.width}), got {_shape(state_features)}")
|
| 669 |
+
|
| 670 |
+
def denoise(self, x_t: Any, state_features: Any, kv: Optional[Mapping[int, Any]] = None) -> Any:
|
| 671 |
+
"""Run the 4 unrolled steps from ``x_t`` (``[1,1,64,A_pad]`` in the version dtype; the persistent buffer or a
|
| 672 |
+
tensor that is copied into it) with ``state_features [1,1,1,1536]`` bf16 (``VLAdapterTT`` ``HoistOut``) and
|
| 673 |
+
the hoisted ``kv`` (default: the bound handles) -> the persistent ``action_pred`` buffer. Both backends return
|
| 674 |
+
the same buffer (module docstring)."""
|
| 675 |
+
self._check_inputs(x_t, state_features)
|
| 676 |
+
if self.backend == "megakernel":
|
| 677 |
+
return self._denoise_megakernel(x_t, state_features, kv)
|
| 678 |
+
return self._denoise_ttnn(x_t, state_features, kv)
|
| 679 |
+
|
| 680 |
+
def _denoise_ttnn(self, x_t: Any, state_features: Any, kv: Optional[Mapping[int, Any]]) -> Any:
|
| 681 |
+
import ttnn
|
| 682 |
+
|
| 683 |
+
step, enc = self.step, self.encoders
|
| 684 |
+
if step is None or enc is None:
|
| 685 |
+
raise RuntimeError("ttnn backend without its modules")
|
| 686 |
+
if kv is not None and kv is not step.kv:
|
| 687 |
+
step.bind_kv(kv)
|
| 688 |
+
if step.kv is None:
|
| 689 |
raise ValueError("no hoisted K/V bound (VLAdapterTT HoistOut.kv): pass kv= or call bind_kv first")
|
| 690 |
if x_t is not self.x_t:
|
| 691 |
ttnn.copy(x_t, self.x_t)
|
| 692 |
self.taps = {}
|
| 693 |
self._tap_pool = []
|
|
|
|
| 694 |
x = self.x_t
|
| 695 |
for k in range(self.n_steps):
|
| 696 |
a = enc.action_encoder.forward(x, k)
|
|
|
|
| 708 |
ttnn.deallocate(sa)
|
| 709 |
sa = self.sa_embs
|
| 710 |
self._tap(f"sa_embs[k={k}]", sa)
|
| 711 |
+
y = step.forward(sa, k)
|
| 712 |
+
for name, t in step.debug_taps.items():
|
| 713 |
self._tap(f"{name}[k={k}]", t)
|
| 714 |
if sa is not self.sa_embs:
|
| 715 |
self._free(sa)
|
|
|
|
| 730 |
self.n_calls += 1
|
| 731 |
return self.action_pred
|
| 732 |
|
| 733 |
+
def _denoise_megakernel(self, x_t: Any, state_features: Any, kv: Optional[Mapping[int, Any]]) -> Any:
|
| 734 |
+
"""The module docstring's four-op sequence around the ``steps4`` generic_op (all on persistent tensors).
|
| 735 |
+
|
| 736 |
+
The program descriptor is rebuilt on every *untraced* call (``mk_program`` -> ``steps4_program()`` on the
|
| 737 |
+
fresh L1 resident addresses, ~1.4 ms host per ``head_timing.mk_program_build_s`` of the K5 benches); inside
|
| 738 |
+
a trace replay nothing runs on the host, so the cost never reaches the production path (WP-K5 review NB 8).
|
| 739 |
+
"""
|
| 740 |
+
import ttnn
|
| 741 |
+
|
| 742 |
+
if kv is not None and kv is not self._kv:
|
| 743 |
+
bound = {int(i): (k, v) for i, (k, v) in kv.items()}
|
| 744 |
+
if self._kv is None or any(
|
| 745 |
+
bound[i][0] is not self._kv[i][0] or bound[i][1] is not self._kv[i][1] for i in bound
|
| 746 |
+
):
|
| 747 |
+
self.bind_kv(bound)
|
| 748 |
+
io = self.mk_io
|
| 749 |
+
self.taps = {}
|
| 750 |
+
# 1. the version-dtype x_t -> the kernel's fp32 Euler state (a TILE copy; bf16 -> fp32 exact)
|
| 751 |
+
ttnn.copy(x_t, io["x_t"])
|
| 752 |
+
# 2. state_features [1,1,1,1536] bf16 TILE -> the one ROW_MAJOR row the hub reads
|
| 753 |
+
sf_rm = ttnn.to_layout(state_features, ttnn.ROW_MAJOR_LAYOUT)
|
| 754 |
+
ttnn.copy(sf_rm, io["state_features"])
|
| 755 |
+
ttnn.deallocate(sf_rm)
|
| 756 |
+
# 3. the denoise on freshly allocated L1 residents (module docstring)
|
| 757 |
+
self._alloc_residents()
|
| 758 |
+
try:
|
| 759 |
+
program, tensors = self.mk_program()
|
| 760 |
+
ttnn.generic_op(tensors, program)
|
| 761 |
+
finally:
|
| 762 |
+
self._free_residents()
|
| 763 |
+
# 4. fp32 action_pred -> the version dtype buffer (N1.6 / N1.7 values are bf16-rounded per step in-kernel)
|
| 764 |
+
ttnn.copy(io["action_pred"], self.action_pred)
|
| 765 |
+
if self.collect_taps:
|
| 766 |
+
vel = L.to_torch(io["velocity_out"]) # [4*64, A_pad] fp32 host (bf16 values)
|
| 767 |
+
for k in range(self.n_steps):
|
| 768 |
+
rows = vel[k * DIT_QUERY_ROWS_PAD : (k + 1) * DIT_QUERY_ROWS_PAD]
|
| 769 |
+
self._tap(f"action_decoder_out[k={k}]", rows.reshape(1, 1, DIT_QUERY_ROWS_PAD, self.a_pad).contiguous())
|
| 770 |
+
self._tap("action_pred_normalized", self.action_pred)
|
| 771 |
+
self.n_calls += 1
|
| 772 |
+
return self.action_pred
|
| 773 |
+
|
| 774 |
__call__ = denoise
|
| 775 |
|
| 776 |
# ------------------------------------------------------------------ host side
|
|
|
|
| 782 |
""":func:`taps_to_golden` of the last ``denoise`` call's taps."""
|
| 783 |
return taps_to_golden(self.taps, self.cfg, self.dlayout, b3=b3, block_taps=block_taps)
|
| 784 |
|
| 785 |
+
def velocity_taps_host(self) -> torch.Tensor:
|
| 786 |
+
"""Megakernel backend: the last launch's per-step decoder outputs ``[n_steps, 64, A_pad]`` fp32 (bf16 values;
|
| 787 |
+
the kernel's ``velocity_out`` tensor, written on every launch, traced or not)."""
|
| 788 |
+
vel = L.to_torch(self.mk_io["velocity_out"])
|
| 789 |
+
return vel.reshape(self.n_steps, DIT_QUERY_ROWS_PAD, self.a_pad)
|
| 790 |
+
|
| 791 |
def free_taps(self) -> None:
|
| 792 |
+
"""Deallocate every collected device tap that is not a persistent buffer (each object once; host copies skipped)."""
|
| 793 |
import ttnn
|
| 794 |
|
| 795 |
keep = {id(self.x_t), id(self.action_pred), id(self.sa_embs)}
|
| 796 |
seen: set = set()
|
| 797 |
+
step_taps = list(self.step.debug_taps.values()) if self.step is not None else []
|
| 798 |
+
for t in list(self.taps.values()) + step_taps + self._tap_pool:
|
| 799 |
+
if t is None or isinstance(t, torch.Tensor) or id(t) in seen or id(t) in keep:
|
| 800 |
continue
|
| 801 |
seen.add(id(t))
|
| 802 |
ttnn.deallocate(t)
|
| 803 |
self.taps = {}
|
| 804 |
self._tap_pool = []
|
| 805 |
+
if self.step is not None:
|
| 806 |
+
self.step.debug_taps = {}
|
| 807 |
+
if self.encoders is not None:
|
| 808 |
+
self.encoders.set_collect_taps(self.collect_taps) # clears the modules' debug_taps dicts
|
| 809 |
|
| 810 |
def release(self) -> None:
|
| 811 |
+
"""Deallocate the persistent buffers this head owns (masks only when it allocated them; never the K/V) and,
|
| 812 |
+
under the megakernel backend, the arena, tables, io tensors and L1 residents of the launcher."""
|
| 813 |
import ttnn
|
| 814 |
|
| 815 |
+
if self.megakernel is not None:
|
| 816 |
+
mk = self.megakernel
|
| 817 |
+
# the bound K/V and mask handles belong to the adapter / executor: unhook them before the launcher frees its io
|
| 818 |
+
for name in list(mk.io):
|
| 819 |
+
if name.startswith("kv_") or name.endswith("_mask"):
|
| 820 |
+
mk.io[name] = None
|
| 821 |
+
mk.deallocate()
|
| 822 |
+
self.megakernel = None
|
| 823 |
for t in (self.x_t, self.action_pred, self.sa_embs):
|
| 824 |
if t is not None:
|
| 825 |
ttnn.deallocate(t)
|
|
|
|
| 830 |
self.masks = {}
|
| 831 |
|
| 832 |
def describe(self) -> Dict[str, Any]:
|
| 833 |
+
d: Dict[str, Any] = self.step.describe() if self.step is not None else self.plan.describe()
|
| 834 |
d.update(
|
| 835 |
{
|
| 836 |
"backend": self.backend,
|
|
|
|
| 839 |
"buckets": list(self.buckets),
|
| 840 |
"persist_sa_embs": self.sa_embs is not None,
|
| 841 |
"owns_masks": self.owns_masks,
|
| 842 |
+
"fuse_swish": self.encoders.action_encoder.fuse_swish if self.encoders is not None else None,
|
| 843 |
"buffer_mem": getattr(self.buffer_mem, "kind", str(self.buffer_mem)),
|
| 844 |
+
"timing": dict(self.timing),
|
| 845 |
}
|
| 846 |
)
|
| 847 |
+
if self.megakernel is not None:
|
| 848 |
+
mk = self.megakernel
|
| 849 |
+
st = mk.stats() if mk.last_spec is not None else None
|
| 850 |
+
d["megakernel"] = {
|
| 851 |
+
"arena_dtype": self.mk_arena_dtype,
|
| 852 |
+
"reader_mode": self.arena.reader_mode,
|
| 853 |
+
"n_blocks": int(self.arena.n_blocks),
|
| 854 |
+
"key_sets": sorted(set(self.arena.key_sets)),
|
| 855 |
+
"arena_bytes_per_step": int(mk.arena_bytes_per_step()),
|
| 856 |
+
"cb_w_slots": dict(mk.k4.cb_w_slots),
|
| 857 |
+
"stats": st.__dict__ if st is not None else None,
|
| 858 |
+
"compile_s": {str(k): v for k, v in mk.compile_seconds.items()},
|
| 859 |
+
"io": {k: (list(int(x) for x in v.shape) if v is not None else None) for k, v in mk.k4_io.items()},
|
| 860 |
+
"launches": int(mk.launches),
|
| 861 |
+
}
|
| 862 |
return d
|
| 863 |
|
| 864 |
|
| 865 |
__all__ = [
|
| 866 |
+
"DERIVED_STEP_TAPS",
|
| 867 |
"FINAL_TAPS",
|
| 868 |
+
"MK_IO_RANK4",
|
| 869 |
+
"MK_STEP_TAPS",
|
| 870 |
"STEP_TAPS",
|
| 871 |
"ActionHeadTT",
|
| 872 |
"HeadEncoders",
|
| 873 |
"allocate_mask_buffers",
|
| 874 |
+
"head_tap_names",
|
| 875 |
"placeholder_masks",
|
| 876 |
"taps_to_golden",
|
| 877 |
]
|
code/models/experimental/gr00t/tt/device.py
CHANGED
|
@@ -141,20 +141,28 @@ def open_gr00t_device(
|
|
| 141 |
l1_small_size: int = DEFAULT_L1_SMALL_SIZE,
|
| 142 |
num_command_queues: int = 1,
|
| 143 |
device_id: int = 0,
|
|
|
|
| 144 |
) -> Any:
|
| 145 |
"""Open device ``device_id`` with the program cache enabled and assert the p150a facts this plan relies on.
|
| 146 |
|
| 147 |
-
``ttnn.open_device(device_id, l1_small_size, trace_region_size[, num_command_queues])``
|
| 148 |
(``ttnn-nanobind/device.cpp:71-80``, mb2 §5) returns the 1x1 ``MeshDevice``; the compute grid must be 11x10 and
|
| 149 |
the DRAM grid 8x1 (mb1 Table 5a) -- every core map and every program config in ``tt.layers`` is sized for that
|
| 150 |
-
card, so anything else raises instead of silently running a wrong configuration.
|
| 151 |
-
|
|
|
|
|
|
|
|
|
|
| 152 |
"""
|
| 153 |
import ttnn
|
| 154 |
|
| 155 |
kwargs: Dict[str, Any] = dict(device_id=device_id, l1_small_size=l1_small_size, trace_region_size=trace_region_size)
|
| 156 |
if num_command_queues != 1:
|
| 157 |
-
kwargs["num_command_queues"] = num_command_queues
|
|
|
|
|
|
|
|
|
|
|
|
|
| 158 |
device = ttnn.open_device(**kwargs)
|
| 159 |
try:
|
| 160 |
device.enable_program_cache()
|
|
|
|
| 141 |
l1_small_size: int = DEFAULT_L1_SMALL_SIZE,
|
| 142 |
num_command_queues: int = 1,
|
| 143 |
device_id: int = 0,
|
| 144 |
+
worker_l1_size: Optional[int] = None,
|
| 145 |
) -> Any:
|
| 146 |
"""Open device ``device_id`` with the program cache enabled and assert the p150a facts this plan relies on.
|
| 147 |
|
| 148 |
+
``ttnn.open_device(device_id, l1_small_size, trace_region_size[, num_command_queues][, worker_l1_size])``
|
| 149 |
(``ttnn-nanobind/device.cpp:71-80``, mb2 §5) returns the 1x1 ``MeshDevice``; the compute grid must be 11x10 and
|
| 150 |
the DRAM grid 8x1 (mb1 Table 5a) -- every core map and every program config in ``tt.layers`` is sized for that
|
| 151 |
+
card, so anything else raises instead of silently running a wrong configuration. ``worker_l1_size`` (bytes of
|
| 152 |
+
allocatable L1 per worker core; ``None`` = the firmware default, 1,461,248 B on the p150a) is what
|
| 153 |
+
``tt.model.open_model_device`` reduces for the megakernel DiT backend (the kernel-config ring buffer grows by the
|
| 154 |
+
same amount, ``mk_k4_summary.md`` §7.2 fact 2); the value actually in force is checked there. Callers hold the
|
| 155 |
+
device lock (``bin/with-device.sh``).
|
| 156 |
"""
|
| 157 |
import ttnn
|
| 158 |
|
| 159 |
kwargs: Dict[str, Any] = dict(device_id=device_id, l1_small_size=l1_small_size, trace_region_size=trace_region_size)
|
| 160 |
if num_command_queues != 1:
|
| 161 |
+
kwargs["num_command_queues"] = int(num_command_queues)
|
| 162 |
+
if worker_l1_size is not None:
|
| 163 |
+
if int(worker_l1_size) <= 0:
|
| 164 |
+
raise ValueError(f"worker_l1_size must be a positive number of bytes, got {worker_l1_size}")
|
| 165 |
+
kwargs["worker_l1_size"] = int(worker_l1_size)
|
| 166 |
device = ttnn.open_device(**kwargs)
|
| 167 |
try:
|
| 168 |
device.enable_program_cache()
|
code/models/experimental/gr00t/tt/model.py
CHANGED
|
@@ -27,7 +27,19 @@ layout groups them into ``A`` / ``B``)::
|
|
| 27 |
side effect: hoisted K_i / V_i in the adapter's persistent buffers
|
| 28 |
denoise (B) ActionHeadTT.denoise inputs: x_t, cross_text_mask | cross_all_mask; constants self_mask,
|
| 29 |
cross_image_mask (when L_img_pad != L_img_real) -> action_pred (a clone
|
| 30 |
-
of the head's persistent buffer, so the executor owns the trace output)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 31 |
|
| 32 |
Every per-call tensor -- the masks included -- is an executor buffer (static-shape §9 / §10: written unconditionally
|
| 33 |
before every replay, address-checked, presence fixed per layout), so ``ActionHeadTT`` is built on the executor's
|
|
@@ -69,7 +81,7 @@ import torch
|
|
| 69 |
|
| 70 |
from models.experimental.gr00t.common.configs import DIT_QUERY_ROWS_PAD, GR00TConfig, LayoutConfig, get_config
|
| 71 |
from models.experimental.gr00t.tt import layers as L
|
| 72 |
-
from models.experimental.gr00t.tt.action_head import ActionHeadTT
|
| 73 |
from models.experimental.gr00t.tt.adapter import STATE_ROWS, HoistOut, VLAdapterTT, slots_to_rows, state_rows_host
|
| 74 |
from models.experimental.gr00t.tt.backbone import (
|
| 75 |
IN_COS,
|
|
@@ -86,7 +98,15 @@ from models.experimental.gr00t.tt.backbone import (
|
|
| 86 |
BackboneRecipe,
|
| 87 |
BackboneTT,
|
| 88 |
)
|
| 89 |
-
from models.experimental.gr00t.tt.device import
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 90 |
from models.experimental.gr00t.tt.encoders import noise_to_xt, xt_dtype, xt_to_action_pred
|
| 91 |
from models.experimental.gr00t.tt.layout import MASK_NAMES, DeviceLayout, StaticShapeError
|
| 92 |
from models.experimental.gr00t.tt.policy import TTPolicy
|
|
@@ -203,6 +223,7 @@ class ModelRecipe:
|
|
| 203 |
trace_layout: str
|
| 204 |
dit_backend: str
|
| 205 |
dtype_policy: str
|
|
|
|
| 206 |
|
| 207 |
@classmethod
|
| 208 |
def from_config(cls, cfg: GR00TConfig, dlayout: DeviceLayout, policy: TTPolicy) -> "ModelRecipe":
|
|
@@ -239,6 +260,7 @@ class ModelRecipe:
|
|
| 239 |
trace_layout=policy.trace_layout,
|
| 240 |
dit_backend=policy.dit_backend,
|
| 241 |
dtype_policy=policy.dtype_policy,
|
|
|
|
| 242 |
)
|
| 243 |
|
| 244 |
@property
|
|
@@ -308,10 +330,14 @@ def model_input_specs(
|
|
| 308 |
return specs
|
| 309 |
|
| 310 |
|
| 311 |
-
def expected_tap_keys(cfg: GR00TConfig, dlayout: DeviceLayout) -> Tuple[str, ...]:
|
| 312 |
"""The golden-named keys ``run_untraced(return_taps=True)`` produces for this version (plan §7.2 matrix; keys in
|
| 313 |
-
``GoldenSet.threshold_key`` form). Per-block DiT taps cover every block; the goldens hold blocks 0 / 1 / last.
|
| 314 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 315 |
keys: List[str] = ["vit_embeddings", "vit_block_first", "vit_block_last"]
|
| 316 |
if r.tower == "siglip":
|
| 317 |
keys.append("vit_post_ln")
|
|
@@ -325,22 +351,161 @@ def expected_tap_keys(cfg: GR00TConfig, dlayout: DeviceLayout) -> Tuple[str, ...
|
|
| 325 |
if r.has_vlsa:
|
| 326 |
keys.append("vl_self_attention_out")
|
| 327 |
keys += ["dit_encoder_hidden_states", "state_features"]
|
| 328 |
-
|
| 329 |
-
keys += [f"action_encoder_out[k={k}]", f"sa_embs[k={k}]"]
|
| 330 |
-
keys += [f"dit_block{i}_out[k={k}]" for i in range(r.n_blocks)]
|
| 331 |
-
keys += [
|
| 332 |
-
f"dit_block_last_out[k={k}]",
|
| 333 |
-
f"dit_out[k={k}]",
|
| 334 |
-
f"action_decoder_out[k={k}]",
|
| 335 |
-
f"pred_velocity[k={k}]",
|
| 336 |
-
f"actions_out[k={k}]",
|
| 337 |
-
]
|
| 338 |
-
keys += ["action_pred_normalized", "action_pred_valid"]
|
| 339 |
if len(set(keys)) != len(keys):
|
| 340 |
raise AssertionError("duplicate tap keys")
|
| 341 |
return tuple(keys)
|
| 342 |
|
| 343 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 344 |
# --------------------------------------------------------------------------------------------------------------------
|
| 345 |
# bench / serve replayer contract (benchmarks/bench_e2e.py REPLAYER_MEMBERS)
|
| 346 |
# --------------------------------------------------------------------------------------------------------------------
|
|
@@ -438,6 +603,7 @@ class Gr00tTT:
|
|
| 438 |
backbone: Optional[BackboneTT] = None,
|
| 439 |
embed_gather: str = "host",
|
| 440 |
cq1_uploads: Optional[bool] = None,
|
|
|
|
| 441 |
) -> None:
|
| 442 |
"""Assemble the modules on an open device (``from_pretrained`` is the public entry point).
|
| 443 |
|
|
@@ -447,6 +613,8 @@ class Gr00tTT:
|
|
| 447 |
``"host"`` (the WP-12 host gather; the constructor default, so stand-in weights keep working -- the production
|
| 448 |
factory :meth:`from_pretrained` defaults to :data:`DEFAULT_EMBED_GATHER`; module docstring).
|
| 449 |
``cq1_uploads``: per-call input writes on CQ 1 (``None`` = on iff the device has two command queues).
|
|
|
|
|
|
|
| 450 |
"""
|
| 451 |
if not isinstance(cfg, GR00TConfig) or not isinstance(dlayout, DeviceLayout):
|
| 452 |
raise TypeError("Gr00tTT(cfg: GR00TConfig, dlayout: DeviceLayout, policy: TTPolicy, W, device)")
|
|
@@ -467,6 +635,17 @@ class Gr00tTT:
|
|
| 467 |
f"weights are for {getattr(W, 'version', '?')}/slot {getattr(W, 'embodiment_id', '?')}, "
|
| 468 |
f"model is {cfg.version}/slot {e}"
|
| 469 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 470 |
self.cfg, self.dlayout, self.policy, self.W, self.device = cfg, dlayout, policy, W, device
|
| 471 |
self.version: str = cfg.version
|
| 472 |
self.layout: LayoutConfig = dlayout.layout
|
|
@@ -479,6 +658,7 @@ class Gr00tTT:
|
|
| 479 |
self.embed_gather = embed_gather
|
| 480 |
self.specs: Dict[str, InputSpec] = model_input_specs(cfg, dlayout, policy, embed_gather)
|
| 481 |
self.timing: Dict[str, float] = {}
|
|
|
|
| 482 |
self.n_predict = 0
|
| 483 |
self._b3: Optional[torch.Tensor] = None
|
| 484 |
self._hoist: Optional[HoistOut] = None
|
|
@@ -497,7 +677,19 @@ class Gr00tTT:
|
|
| 497 |
raise ValueError("the injected backbone must be built with collect_taps=False and hold no taps")
|
| 498 |
self.backbone = backbone
|
| 499 |
t1 = time.perf_counter()
|
| 500 |
-
self.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 501 |
t2 = time.perf_counter()
|
| 502 |
if embed_gather == "device":
|
| 503 |
if not hasattr(W, "_embed_table_tensor"):
|
|
@@ -621,6 +813,8 @@ class Gr00tTT:
|
|
| 621 |
embed_gather: str = DEFAULT_EMBED_GATHER,
|
| 622 |
cq1_uploads: Optional[bool] = None,
|
| 623 |
num_command_queues: int = 2,
|
|
|
|
|
|
|
| 624 |
) -> "Gr00tTT":
|
| 625 |
"""Load the weights and assemble the model (plan §2.2 row ``tt/model.py``).
|
| 626 |
|
|
@@ -628,14 +822,22 @@ class Gr00tTT:
|
|
| 628 |
version: ``"n15" | "n16" | "n17"``.
|
| 629 |
embodiment: embodiment tag (default: the layout's ``embodiment_tag``); must be the slot the layout is baked for.
|
| 630 |
layout_name: one of ``cfg.layouts`` (default ``cfg.canonical_layout_name``, D2).
|
| 631 |
-
policy: ``TTPolicy`` (default ``TTPolicy()``: mixed_dit,
|
| 632 |
-
|
| 633 |
-
|
|
|
|
|
|
|
|
|
|
| 634 |
cache_root: ``HostPlanCache`` root for ``TTWeights.load`` (default: the project cache).
|
| 635 |
warm_inputs: a ``ModelInputs`` (or ``Observation``) to warm and capture on immediately; otherwise the
|
| 636 |
first ``predict`` / ``predict_normalized`` / ``tracer.upload`` call warms and captures on its inputs.
|
| 637 |
noise: the noise for that warm pass (default: the version's seeded noise).
|
| 638 |
-
embed_gather, cq1_uploads: see the constructor (module docstring
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 639 |
"""
|
| 640 |
from models.experimental.gr00t.common.checkpoint import LazyCheckpoint
|
| 641 |
from models.experimental.gr00t.tt.weights import TTWeights
|
|
@@ -653,11 +855,17 @@ class Gr00tTT:
|
|
| 653 |
f"embodiment {tag!r} is slot {e} -- pick a layout of that embodiment"
|
| 654 |
)
|
| 655 |
dlayout = DeviceLayout.from_layout(cfg, layout)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 656 |
owns = device is None
|
| 657 |
-
dev =
|
| 658 |
try:
|
| 659 |
t0 = time.perf_counter()
|
| 660 |
-
W = TTWeights.load(version, e, pol, dev, cache_root=cache_root, layout_name=layout.name)
|
| 661 |
load_s = time.perf_counter() - t0
|
| 662 |
ck = LazyCheckpoint(version)
|
| 663 |
model = cls(
|
|
@@ -671,12 +879,14 @@ class Gr00tTT:
|
|
| 671 |
owns_device=owns,
|
| 672 |
embed_gather=embed_gather,
|
| 673 |
cq1_uploads=cq1_uploads,
|
|
|
|
| 674 |
)
|
| 675 |
except Exception:
|
| 676 |
if owns:
|
| 677 |
close_gr00t_device(dev)
|
| 678 |
raise
|
| 679 |
model.timing["weights_load_s"] = load_s
|
|
|
|
| 680 |
if warm_inputs is not None:
|
| 681 |
mi = warm_inputs if not _is_observation(warm_inputs) else model.encode(warm_inputs)
|
| 682 |
model.warm_and_capture(mi, noise)
|
|
@@ -1005,7 +1215,7 @@ class Gr00tTT:
|
|
| 1005 |
else:
|
| 1006 |
taps[name] = slots_to_rows(L.to_torch(t), mi.slots)
|
| 1007 |
taps.update(self.head.golden_taps(b3=self.b3_bias(), block_taps=True))
|
| 1008 |
-
missing = [k for k in expected_tap_keys(cfg, dl) if k not in taps]
|
| 1009 |
if missing:
|
| 1010 |
raise RuntimeError(f"run_untraced produced no tap for {missing}")
|
| 1011 |
return taps
|
|
@@ -1077,6 +1287,8 @@ class Gr00tTT:
|
|
| 1077 |
"policy": self.policy.to_dict(),
|
| 1078 |
"recipe": self.recipe.describe(),
|
| 1079 |
"embed_gather": self.embed_gather,
|
|
|
|
|
|
|
| 1080 |
"gather": None if self.gather is None else self.gather.describe(),
|
| 1081 |
"input_specs": {n: s.describe() for n, s in self.specs.items()},
|
| 1082 |
"executor": self.executor.describe(),
|
|
@@ -1087,6 +1299,7 @@ class Gr00tTT:
|
|
| 1087 |
"device_bytes": float(self.W.device_bytes),
|
| 1088 |
"load_stats": dict(getattr(self.W, "load_stats", {})),
|
| 1089 |
"dtype_histogram": self.W.dtype_histogram() if hasattr(self.W, "dtype_histogram") else None,
|
|
|
|
| 1090 |
},
|
| 1091 |
}
|
| 1092 |
for name in ("backbone", "adapter", "head"):
|
|
@@ -1101,6 +1314,14 @@ class Gr00tTT:
|
|
| 1101 |
# --------------------------------------------------------------------------------------------------------------------
|
| 1102 |
|
| 1103 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1104 |
def _is_observation(x: Any) -> bool:
|
| 1105 |
return hasattr(x, "images") and hasattr(x, "instruction") and hasattr(x, "embodiment_tag")
|
| 1106 |
|
|
@@ -1139,6 +1360,7 @@ __all__ = [
|
|
| 1139 |
"IN_VL_MASK",
|
| 1140 |
"IN_XT",
|
| 1141 |
"MASK_BUFFER_NAMES",
|
|
|
|
| 1142 |
"NOISE_SEED",
|
| 1143 |
"N_STEPS",
|
| 1144 |
"OUT_ACTION_PRED",
|
|
@@ -1150,8 +1372,17 @@ __all__ = [
|
|
| 1150 |
"Gr00tTT",
|
| 1151 |
"ModelRecipe",
|
| 1152 |
"TraceReplayer",
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1153 |
"expected_tap_keys",
|
|
|
|
|
|
|
|
|
|
| 1154 |
"model_input_specs",
|
|
|
|
| 1155 |
"text_ids_spec",
|
| 1156 |
"unproduced_taps",
|
|
|
|
| 1157 |
]
|
|
|
|
| 27 |
side effect: hoisted K_i / V_i in the adapter's persistent buffers
|
| 28 |
denoise (B) ActionHeadTT.denoise inputs: x_t, cross_text_mask | cross_all_mask; constants self_mask,
|
| 29 |
cross_image_mask (when L_img_pad != L_img_real) -> action_pred (a clone
|
| 30 |
+
of the head's persistent buffer, so the executor owns the trace output).
|
| 31 |
+
``policy.dit_backend``: the Stage-1 ttnn ops or the ``steps4``
|
| 32 |
+
megakernel ``generic_op`` (WP-K5) -- same inputs, same output buffer.
|
| 33 |
+
|
| 34 |
+
The megakernel backend -- the ``TTPolicy`` **default** since 2026-09-18 (``tt.policy.DEFAULT_DIT_BACKEND``; the ttnn
|
| 35 |
+
Stage-1 path stays selectable through ``TTPolicy(dit_backend="ttnn")`` / ``GR00T_DIT_BACKEND=ttnn``) -- needs the
|
| 36 |
+
device opened with a smaller ``worker_l1_size`` (the kernel-config ring buffer of ``mk_k4_summary.md`` §7.2 fact 2;
|
| 37 |
+
:data:`MK_L1_CUT_BYTES_DEFAULT`): :func:`open_model_device` does that from the policy and
|
| 38 |
+
``from_pretrained(device=None)`` uses it; a device that was opened another way is refused at build
|
| 39 |
+
(:func:`device_worker_l1_size`). Under the megakernel the ttnn DiT weights are not uploaded
|
| 40 |
+
(:func:`megakernel_weight_names`: the arena replaces the ``dit`` / ``dit_precompute`` / ``encoders`` tensors except
|
| 41 |
+
the adapter's ``dit.block.{i}.Wkv`` / ``bkv``, ``enc.state.*`` and ``enc.action.b3_pos`` that ``TTWeights.constants``
|
| 42 |
+
reads); ``from_pretrained(load_ttnn_dit_weights=True)`` keeps the full Stage-1 upload.
|
| 43 |
|
| 44 |
Every per-call tensor -- the masks included -- is an executor buffer (static-shape §9 / §10: written unconditionally
|
| 45 |
before every replay, address-checked, presence fixed per layout), so ``ActionHeadTT`` is built on the executor's
|
|
|
|
| 81 |
|
| 82 |
from models.experimental.gr00t.common.configs import DIT_QUERY_ROWS_PAD, GR00TConfig, LayoutConfig, get_config
|
| 83 |
from models.experimental.gr00t.tt import layers as L
|
| 84 |
+
from models.experimental.gr00t.tt.action_head import ActionHeadTT, head_tap_names
|
| 85 |
from models.experimental.gr00t.tt.adapter import STATE_ROWS, HoistOut, VLAdapterTT, slots_to_rows, state_rows_host
|
| 86 |
from models.experimental.gr00t.tt.backbone import (
|
| 87 |
IN_COS,
|
|
|
|
| 98 |
BackboneRecipe,
|
| 99 |
BackboneTT,
|
| 100 |
)
|
| 101 |
+
from models.experimental.gr00t.tt.device import (
|
| 102 |
+
DEFAULT_L1_SMALL_SIZE,
|
| 103 |
+
DEFAULT_TRACE_REGION_SIZE,
|
| 104 |
+
DRAM,
|
| 105 |
+
L1,
|
| 106 |
+
close_gr00t_device,
|
| 107 |
+
open_gr00t_device,
|
| 108 |
+
resolve_mem,
|
| 109 |
+
)
|
| 110 |
from models.experimental.gr00t.tt.encoders import noise_to_xt, xt_dtype, xt_to_action_pred
|
| 111 |
from models.experimental.gr00t.tt.layout import MASK_NAMES, DeviceLayout, StaticShapeError
|
| 112 |
from models.experimental.gr00t.tt.policy import TTPolicy
|
|
|
|
| 223 |
trace_layout: str
|
| 224 |
dit_backend: str
|
| 225 |
dtype_policy: str
|
| 226 |
+
mk_arena_dtype: Optional[str] = None # the megakernel's weight arena (resolved per version); None for "ttnn"
|
| 227 |
|
| 228 |
@classmethod
|
| 229 |
def from_config(cls, cfg: GR00TConfig, dlayout: DeviceLayout, policy: TTPolicy) -> "ModelRecipe":
|
|
|
|
| 260 |
trace_layout=policy.trace_layout,
|
| 261 |
dit_backend=policy.dit_backend,
|
| 262 |
dtype_policy=policy.dtype_policy,
|
| 263 |
+
mk_arena_dtype=policy.mk_arena_dtype_for(cfg.version) if policy.dit_backend == "megakernel" else None,
|
| 264 |
)
|
| 265 |
|
| 266 |
@property
|
|
|
|
| 330 |
return specs
|
| 331 |
|
| 332 |
|
| 333 |
+
def expected_tap_keys(cfg: GR00TConfig, dlayout: DeviceLayout, dit_backend: str = "ttnn") -> Tuple[str, ...]:
|
| 334 |
"""The golden-named keys ``run_untraced(return_taps=True)`` produces for this version (plan §7.2 matrix; keys in
|
| 335 |
+
``GoldenSet.threshold_key`` form). Per-block DiT taps cover every block; the goldens hold blocks 0 / 1 / last.
|
| 336 |
+
``dit_backend`` selects the head's tap set (``tt.action_head.head_tap_names``): the megakernel exposes only the
|
| 337 |
+
per-step decoder output and the final prediction (WP-K5). The parameter default stays ``"ttnn"`` on purpose:
|
| 338 |
+
it is the **full** Stage-1 set (the superset the golden tap map is checked against on CPU); model-level callers
|
| 339 |
+
pass ``model.policy.dit_backend``."""
|
| 340 |
+
r = ModelRecipe.from_config(cfg, dlayout, TTPolicy(dit_backend=dit_backend))
|
| 341 |
keys: List[str] = ["vit_embeddings", "vit_block_first", "vit_block_last"]
|
| 342 |
if r.tower == "siglip":
|
| 343 |
keys.append("vit_post_ln")
|
|
|
|
| 351 |
if r.has_vlsa:
|
| 352 |
keys.append("vl_self_attention_out")
|
| 353 |
keys += ["dit_encoder_hidden_states", "state_features"]
|
| 354 |
+
keys += list(head_tap_names(r.dit_backend, r.n_steps, r.n_blocks))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 355 |
if len(set(keys)) != len(keys):
|
| 356 |
raise AssertionError("duplicate tap keys")
|
| 357 |
return tuple(keys)
|
| 358 |
|
| 359 |
|
| 360 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 361 |
+
# device opening per policy (WP-K5: the megakernel's worker_l1_size)
|
| 362 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 363 |
+
|
| 364 |
+
|
| 365 |
+
#: Bytes taken off every core's allocatable L1 (``worker_l1_size``) for the megakernel backend: the kernel-config
|
| 366 |
+
#: ring buffer grows by the same amount (``mk_k4_summary.md`` §7.2 fact 2). Measured compute-core K4 program sizes
|
| 367 |
+
#: (``program.cpp:3224``): N1.6 91,168 B, N1.5 108,912 B (its future-token / mask paths) against the p150a default
|
| 368 |
+
#: ring of 70,656 B, so >= 38,256 B are needed; 64 KiB leaves 27 KB of headroom (K4's own tests use 96 KiB,
|
| 369 |
+
#: ``dit_program.K4_WORKER_L1_SIZE``). Any cut >= 32 KiB breaks the N1.5 / N1.7 **adapter** stage with L1
|
| 370 |
+
#: intermediates: its VL-SA programs (S_pad 544 / 480) allocate static CBs up to 1,405,952 B and the stage's L1
|
| 371 |
+
#: activations extend down to <= 1,339,392 B (``program.cpp:2149``; WP-K5 bring-up sweeps at 96 / 48 / 40 KiB), i.e.
|
| 372 |
+
#: the two requirements leave no common value -- hence :data:`ADAPTER_MEM_BY_BACKEND` moves the adapter's
|
| 373 |
+
#: intermediates to DRAM under the megakernel backend. ``GR00T_MK_L1_CUT_KIB`` overrides the value (sweeps).
|
| 374 |
+
MK_L1_CUT_BYTES_DEFAULT: int = 64 * 1024
|
| 375 |
+
#: Memory of the ``VLAdapterTT`` intermediates per DiT backend (``VLAdapterTT(mem=)``; the persistent K/V stay in
|
| 376 |
+
#: DRAM either way). ``"DRAM"`` under the megakernel frees the ~230 KB of per-core L1 the VL-SA stage needs next to
|
| 377 |
+
#: its static CBs (comment above); the values are identical (same ops, same dtypes), only the placement changes.
|
| 378 |
+
ADAPTER_MEM_BY_BACKEND: Dict[str, str] = {"ttnn": "L1", "megakernel": "DRAM"}
|
| 379 |
+
MK_L1_CUT_ENV = "GR00T_MK_L1_CUT_KIB"
|
| 380 |
+
|
| 381 |
+
|
| 382 |
+
def mk_l1_cut_bytes() -> int:
|
| 383 |
+
""":data:`MK_L1_CUT_BYTES_DEFAULT` or the :data:`MK_L1_CUT_ENV` override (KiB, positive integer)."""
|
| 384 |
+
import os
|
| 385 |
+
|
| 386 |
+
raw = os.environ.get(MK_L1_CUT_ENV)
|
| 387 |
+
if not raw:
|
| 388 |
+
return MK_L1_CUT_BYTES_DEFAULT
|
| 389 |
+
kib = int(raw)
|
| 390 |
+
if kib <= 0:
|
| 391 |
+
raise ValueError(f"{MK_L1_CUT_ENV}={raw!r} must be a positive number of KiB")
|
| 392 |
+
return kib * 1024
|
| 393 |
+
|
| 394 |
+
|
| 395 |
+
def worker_l1_size_for(policy: TTPolicy) -> Optional[int]:
|
| 396 |
+
"""The ``worker_l1_size`` the device must be opened with for ``policy``: ``descriptors.L1_USABLE_BYTES`` (the
|
| 397 |
+
p150a default, 1,461,248 B) minus :func:`mk_l1_cut_bytes` for the megakernel backend (the K4 binaries need the
|
| 398 |
+
larger kernel-config ring buffer that the cut buys, module constant above), ``None`` (the ttnn default)
|
| 399 |
+
otherwise."""
|
| 400 |
+
if not isinstance(policy, TTPolicy):
|
| 401 |
+
raise TypeError(f"policy must be a TTPolicy, got {type(policy).__name__}")
|
| 402 |
+
if policy.dit_backend != "megakernel":
|
| 403 |
+
return None
|
| 404 |
+
from models.experimental.gr00t.tt.megakernel.descriptors import L1_USABLE_BYTES
|
| 405 |
+
|
| 406 |
+
return int(L1_USABLE_BYTES) - mk_l1_cut_bytes()
|
| 407 |
+
|
| 408 |
+
|
| 409 |
+
def device_worker_l1_size(device: Any) -> int:
|
| 410 |
+
"""The user-allocatable L1 per worker core the device was opened with: ``ttnn.reports.get_device_info(device)``
|
| 411 |
+
``.cb_limit`` (= the ``worker_l1_size`` of ``open_device``; the ``.worker_l1_size`` field there is the physical
|
| 412 |
+
1.5 MiB). Measured on the p150a: 1,461,248 B by default, 1,362,944 B under ``K4_WORKER_L1_SIZE``."""
|
| 413 |
+
import ttnn
|
| 414 |
+
|
| 415 |
+
return int(ttnn._ttnn.reports.get_device_info(device).cb_limit)
|
| 416 |
+
|
| 417 |
+
|
| 418 |
+
def open_model_device(
|
| 419 |
+
policy: TTPolicy,
|
| 420 |
+
*,
|
| 421 |
+
trace_region_size: int = DEFAULT_TRACE_REGION_SIZE,
|
| 422 |
+
l1_small_size: int = DEFAULT_L1_SMALL_SIZE,
|
| 423 |
+
num_command_queues: int = 1,
|
| 424 |
+
device_id: int = 0,
|
| 425 |
+
) -> Any:
|
| 426 |
+
"""``tt.device.open_gr00t_device`` (program cache, p150a grid asserts) with the ``worker_l1_size`` the policy's
|
| 427 |
+
DiT backend needs: the firmware default for the ttnn backend, :func:`worker_l1_size_for` (the 64 KiB cut) for the
|
| 428 |
+
megakernel -- checked against :func:`device_worker_l1_size` after the open. ``num_command_queues=2`` is the
|
| 429 |
+
production path (``Gr00tTT.from_pretrained`` default; CQ-1 input writes overlap the traces). Callers hold the
|
| 430 |
+
device lock (``bin/with-device.sh``)."""
|
| 431 |
+
wl1 = worker_l1_size_for(policy)
|
| 432 |
+
device = open_gr00t_device(
|
| 433 |
+
trace_region_size=trace_region_size,
|
| 434 |
+
l1_small_size=l1_small_size,
|
| 435 |
+
num_command_queues=num_command_queues,
|
| 436 |
+
device_id=device_id,
|
| 437 |
+
worker_l1_size=wl1,
|
| 438 |
+
)
|
| 439 |
+
if wl1 is not None:
|
| 440 |
+
got = device_worker_l1_size(device)
|
| 441 |
+
if got != wl1:
|
| 442 |
+
close_gr00t_device(device)
|
| 443 |
+
raise RuntimeError(f"device opened with worker_l1_size {got}, requested {wl1} (megakernel backend)")
|
| 444 |
+
return device
|
| 445 |
+
|
| 446 |
+
|
| 447 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 448 |
+
# ttnn weights the megakernel backend does not need (startup cost, WP-K5 review NB 7)
|
| 449 |
+
# --------------------------------------------------------------------------------------------------------------------
|
| 450 |
+
|
| 451 |
+
|
| 452 |
+
#: Plan categories the Stage-1 head reads from ``TTWeights`` and the megakernel arena replaces
|
| 453 |
+
#: (``tt.megakernel.dit_program.real_arena_plan`` packs the same host plan tensors into the DRAM arena).
|
| 454 |
+
MK_ARENA_CATEGORIES: Tuple[str, ...] = ("dit", "dit_precompute", "encoders")
|
| 455 |
+
|
| 456 |
+
|
| 457 |
+
def megakernel_needs_ttnn_tensor(name: str, category: str) -> bool:
|
| 458 |
+
"""Whether plan tensor ``name`` of ``category`` must stay in ``TTWeights`` under the megakernel backend: every
|
| 459 |
+
tensor outside :data:`MK_ARENA_CATEGORIES`, plus the ones the *other* modules read from those categories --
|
| 460 |
+
``dit.block.{i}.Wkv`` / ``bkv`` (``VLAdapterTT`` hoists the cross K/V with them), ``enc.state.*`` (the adapter's
|
| 461 |
+
state encoder) and ``enc.action.b3_pos`` (``TTWeights.constants`` builds ``b3_pos_rows`` from the resident
|
| 462 |
+
tensor)."""
|
| 463 |
+
if category not in MK_ARENA_CATEGORIES:
|
| 464 |
+
return True
|
| 465 |
+
if name.startswith("dit.block.") and (name.endswith(".Wkv") or name.endswith(".bkv")):
|
| 466 |
+
return True
|
| 467 |
+
if name.startswith("enc.state."):
|
| 468 |
+
return True
|
| 469 |
+
return name == "enc.action.b3_pos"
|
| 470 |
+
|
| 471 |
+
|
| 472 |
+
def megakernel_weight_names(
|
| 473 |
+
version: str, embodiment_id: int, policy: TTPolicy, layout_name: Optional[str] = None
|
| 474 |
+
) -> Tuple[Tuple[str, ...], Dict[str, Any]]:
|
| 475 |
+
"""``(names to upload, facts)`` for ``TTWeights.load(names=...)`` under the megakernel backend: the full plan of
|
| 476 |
+
``(version, embodiment, policy)`` minus the tensors :func:`megakernel_needs_ttnn_tensor` rejects. ``facts``
|
| 477 |
+
records what is skipped (``n_tensors``, ``device_bytes``, per category) for ``Gr00tTT.describe``."""
|
| 478 |
+
from models.experimental.gr00t.tt.policy import plan_options
|
| 479 |
+
from models.experimental.gr00t.tt.weights import upload_plan
|
| 480 |
+
|
| 481 |
+
entries = upload_plan(version, int(embodiment_id), policy, plan_options(policy, version, layout_name))
|
| 482 |
+
keep: List[str] = []
|
| 483 |
+
skipped: Dict[str, Dict[str, float]] = {}
|
| 484 |
+
n_skip, bytes_skip = 0, 0.0
|
| 485 |
+
for name, e in entries.items():
|
| 486 |
+
if megakernel_needs_ttnn_tensor(name, e.category):
|
| 487 |
+
keep.append(name)
|
| 488 |
+
continue
|
| 489 |
+
c = skipped.setdefault(e.category, {"n_tensors": 0, "device_bytes": 0.0})
|
| 490 |
+
c["n_tensors"] += 1
|
| 491 |
+
c["device_bytes"] += float(e.nbytes)
|
| 492 |
+
n_skip += 1
|
| 493 |
+
bytes_skip += float(e.nbytes)
|
| 494 |
+
if not keep:
|
| 495 |
+
raise AssertionError("megakernel_weight_names selected no plan tensor")
|
| 496 |
+
facts = {
|
| 497 |
+
"n_tensors": n_skip,
|
| 498 |
+
"device_bytes": bytes_skip,
|
| 499 |
+
"categories": skipped,
|
| 500 |
+
"n_kept": len(keep),
|
| 501 |
+
"n_full": len(entries),
|
| 502 |
+
"kept_from_arena_categories": sorted(
|
| 503 |
+
n for n in keep if entries[n].category in MK_ARENA_CATEGORIES and not n.startswith("dit.block.")
|
| 504 |
+
),
|
| 505 |
+
}
|
| 506 |
+
return tuple(keep), facts
|
| 507 |
+
|
| 508 |
+
|
| 509 |
# --------------------------------------------------------------------------------------------------------------------
|
| 510 |
# bench / serve replayer contract (benchmarks/bench_e2e.py REPLAYER_MEMBERS)
|
| 511 |
# --------------------------------------------------------------------------------------------------------------------
|
|
|
|
| 603 |
backbone: Optional[BackboneTT] = None,
|
| 604 |
embed_gather: str = "host",
|
| 605 |
cq1_uploads: Optional[bool] = None,
|
| 606 |
+
adapter_mem: Optional[str] = None,
|
| 607 |
) -> None:
|
| 608 |
"""Assemble the modules on an open device (``from_pretrained`` is the public entry point).
|
| 609 |
|
|
|
|
| 613 |
``"host"`` (the WP-12 host gather; the constructor default, so stand-in weights keep working -- the production
|
| 614 |
factory :meth:`from_pretrained` defaults to :data:`DEFAULT_EMBED_GATHER`; module docstring).
|
| 615 |
``cq1_uploads``: per-call input writes on CQ 1 (``None`` = on iff the device has two command queues).
|
| 616 |
+
``adapter_mem``: ``"L1"`` | ``"DRAM"`` for the adapter's intermediates (``None`` = :data:`ADAPTER_MEM_BY_BACKEND`
|
| 617 |
+
of the policy's DiT backend; module docstring).
|
| 618 |
"""
|
| 619 |
if not isinstance(cfg, GR00TConfig) or not isinstance(dlayout, DeviceLayout):
|
| 620 |
raise TypeError("Gr00tTT(cfg: GR00TConfig, dlayout: DeviceLayout, policy: TTPolicy, W, device)")
|
|
|
|
| 635 |
f"weights are for {getattr(W, 'version', '?')}/slot {getattr(W, 'embodiment_id', '?')}, "
|
| 636 |
f"model is {cfg.version}/slot {e}"
|
| 637 |
)
|
| 638 |
+
wl1 = worker_l1_size_for(policy)
|
| 639 |
+
if wl1 is not None:
|
| 640 |
+
got = device_worker_l1_size(device)
|
| 641 |
+
if got > wl1:
|
| 642 |
+
raise RuntimeError(
|
| 643 |
+
f"dit_backend='megakernel' (the TTPolicy default) needs the device opened with worker_l1_size <= "
|
| 644 |
+
f"{wl1}; this device was opened with {got} (the K4 kernel binaries do not fit its kernel-config "
|
| 645 |
+
"ring buffer, mk_k4_summary.md section 7.2). Open it with tt.model.open_model_device(policy, ...) "
|
| 646 |
+
"or let Gr00tTT.from_pretrained(device=None) open it; for the Stage-1 path on a default device "
|
| 647 |
+
"use TTPolicy(dit_backend='ttnn') / GR00T_DIT_BACKEND=ttnn"
|
| 648 |
+
)
|
| 649 |
self.cfg, self.dlayout, self.policy, self.W, self.device = cfg, dlayout, policy, W, device
|
| 650 |
self.version: str = cfg.version
|
| 651 |
self.layout: LayoutConfig = dlayout.layout
|
|
|
|
| 658 |
self.embed_gather = embed_gather
|
| 659 |
self.specs: Dict[str, InputSpec] = model_input_specs(cfg, dlayout, policy, embed_gather)
|
| 660 |
self.timing: Dict[str, float] = {}
|
| 661 |
+
self.weights_skipped: Dict[str, Any] = {} # set by from_pretrained under the megakernel (module docstring)
|
| 662 |
self.n_predict = 0
|
| 663 |
self._b3: Optional[torch.Tensor] = None
|
| 664 |
self._hoist: Optional[HoistOut] = None
|
|
|
|
| 677 |
raise ValueError("the injected backbone must be built with collect_taps=False and hold no taps")
|
| 678 |
self.backbone = backbone
|
| 679 |
t1 = time.perf_counter()
|
| 680 |
+
self.adapter_mem: str = ADAPTER_MEM_BY_BACKEND[policy.dit_backend] if adapter_mem is None else str(adapter_mem)
|
| 681 |
+
if self.adapter_mem not in ("L1", "DRAM"):
|
| 682 |
+
raise ValueError(f"adapter_mem must be 'L1' or 'DRAM', got {adapter_mem!r}")
|
| 683 |
+
self.adapter = VLAdapterTT(
|
| 684 |
+
cfg,
|
| 685 |
+
dlayout,
|
| 686 |
+
W,
|
| 687 |
+
policy,
|
| 688 |
+
device,
|
| 689 |
+
gather_mode=gather_mode,
|
| 690 |
+
collect_taps=False,
|
| 691 |
+
mem=DRAM if self.adapter_mem == "DRAM" else L1,
|
| 692 |
+
)
|
| 693 |
t2 = time.perf_counter()
|
| 694 |
if embed_gather == "device":
|
| 695 |
if not hasattr(W, "_embed_table_tensor"):
|
|
|
|
| 813 |
embed_gather: str = DEFAULT_EMBED_GATHER,
|
| 814 |
cq1_uploads: Optional[bool] = None,
|
| 815 |
num_command_queues: int = 2,
|
| 816 |
+
adapter_mem: Optional[str] = None,
|
| 817 |
+
load_ttnn_dit_weights: Optional[bool] = None,
|
| 818 |
) -> "Gr00tTT":
|
| 819 |
"""Load the weights and assemble the model (plan §2.2 row ``tt/model.py``).
|
| 820 |
|
|
|
|
| 822 |
version: ``"n15" | "n16" | "n17"``.
|
| 823 |
embodiment: embodiment tag (default: the layout's ``embodiment_tag``); must be the slot the layout is baked for.
|
| 824 |
layout_name: one of ``cfg.layouts`` (default ``cfg.canonical_layout_name``, D2).
|
| 825 |
+
policy: ``TTPolicy`` (default ``TTPolicy()``: mixed_dit, **megakernel** backend with the bfp8_b arena, meta
|
| 826 |
+
RoPE, per_stage traces; ``TTPolicy().with_env_overrides()`` honours ``GR00T_DIT_BACKEND`` /
|
| 827 |
+
``GR00T_MK_ARENA_DTYPE`` -- the factory itself never reads the environment).
|
| 828 |
+
device: an open mesh device; ``None`` opens one with :func:`open_model_device` (``num_command_queues``;
|
| 829 |
+
the megakernel backend's ``worker_l1_size``), closed by ``release``. A device opened another way is
|
| 830 |
+
refused for the megakernel backend (constructor).
|
| 831 |
cache_root: ``HostPlanCache`` root for ``TTWeights.load`` (default: the project cache).
|
| 832 |
warm_inputs: a ``ModelInputs`` (or ``Observation``) to warm and capture on immediately; otherwise the
|
| 833 |
first ``predict`` / ``predict_normalized`` / ``tracer.upload`` call warms and captures on its inputs.
|
| 834 |
noise: the noise for that warm pass (default: the version's seeded noise).
|
| 835 |
+
embed_gather, cq1_uploads, adapter_mem: see the constructor (module docstring; ``adapter_mem="DRAM"`` is
|
| 836 |
+
what a ttnn-backend model needs on a device opened for the megakernel backend).
|
| 837 |
+
load_ttnn_dit_weights: ``None`` (default) uploads the Stage-1 DiT / precompute / encoder tensors only for
|
| 838 |
+
the ttnn backend (the megakernel packs them into its arena instead, :func:`megakernel_weight_names`;
|
| 839 |
+
``describe()["weights"]["skipped"]`` and :attr:`weights_skipped` record what was left out);
|
| 840 |
+
``True`` forces the full upload under either backend, ``False`` is an error for the ttnn backend.
|
| 841 |
"""
|
| 842 |
from models.experimental.gr00t.common.checkpoint import LazyCheckpoint
|
| 843 |
from models.experimental.gr00t.tt.weights import TTWeights
|
|
|
|
| 855 |
f"embodiment {tag!r} is slot {e} -- pick a layout of that embodiment"
|
| 856 |
)
|
| 857 |
dlayout = DeviceLayout.from_layout(cfg, layout)
|
| 858 |
+
skip_facts: Dict[str, Any] = {}
|
| 859 |
+
names: Optional[Tuple[str, ...]] = None
|
| 860 |
+
if load_ttnn_dit_weights is False and pol.dit_backend != "megakernel":
|
| 861 |
+
raise ValueError("load_ttnn_dit_weights=False needs dit_backend='megakernel' (the ttnn head reads them)")
|
| 862 |
+
if pol.dit_backend == "megakernel" and not load_ttnn_dit_weights:
|
| 863 |
+
names, skip_facts = megakernel_weight_names(version, e, pol, layout.name)
|
| 864 |
owns = device is None
|
| 865 |
+
dev = open_model_device(pol, num_command_queues=int(num_command_queues)) if owns else device
|
| 866 |
try:
|
| 867 |
t0 = time.perf_counter()
|
| 868 |
+
W = TTWeights.load(version, e, pol, dev, cache_root=cache_root, names=names, layout_name=layout.name)
|
| 869 |
load_s = time.perf_counter() - t0
|
| 870 |
ck = LazyCheckpoint(version)
|
| 871 |
model = cls(
|
|
|
|
| 879 |
owns_device=owns,
|
| 880 |
embed_gather=embed_gather,
|
| 881 |
cq1_uploads=cq1_uploads,
|
| 882 |
+
adapter_mem=adapter_mem,
|
| 883 |
)
|
| 884 |
except Exception:
|
| 885 |
if owns:
|
| 886 |
close_gr00t_device(dev)
|
| 887 |
raise
|
| 888 |
model.timing["weights_load_s"] = load_s
|
| 889 |
+
model.weights_skipped = dict(skip_facts)
|
| 890 |
if warm_inputs is not None:
|
| 891 |
mi = warm_inputs if not _is_observation(warm_inputs) else model.encode(warm_inputs)
|
| 892 |
model.warm_and_capture(mi, noise)
|
|
|
|
| 1215 |
else:
|
| 1216 |
taps[name] = slots_to_rows(L.to_torch(t), mi.slots)
|
| 1217 |
taps.update(self.head.golden_taps(b3=self.b3_bias(), block_taps=True))
|
| 1218 |
+
missing = [k for k in expected_tap_keys(cfg, dl, self.policy.dit_backend) if k not in taps]
|
| 1219 |
if missing:
|
| 1220 |
raise RuntimeError(f"run_untraced produced no tap for {missing}")
|
| 1221 |
return taps
|
|
|
|
| 1287 |
"policy": self.policy.to_dict(),
|
| 1288 |
"recipe": self.recipe.describe(),
|
| 1289 |
"embed_gather": self.embed_gather,
|
| 1290 |
+
"adapter_mem": self.adapter_mem,
|
| 1291 |
+
"worker_l1_size": _worker_l1_size_or_none(self.device),
|
| 1292 |
"gather": None if self.gather is None else self.gather.describe(),
|
| 1293 |
"input_specs": {n: s.describe() for n, s in self.specs.items()},
|
| 1294 |
"executor": self.executor.describe(),
|
|
|
|
| 1299 |
"device_bytes": float(self.W.device_bytes),
|
| 1300 |
"load_stats": dict(getattr(self.W, "load_stats", {})),
|
| 1301 |
"dtype_histogram": self.W.dtype_histogram() if hasattr(self.W, "dtype_histogram") else None,
|
| 1302 |
+
"skipped": dict(self.weights_skipped),
|
| 1303 |
},
|
| 1304 |
}
|
| 1305 |
for name in ("backbone", "adapter", "head"):
|
|
|
|
| 1314 |
# --------------------------------------------------------------------------------------------------------------------
|
| 1315 |
|
| 1316 |
|
| 1317 |
+
def _worker_l1_size_or_none(device: Any) -> Optional[int]:
|
| 1318 |
+
""":func:`device_worker_l1_size` for a real mesh device; ``None`` for the CPU tests' stand-in devices."""
|
| 1319 |
+
try:
|
| 1320 |
+
return device_worker_l1_size(device)
|
| 1321 |
+
except (ImportError, TypeError, AttributeError, RuntimeError):
|
| 1322 |
+
return None
|
| 1323 |
+
|
| 1324 |
+
|
| 1325 |
def _is_observation(x: Any) -> bool:
|
| 1326 |
return hasattr(x, "images") and hasattr(x, "instruction") and hasattr(x, "embodiment_tag")
|
| 1327 |
|
|
|
|
| 1360 |
"IN_VL_MASK",
|
| 1361 |
"IN_XT",
|
| 1362 |
"MASK_BUFFER_NAMES",
|
| 1363 |
+
"MK_ARENA_CATEGORIES",
|
| 1364 |
"NOISE_SEED",
|
| 1365 |
"N_STEPS",
|
| 1366 |
"OUT_ACTION_PRED",
|
|
|
|
| 1372 |
"Gr00tTT",
|
| 1373 |
"ModelRecipe",
|
| 1374 |
"TraceReplayer",
|
| 1375 |
+
"ADAPTER_MEM_BY_BACKEND",
|
| 1376 |
+
"MK_L1_CUT_BYTES_DEFAULT",
|
| 1377 |
+
"MK_L1_CUT_ENV",
|
| 1378 |
+
"device_worker_l1_size",
|
| 1379 |
"expected_tap_keys",
|
| 1380 |
+
"megakernel_needs_ttnn_tensor",
|
| 1381 |
+
"megakernel_weight_names",
|
| 1382 |
+
"mk_l1_cut_bytes",
|
| 1383 |
"model_input_specs",
|
| 1384 |
+
"open_model_device",
|
| 1385 |
"text_ids_spec",
|
| 1386 |
"unproduced_taps",
|
| 1387 |
+
"worker_l1_size_for",
|
| 1388 |
]
|
code/models/experimental/gr00t/tt/policy.py
CHANGED
|
@@ -42,6 +42,29 @@ from typing import Any, Dict, Mapping, Optional, Tuple
|
|
| 42 |
|
| 43 |
DTYPE_POLICIES: Tuple[str, ...] = ("bf16", "mixed_dit", "mixed")
|
| 44 |
DIT_BACKENDS: Tuple[str, ...] = ("ttnn", "megakernel")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 45 |
READER_MODES: Tuple[str, ...] = ("relay", "direct")
|
| 46 |
ROPE_LAYOUTS: Tuple[str, ...] = ("meta", "hf")
|
| 47 |
TRACE_LAYOUTS: Tuple[str, ...] = ("per_stage", "two")
|
|
@@ -109,8 +132,19 @@ class TTPolicy:
|
|
| 109 |
"""Run-time policy of the device port (plan §2.2 row ``tt/policy.py``).
|
| 110 |
|
| 111 |
* ``dtype_policy``: weight dtype class table, see the module docstring.
|
| 112 |
-
* ``dit_backend``: ``"
|
| 113 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 114 |
* ``rope_layout``: ``"meta"`` (``rotary_embedding_llama`` on the permuted layout, D8 default) | ``"hf"``
|
| 115 |
(``rotary_embedding_hf`` on the HF layout). Byte-changing: part of ``TransformOptions``.
|
| 116 |
* ``trace_layout``: ``"per_stage"`` (vision / llm / adapter / denoise traces, bring-up) | ``"two"`` (D1 end state).
|
|
@@ -134,7 +168,8 @@ class TTPolicy:
|
|
| 134 |
"""
|
| 135 |
|
| 136 |
dtype_policy: str = "mixed_dit"
|
| 137 |
-
dit_backend: str =
|
|
|
|
| 138 |
reader_mode: str = "relay"
|
| 139 |
rope_layout: str = "meta"
|
| 140 |
trace_layout: str = "per_stage"
|
|
@@ -149,6 +184,7 @@ class TTPolicy:
|
|
| 149 |
def __post_init__(self) -> None:
|
| 150 |
_check("dtype_policy", self.dtype_policy, DTYPE_POLICIES)
|
| 151 |
_check("dit_backend", self.dit_backend, DIT_BACKENDS)
|
|
|
|
| 152 |
_check("reader_mode", self.reader_mode, READER_MODES)
|
| 153 |
_check("rope_layout", self.rope_layout, ROPE_LAYOUTS)
|
| 154 |
_check("trace_layout", self.trace_layout, TRACE_LAYOUTS)
|
|
@@ -184,6 +220,31 @@ class TTPolicy:
|
|
| 184 |
raise ValueError(f"unknown vision family {vision_family!r}; expected one of {VISION_FAMILIES}")
|
| 185 |
return bool(self.vit_fork) and vision_family == VIT_FORK_FAMILY
|
| 186 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 187 |
def admitted(self) -> Dict[str, str]:
|
| 188 |
"""The category -> dtype table in force for this policy: ``admitted_classes`` when given, else the module
|
| 189 |
table :data:`ADMITTED_CLASSES` (read at call time so a table change is seen by every later ``dtype_of``)."""
|
|
@@ -211,6 +272,7 @@ class TTPolicy:
|
|
| 211 |
return {
|
| 212 |
"dtype_policy": self.dtype_policy,
|
| 213 |
"dit_backend": self.dit_backend,
|
|
|
|
| 214 |
"reader_mode": self.reader_mode,
|
| 215 |
"rope_layout": self.rope_layout,
|
| 216 |
"trace_layout": self.trace_layout,
|
|
@@ -358,6 +420,11 @@ def dtype_histogram(decisions: Mapping[str, str]) -> Dict[str, int]:
|
|
| 358 |
__all__ = [
|
| 359 |
"DTYPE_POLICIES",
|
| 360 |
"DIT_BACKENDS",
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 361 |
"READER_MODES",
|
| 362 |
"ROPE_LAYOUTS",
|
| 363 |
"TRACE_LAYOUTS",
|
|
|
|
| 42 |
|
| 43 |
DTYPE_POLICIES: Tuple[str, ...] = ("bf16", "mixed_dit", "mixed")
|
| 44 |
DIT_BACKENDS: Tuple[str, ...] = ("ttnn", "megakernel")
|
| 45 |
+
#: ``TTPolicy.dit_backend`` default. **``"megakernel"`` since 2026-09-18** (plan-owner ratification in
|
| 46 |
+
#: ``docs/plan/reviews/WP-K5-review-round1.md`` "Plan-owner decisions" item 1, overriding plan §5.7's "G3 fails ->
|
| 47 |
+
#: Stage 1 stays production"): the whole DiT denoise runs as one ``steps4`` ``generic_op`` inside trace B
|
| 48 |
+
#: (``tests/tt/results/mk_k5_summary.md``: gate outcomes equal to the ttnn backend's on 21 / 23 (version, sample)
|
| 49 |
+
#: pairs and every canonical sample, mk vs ttnn ``action_pred`` >= 0.999993, e2e -13 / -23 / -17 %). The ttnn
|
| 50 |
+
#: backend stays fully selectable (``TTPolicy(dit_backend="ttnn")`` or ``GR00T_DIT_BACKEND=ttnn``).
|
| 51 |
+
DEFAULT_DIT_BACKEND: str = "megakernel"
|
| 52 |
+
#: ``TTPolicy.mk_arena_dtype``: the DRAM weight arena of the megakernel backend (plan §5.3 ``dtype_stream``);
|
| 53 |
+
#: ``"auto"`` resolves per version through :data:`MK_ARENA_DTYPE_DEFAULT` (:meth:`TTPolicy.mk_arena_dtype_for`).
|
| 54 |
+
MK_ARENA_DTYPES: Tuple[str, ...] = ("auto", "bf16", "bfp8_b")
|
| 55 |
+
#: Per-version arena dtype behind ``mk_arena_dtype="auto"``: the fastest arena whose e2e gate outcomes equal the
|
| 56 |
+
#: bf16 arena's on the canonical and every multi-sample golden (WP-K5, ``tests/tt/results/mk_k5_summary.md`` §3 / §4.1:
|
| 57 |
+
#: bf16 and bfp8_b have identical pass / fail sets on all 23 (version, sample) pairs and agree to <= 8e-4 PCC on
|
| 58 |
+
#: ``action_pred_normalized`` (<= 1e-6 on N1.6); bfp8_b denoise trace 10.41 / 17.16 / 18.80 ms vs bf16 11.36 / 18.92 /
|
| 59 |
+
#: 20.66 ms, §5). Stage 1 already runs the DiT projections in bfp8 (``ADMITTED_CLASSES["dit"]``), so this is the same
|
| 60 |
+
#: weight dtype.
|
| 61 |
+
MK_ARENA_DTYPE_DEFAULT: Dict[str, str] = {"n15": "bfp8_b", "n16": "bfp8_b", "n17": "bfp8_b"}
|
| 62 |
+
#: Environment overrides of :meth:`TTPolicy.with_env_overrides` (WP-K5; the device test files, the published policy
|
| 63 |
+
#: server and ``Gr00tTT`` callers that want an A/B against the Stage-1 path use them):
|
| 64 |
+
#: ``GR00T_DIT_BACKEND=ttnn|megakernel`` and ``GR00T_MK_ARENA_DTYPE=auto|bf16|bfp8_b``. Unset = the code defaults
|
| 65 |
+
#: (:data:`DEFAULT_DIT_BACKEND`, ``"auto"``).
|
| 66 |
+
DIT_BACKEND_ENV = "GR00T_DIT_BACKEND"
|
| 67 |
+
MK_ARENA_DTYPE_ENV = "GR00T_MK_ARENA_DTYPE"
|
| 68 |
READER_MODES: Tuple[str, ...] = ("relay", "direct")
|
| 69 |
ROPE_LAYOUTS: Tuple[str, ...] = ("meta", "hf")
|
| 70 |
TRACE_LAYOUTS: Tuple[str, ...] = ("per_stage", "two")
|
|
|
|
| 132 |
"""Run-time policy of the device port (plan §2.2 row ``tt/policy.py``).
|
| 133 |
|
| 134 |
* ``dtype_policy``: weight dtype class table, see the module docstring.
|
| 135 |
+
* ``dit_backend``: ``"megakernel"`` (**default** since 2026-09-18, :data:`DEFAULT_DIT_BACKEND`: Stage 2, the whole
|
| 136 |
+
denoise as one ``steps4`` ``generic_op`` inside trace B, ``tt.action_head.ActionHeadTT`` +
|
| 137 |
+
``tt.megakernel.dit_program``; D1/D10, WP-K5) | ``"ttnn"`` (Stage 1: trace B of ttnn ops, the previous default;
|
| 138 |
+
selectable for A/B). The megakernel needs the device opened with a ``worker_l1_size`` reduced by
|
| 139 |
+
``tt.model.MK_L1_CUT_BYTES_DEFAULT`` (``tt.model.open_model_device``; ``Gr00tTT.from_pretrained(device=None)``
|
| 140 |
+
does it) and the VL adapter's intermediates in DRAM; ``Gr00tTT`` refuses a device opened otherwise.
|
| 141 |
+
* ``mk_arena_dtype``: the megakernel's DRAM weight arena, ``"bf16"`` | ``"bfp8_b"`` | ``"auto"`` (default: the
|
| 142 |
+
per-version table :data:`MK_ARENA_DTYPE_DEFAULT`, :meth:`mk_arena_dtype_for`). Independent of ``dtype_policy``
|
| 143 |
+
(which governs the ttnn weights ``TTWeights`` uploads); not part of the weights cache key (the arena is packed
|
| 144 |
+
from the bf16 host plan at model build). Ignored by the ttnn backend.
|
| 145 |
+
* ``reader_mode``: K1 megakernel weight-stream topology, ``"relay"`` ring (default, plan §5.3) | ``"direct"``;
|
| 146 |
+
the K4 / K5 ``steps4`` program always streams in direct mode (``mk_k4_summary.md`` §8), so this knob only
|
| 147 |
+
reaches the K1 bench / unit-test launches.
|
| 148 |
* ``rope_layout``: ``"meta"`` (``rotary_embedding_llama`` on the permuted layout, D8 default) | ``"hf"``
|
| 149 |
(``rotary_embedding_hf`` on the HF layout). Byte-changing: part of ``TransformOptions``.
|
| 150 |
* ``trace_layout``: ``"per_stage"`` (vision / llm / adapter / denoise traces, bring-up) | ``"two"`` (D1 end state).
|
|
|
|
| 168 |
"""
|
| 169 |
|
| 170 |
dtype_policy: str = "mixed_dit"
|
| 171 |
+
dit_backend: str = DEFAULT_DIT_BACKEND
|
| 172 |
+
mk_arena_dtype: str = "auto"
|
| 173 |
reader_mode: str = "relay"
|
| 174 |
rope_layout: str = "meta"
|
| 175 |
trace_layout: str = "per_stage"
|
|
|
|
| 184 |
def __post_init__(self) -> None:
|
| 185 |
_check("dtype_policy", self.dtype_policy, DTYPE_POLICIES)
|
| 186 |
_check("dit_backend", self.dit_backend, DIT_BACKENDS)
|
| 187 |
+
_check("mk_arena_dtype", self.mk_arena_dtype, MK_ARENA_DTYPES)
|
| 188 |
_check("reader_mode", self.reader_mode, READER_MODES)
|
| 189 |
_check("rope_layout", self.rope_layout, ROPE_LAYOUTS)
|
| 190 |
_check("trace_layout", self.trace_layout, TRACE_LAYOUTS)
|
|
|
|
| 220 |
raise ValueError(f"unknown vision family {vision_family!r}; expected one of {VISION_FAMILIES}")
|
| 221 |
return bool(self.vit_fork) and vision_family == VIT_FORK_FAMILY
|
| 222 |
|
| 223 |
+
def mk_arena_dtype_for(self, version: str) -> str:
|
| 224 |
+
"""The megakernel arena dtype of ``version`` (``"bf16"`` | ``"bfp8_b"``): the explicit field, or the
|
| 225 |
+
per-version default :data:`MK_ARENA_DTYPE_DEFAULT` behind ``"auto"``. Unknown versions raise."""
|
| 226 |
+
if self.mk_arena_dtype != "auto":
|
| 227 |
+
return self.mk_arena_dtype
|
| 228 |
+
try:
|
| 229 |
+
return MK_ARENA_DTYPE_DEFAULT[version]
|
| 230 |
+
except KeyError:
|
| 231 |
+
raise KeyError(f"no default megakernel arena dtype for version {version!r}") from None
|
| 232 |
+
|
| 233 |
+
def with_env_overrides(self, env: Optional[Mapping[str, str]] = None) -> "TTPolicy":
|
| 234 |
+
"""Copy with ``dit_backend`` / ``mk_arena_dtype`` taken from :data:`DIT_BACKEND_ENV` /
|
| 235 |
+
:data:`MK_ARENA_DTYPE_ENV` when set (the device test files and the published policy server parametrise the
|
| 236 |
+
DiT backend this way, WP-K5: ``GR00T_DIT_BACKEND=ttnn pytest ...`` runs the Stage-1 path, unset = the
|
| 237 |
+
megakernel default). Unset variables change nothing; an unknown value raises through ``__post_init__``."""
|
| 238 |
+
import os
|
| 239 |
+
|
| 240 |
+
src = os.environ if env is None else env
|
| 241 |
+
changes: Dict[str, str] = {}
|
| 242 |
+
if src.get(DIT_BACKEND_ENV):
|
| 243 |
+
changes["dit_backend"] = str(src[DIT_BACKEND_ENV])
|
| 244 |
+
if src.get(MK_ARENA_DTYPE_ENV):
|
| 245 |
+
changes["mk_arena_dtype"] = str(src[MK_ARENA_DTYPE_ENV])
|
| 246 |
+
return replace(self, **changes) if changes else self
|
| 247 |
+
|
| 248 |
def admitted(self) -> Dict[str, str]:
|
| 249 |
"""The category -> dtype table in force for this policy: ``admitted_classes`` when given, else the module
|
| 250 |
table :data:`ADMITTED_CLASSES` (read at call time so a table change is seen by every later ``dtype_of``)."""
|
|
|
|
| 272 |
return {
|
| 273 |
"dtype_policy": self.dtype_policy,
|
| 274 |
"dit_backend": self.dit_backend,
|
| 275 |
+
"mk_arena_dtype": self.mk_arena_dtype,
|
| 276 |
"reader_mode": self.reader_mode,
|
| 277 |
"rope_layout": self.rope_layout,
|
| 278 |
"trace_layout": self.trace_layout,
|
|
|
|
| 420 |
__all__ = [
|
| 421 |
"DTYPE_POLICIES",
|
| 422 |
"DIT_BACKENDS",
|
| 423 |
+
"DEFAULT_DIT_BACKEND",
|
| 424 |
+
"MK_ARENA_DTYPES",
|
| 425 |
+
"MK_ARENA_DTYPE_DEFAULT",
|
| 426 |
+
"DIT_BACKEND_ENV",
|
| 427 |
+
"MK_ARENA_DTYPE_ENV",
|
| 428 |
"READER_MODES",
|
| 429 |
"ROPE_LAYOUTS",
|
| 430 |
"TRACE_LAYOUTS",
|