changh95 commited on
Commit
0c2d7e9
·
verified ·
1 Parent(s): 78d7612

Add files using upload-large-folder tool

Browse files
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
- setup = E2E.resolve_setup(version, policy=policy, layout_name=layout_name, embodiment=embodiment)
 
 
 
 
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, policy=policy, trace_layout=trace_layout, layout_name=layout_name, embodiment=embodiment
 
 
 
 
 
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(dtype_policy=policy, trace_layout=trace_layout) # raises on unknown values
 
 
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, policy=policy, trace_layout=trace_layout, layout_name=layout_name, embodiment=embodiment
 
 
 
 
 
 
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(setup, device, cache_root, embed_gather=embed_gather, cq1_uploads=cq1_uploads)
 
 
 
 
 
 
 
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, policy=policy, trace_layout=trace_layout, layout_name=layout_name, embodiment=embodiment
 
 
 
 
 
 
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} | stages {st} | embed_gather {g} | "
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
- device = D.open_gr00t_device(
 
 
 
 
 
 
 
 
 
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 = Gr00tTT.from_pretrained(version, policy=TTPolicy(), device=device)
 
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,) * 5
 
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
- assert len(spec.semaphores) == 9 and sorted(s.sem_id for s in spec.semaphores) == list(range(9))
 
 
 
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
- one ``generic_op`` for the whole denoise, plan §5) raises ``NotImplementedError`` from :meth:`denoise` until WP-K5
28
- wires ``DiTMegakernel`` -- construction succeeds so that ``Gr00tTT`` can be assembled under either policy.
29
-
30
- Taps (``collect_taps=True``) are device tensors under the golden keys (``sa_embs[k=0]``, ``dit_block{i}_out[k=0]``,
31
- ``dit_block_last_out[k=0]``, ``dit_out[k=0]``, ``action_decoder_out[k=0]``, ``actions_out[k=0]``,
32
- ``action_pred_normalized``; plus ``action_encoder_pre_pos[k=0]``); :func:`taps_to_golden` slices them to the goldens'
33
- shapes for ``tests.tt.harness.check_taps``. ``ttnn`` is imported lazily (constructor and forwards).
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- device: the open mesh device (buffers are allocated here).
241
- encoders: a prebuilt :class:`HeadEncoders` (default: built from ``W``).
 
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
- encoders if encoders is not None else HeadEncoders.build(cfg, W, dlayout, mem=mem, fuse_swish=fuse_swish)
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 # WP-K5 builds DiTMegakernel here for backend == "megakernel"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- return self.step.bind_kv(kv)
 
 
 
 
 
 
 
 
 
 
 
 
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.set_collect_taps(flag)
382
- self.encoders.set_collect_taps(flag)
 
 
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 denoise(self, x_t: Any, state_features: Any, kv: Optional[Mapping[int, Any]] = None) -> Any:
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
- if kv is not None and kv is not self.step.kv:
424
- self.step.bind_kv(kv)
425
- if self.step.kv is None:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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 = self.step.forward(sa, k)
450
- for name, t in self.step.debug_taps.items():
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
- for t in list(self.taps.values()) + list(self.step.debug_taps.values()) + self._tap_pool:
489
- if t is None or id(t) in seen or id(t) in keep:
 
490
  continue
491
  seen.add(id(t))
492
  ttnn.deallocate(t)
493
  self.taps = {}
494
  self._tap_pool = []
495
- self.step.debug_taps = {}
496
- self.encoders.set_collect_taps(self.collect_taps) # clears the modules' debug_taps dicts
 
 
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. Callers hold the device lock
151
- (``bin/with-device.sh``).
 
 
 
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 DRAM, close_gr00t_device, open_gr00t_device, resolve_mem
 
 
 
 
 
 
 
 
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
- r = ModelRecipe.from_config(cfg, dlayout, TTPolicy())
 
 
 
 
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
- for k in range(r.n_steps):
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.adapter = VLAdapterTT(cfg, dlayout, W, policy, device, gather_mode=gather_mode, collect_taps=False)
 
 
 
 
 
 
 
 
 
 
 
 
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, ttnn backend, meta RoPE, per_stage traces).
632
- device: an open mesh device; ``None`` opens one with ``tt.device.open_gr00t_device(num_command_queues=
633
- num_command_queues)`` (closed by ``release``).
 
 
 
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 = open_gr00t_device(num_command_queues=int(num_command_queues)) if owns else device
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``: ``"ttnn"`` (Stage 1 trace B of ttnn ops) | ``"megakernel"`` (Stage 2 ``generic_op``; D1/D10).
113
- * ``reader_mode``: megakernel weight-stream topology, ``"relay"`` ring (default, plan §5.3) | ``"direct"``.
 
 
 
 
 
 
 
 
 
 
 
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 = "ttnn"
 
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",