File size: 70,104 Bytes
4d9b003
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
973
974
975
976
977
978
979
980
981
982
983
984
985
986
987
988
989
990
991
992
993
994
995
996
997
998
999
1000
1001
1002
1003
1004
1005
1006
1007
1008
1009
1010
1011
1012
1013
1014
1015
1016
1017
1018
1019
1020
1021
1022
1023
1024
1025
1026
1027
1028
1029
1030
1031
1032
1033
1034
1035
1036
1037
1038
1039
1040
1041
1042
1043
1044
1045
1046
1047
1048
1049
1050
1051
1052
1053
1054
1055
1056
1057
1058
1059
1060
1061
1062
1063
1064
1065
1066
1067
1068
1069
1070
1071
1072
1073
1074
1075
1076
1077
1078
1079
1080
1081
1082
1083
1084
1085
1086
1087
1088
1089
1090
1091
1092
1093
1094
1095
1096
1097
1098
1099
1100
1101
1102
1103
1104
1105
1106
1107
1108
1109
1110
1111
1112
1113
1114
1115
1116
1117
1118
1119
1120
1121
1122
1123
1124
1125
1126
1127
1128
1129
1130
1131
1132
1133
1134
1135
1136
1137
1138
1139
1140
1141
1142
1143
1144
1145
1146
1147
1148
1149
1150
1151
1152
1153
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180
1181
1182
1183
1184
1185
1186
1187
1188
1189
1190
1191
1192
1193
1194
1195
1196
1197
1198
1199
1200
1201
1202
1203
1204
1205
1206
1207
1208
1209
1210
1211
1212
1213
1214
1215
1216
1217
1218
1219
1220
1221
1222
1223
1224
1225
1226
1227
1228
1229
1230
1231
1232
1233
1234
1235
1236
1237
1238
1239
1240
1241
1242
1243
1244
1245
1246
1247
1248
1249
1250
1251
1252
1253
1254
1255
1256
1257
1258
1259
1260
1261
1262
1263
1264
1265
1266
1267
1268
1269
1270
1271
1272
1273
1274
1275
1276
1277
1278
1279
1280
1281
1282
1283
1284
1285
1286
1287
1288
1289
1290
1291
1292
1293
1294
1295
1296
1297
1298
1299
1300
1301
1302
1303
1304
1305
1306
1307
1308
1309
1310
1311
1312
1313
1314
1315
1316
1317
1318
# SPDX-License-Identifier: Apache-2.0
"""C02: ``TraceRunner`` -- every device stage of a port runs inside metal traces.

What it enforces (TT_PLATFORM.md section 3, REFERENCE_PATTERNS.md section 1.4):

1. **Persistent I/O before any capture.** Inputs, RT-dev parameters, state buffers and (optionally) output buffers
   are allocated (and given defined initial values: never uninitialised index buffers) when they are added, i.e.
   before the first capture. A variant returns either the tensors its last ops produce (allocated during capture,
   valid while the trace lives) or persistent outputs written with ``ctx.write_output`` (stable addresses shared by
   several variants, never overwritten by another variant's replay).
2. **Warm-up, then capture.** Every variant runs eagerly ``warmup_runs`` times (kernel JIT, program cache, prepared
   conv weights) before *any* trace is captured. Adding a variant after a capture releases all traces, warms the new
   one and recaptures everything, because a warm-up after a capture could allocate into a trace's freed
   intermediates.
3. **Strict capture.** ``device.set_program_cache_misses_allowed(False)`` during capture (a miss raises with the op
   name instead of aborting with "Writes are not supported during trace capture"); ``end_trace_capture`` runs in a
   ``finally`` block and a failed capture is released (an open capture left the process spinning in close_device:
   RP section 1.4); the program-cache entry count must not change during capture.
4. **Variants.** Several traces keyed by name (shape buckets, modes, segments), chosen by the caller per run.
5. **RT-dev parameters.** Per-frame values that must not change the program (thresholds, timesteps, poses) live
   in persistent device tensors refreshed before ``execute_trace``; unchanged values are not re-uploaded.
6. **1CQ / 2CQ.** With 2 CQs inputs are uploaded on CQ1 and ordered with events, exactly like
   ``common/tools/check_dispatch.py`` (CQ1 waits for the last trace, uploads, records; CQ0 waits, replays, records).
   ``stage_inputs=True`` uses the ``tt_cnn`` executor pattern instead: CQ1 writes a DRAM staging copy while the
   previous trace still runs, and an eager copy on CQ0 moves it into the trace input before the replay. Every CQ0
   write into a trace input outside the protocol (:meth:`TraceRunner.write_input`) re-records the event CQ1 waits for.
7. **State inside the trace.** ``ctx.write_state(name, value)`` is ``ttnn.copy(value, buffer)`` into a persistent
   buffer (FLOAT32 / UINT32 / BFLOAT16 / ... all supported), or no op at all when ``value`` was computed straight
   into ``ctx.write_target(name)`` (``output_tensor=``). ``pingpong=True`` states use two buffers and two traces
   per variant (phase 0 reads A writes B, phase 1 reads B writes A); the phase flips after every run of a variant
   that writes the ping-pong states. Per-stream state (PLAN.md D16): ``add_state(..., banks=n)`` +
   :meth:`TraceRunner.save_state` / :meth:`TraceRunner.load_state`, with the policy in :class:`StreamBanks`.
8. **Readback.** One packed output (:func:`pack_outputs`) gives one D2H; reads go into preallocated host tensors,
   on CQ0 or (segmented pipelines, 2CQ) on CQ1 after a host-side event wait.
9. **Alloc tracking.** With ``TT_METAL_TRACE_ALLOC_TRACKING=1`` set before ``import ttnn``, ``ttnn.execute_trace``
   refuses to replay over live unsafe buffers; the runner acknowledges the outputs of traces captured after the
   first one (they may be overwritten by an earlier trace's replay: read outputs before running another variant).

Example::

    runner = TraceRunner(device, num_command_queues=2)
    runner.add_input("x", shape=(1, 1, 64, 64), dtype="bfloat16", layout=ttnn.TILE_LAYOUT)
    runner.add_param("scale", 1.0)                                  # fp32 [1,1,1,1], TILE
    runner.add_variant("default", lambda ctx: ttnn.relu(ttnn.multiply(ctx["x"], ctx["scale"])))
    runner.capture()                                                # warm-up, then capture
    y = runner("default", inputs={"x": x_np}, params={"scale": 0.5})  # upload, replay, read -> numpy
"""
from __future__ import annotations

import contextlib
import math
import time
from dataclasses import dataclass, field
from typing import Any, Callable, Dict, Iterator, List, Mapping, Optional, Sequence, Tuple

import numpy as np

from .io import InputError
from .tensors import TILE, dtype_name, round_up, to_host_tensor, to_numpy, ttnn_dtype

__all__ = [
    "CQ_COMPUTE",
    "CQ_INPUT",
    "PackEntry",
    "PackLayout",
    "Packed",
    "pack_outputs",
    "SINGLE_ROW_MAX_ELEMS",
    "PACK_ROW_ELEMS",
    "TraceContext",
    "TraceRunner",
    "StreamBanks",
    "alloc_tracking_enabled",
]

CQ_COMPUTE = 0   # programs, traces and (by default) readback
CQ_INPUT = 1     # host -> device uploads when the device has 2 command queues


def alloc_tracking_enabled() -> bool:
    """True when ``TT_METAL_TRACE_ALLOC_TRACKING=1`` was set before ttnn was imported (read from tt-metal)."""
    try:
        from ttnn.tools import trace_allocation_tracker as tracker
    except ImportError:
        return False
    return bool(getattr(tracker, "TRACE_ALLOC_TRACKING", False))


# ----------------------------------------------------------------------------------------------- packing

# Packed layouts (PORT_LOG Q9 of the YOLOX port). A ROW_MAJOR tensor is stored page by page, one page per row, and
# the RM reshape / concat programs stage whole pages in L1 (reshape_rm_program_factory.cpp: 2 x the destination page
# per kernel copy when the pages are not 16-byte aligned), so a single-row pack of more than ~0.6 MB fails with
# "RM reshape dest staging does not fit in L1". Above SINGLE_ROW_MAX_ELEMS the pack is a [1, 1, rows, R] tensor of
# R = PACK_ROW_ELEMS elements per row: every RM page the packing programs touch stays below ROW_PAGE_MAX_BYTES,
# whatever the size of the outputs (tens of MB), and the readback is still ONE device-to-host copy.
SINGLE_ROW_MAX_ELEMS = 131072    # one [1, 1, 1, total] row up to 512 KiB of float32 (the 0.1.0 - 0.14.0 layout)
PACK_ROW_ELEMS = 1024            # elements per row of the multi-row layout (4 KiB float32 pages)
FLAT_MAX_BYTES = 64 << 10        # multi-row: tensors up to this size are flattened to one row, then cut into rows
ROW_PAGE_MAX_BYTES = 128 << 10   # multi-row: largest RM page (last dim x element size) a packed tensor may have
_ELEM_BYTES = {"float32": 4, "uint32": 4, "int32": 4, "bfloat16": 2, "uint16": 2, "uint8": 1}


@dataclass(frozen=True)
class PackEntry:
    """One tensor inside a packed readback: elements ``[offset, offset + numel)`` reshaped to ``shape``.

    ``pitch > 0`` (multi-row layout only): the tensor's rows (its last dim, ``shape[-1]`` elements)
    are stored ``pitch`` elements apart, zero-padded, i.e. elements ``[offset, offset + numel // shape[-1] * pitch)``
    viewed as ``(rows, pitch)`` hold it in their first ``shape[-1]`` columns. ``0``: contiguous."""

    name: str
    offset: int
    numel: int
    shape: Tuple[int, ...]
    pitch: int = 0

    def view(self, flat: np.ndarray) -> np.ndarray:
        """This tensor inside the flat readback ``flat`` (a view, strided when ``pitch`` is set)."""
        if not self.pitch:
            return flat[self.offset:self.offset + self.numel].reshape(self.shape)
        cols = int(self.shape[-1]) if self.shape else 1
        rows = self.numel // cols
        return flat[self.offset:self.offset + rows * self.pitch].reshape(rows, self.pitch)[:, :cols].reshape(
            self.shape)


@dataclass(frozen=True)
class PackLayout:
    """Host-side description of a packed output (built at capture time, applied at every read).

    ``entries`` index the packed tensor's elements in row-major order, so :meth:`unpack` does not depend on the
    device shape ``(1, 1, rows, row_elems)``: one row of ``total`` elements (``rows == 1``) or rows of
    ``row_elems`` elements, every tensor starting on a row boundary (multi-row layout; an entry with a ``pitch``
    stores its rows zero-padded to that many elements, see :class:`PackEntry`)."""

    entries: Tuple[PackEntry, ...]
    total: int
    rows: int = 1
    row_elems: int = 0           # 0: one row of ``total`` elements

    @property
    def shape(self) -> Tuple[int, int, int, int]:
        """Shape of the packed device tensor."""
        return (1, 1, self.rows, self.row_elems or self.total)

    def unpack(self, flat: Any) -> Dict[str, np.ndarray]:
        """Flat readback (any shape with ``total`` elements) -> ``{name: array}`` (views, no copies; the view of an
        entry with a ``pitch`` is strided, not C-contiguous)."""
        a = np.asarray(flat).reshape(-1)
        if a.size != self.total:
            raise ValueError(f"packed readback has {a.size} elements, layout expects {self.total}")
        return {e.name: e.view(a) for e in self.entries}


@dataclass
class Packed:
    """A packed device tensor plus its layout; return it from a variant function to get ``{name: array}`` back."""

    tensor: Any
    layout: PackLayout


def _elem_bytes(dtype: Any) -> int:
    return _ELEM_BYTES.get(dtype_name(dtype), 4)


def _row_major(ttnn: Any, t: Any, dtype: Any) -> Any:
    if t.dtype != dtype:
        t = ttnn.typecast(t, dtype)
    if t.layout != ttnn.ROW_MAJOR_LAYOUT:
        t = ttnn.to_layout(t, ttnn.ROW_MAJOR_LAYOUT)
    return t


def _flat_rows(ttnn: Any, t: Any, shape: Tuple[int, ...], numel: int, row_elems: int) -> Tuple[Any, int]:
    """A small RM tensor -> one row -> zero-padded to whole rows -> ``[1, 1, k, row_elems]``."""
    padded = round_up(numel, row_elems)
    if shape != (1, 1, 1, numel):
        t = ttnn.reshape(t, (1, 1, 1, numel))
    if padded != numel:
        t = ttnn.pad(t, [(0, 0), (0, 0), (0, 0), (0, padded - numel)], 0.0)
    if padded != row_elems:
        t = ttnn.reshape(t, (1, 1, padded // row_elems, row_elems))
    return t, padded


def _fill_rows(rows: int, cols: int, row_elems: int) -> int:
    """Rows of ``cols`` elements a ``[rows, cols]`` tensor is zero-padded to so that it fills whole packed rows."""
    return round_up(rows, row_elems // math.gcd(cols, row_elems))


def _row_pitches(cols: int, row_elems: int, elem: int) -> List[int]:
    """Row pitches > ``cols`` worth trying: ``cols`` rounded up to every power of two dividing ``row_elems`` (rows of
    that pitch fill a packed row every ``row_elems / gcd`` rows) and to ``row_elems`` (every row fills whole packed
    rows), within the page budget."""
    steps, m = {row_elems}, 2
    while row_elems % m == 0:
        steps.add(m)
        m *= 2
    return sorted({p for p in (round_up(cols, m) for m in steps) if p != cols and p * elem <= ROW_PAGE_MAX_BYTES})


def _pack_rows_segment(ttnn: Any, name: str, t: Any, shape: Tuple[int, ...], numel: int, row_elems: int,
                       elem: int) -> Tuple[Any, int, int]:
    """One ROW_MAJOR tensor -> ``[1, 1, k, row_elems]`` holding its elements in order, zero-padded to whole rows.
    Returns ``(segment, k * row_elems, pitch)`` (``pitch``: see :class:`PackEntry`). No RM page above
    ``ROW_PAGE_MAX_BYTES`` is created or reshaped."""
    cols = int(shape[-1]) if shape else 1
    rows = numel // cols
    if cols * elem <= ROW_PAGE_MAX_BYTES:
        # [rows, cols] is a view of the RM tensor; rows * cols fills whole packed rows when rows % unit == 0
        rows_p = _fill_rows(rows, cols, row_elems)
        if rows_p != rows and numel * elem <= FLAT_MAX_BYTES:
            return (*_flat_rows(ttnn, t, shape, numel, row_elems), 0)   # small: pad < row_elems elements, not rows
        if rows_p != rows:
            # Appending zero rows costs up to row_elems / gcd(cols, row_elems) - 1 rows: 80 MB for a [1, 1, 2, 20001]
            # fp32 tensor. Take one flat row (when it fits the page budget) or rows zero-padded to a pitch instead
            # when that is smaller by more than 1/8 of the tensor (typical shapes keep the contiguous layout).
            options = [(p * _fill_rows(rows, p, row_elems), p) for p in _row_pitches(cols, row_elems, elem)]
            if round_up(numel, row_elems) * elem <= ROW_PAGE_MAX_BYTES:
                options.append((round_up(numel, row_elems), 0))
            if options and rows_p * cols - min(options)[0] > numel // 8:
                size, pitch = min(options)                     # ties: the flat row, then the narrower pitch
                if not pitch:
                    return (*_flat_rows(ttnn, t, shape, numel, row_elems), 0)
                prows = size // pitch
                if shape != (1, 1, rows, cols):
                    t = ttnn.reshape(t, (1, 1, rows, cols))
                t = ttnn.pad(t, [(0, 0), (0, 0), (0, prows - rows), (0, pitch - cols)], 0.0)
                if pitch != row_elems:
                    t = ttnn.reshape(t, (1, 1, size // row_elems, row_elems))
                return t, size, pitch
        # zero rows appended, then one RM reshape into rows of row_elems (source pages of cols elements, destination
        # pages of row_elems elements: both small)
        if shape != (1, 1, rows, cols):
            t = ttnn.reshape(t, (1, 1, rows, cols))
        if rows_p != rows:
            t = ttnn.pad(t, [(0, 0), (0, 0), (0, rows_p - rows), (0, 0)], 0.0)
        if cols != row_elems:
            t = ttnn.reshape(t, (1, 1, rows_p * cols // row_elems, row_elems))
        return t, rows_p * cols, 0
    if rows != 1:
        raise ValueError(f"pack_outputs: {name!r} {shape} has rows of {cols * elem} B; the multi-row layout reads "
                         f"RM rows of at most {ROW_PAGE_MAX_BYTES} B: give it a narrower last dim (e.g. "
                         f"[1, 1, -1, {row_elems}]) before packing")
    # one wide row (a flat vector): cut it into chunks of whole packed rows, each well inside the L1 budget
    if shape != (1, 1, 1, numel):
        t = ttnn.reshape(t, (1, 1, 1, numel))
    width = max(row_elems, (ROW_PAGE_MAX_BYTES // elem) // row_elems * row_elems)
    parts, total = [], 0
    for start in range(0, numel, width):
        stop = min(start + width, numel)
        chunk = ttnn.slice(t, [0, 0, 0, start], [1, 1, 1, stop])
        chunk, padded = _flat_rows(ttnn, chunk, (1, 1, 1, stop - start), stop - start, row_elems)
        parts.append(chunk)
        total += padded
    return (parts[0] if len(parts) == 1 else ttnn.concat(parts, dim=2)), total, 0


def pack_outputs(tensors: Mapping[str, Any], *, dtype: Any = "float32", align: int = 32,
                 row_elems: Optional[int] = None) -> Packed:
    """Pack several device tensors into ONE ROW_MAJOR tensor so the host reads them with one device-to-host copy.

    Each tensor is typecast to ``dtype`` (float32 default: exact for bf16 and for integers below 2**24) and converted
    to ROW_MAJOR; call this inside the variant function (the ops become part of the trace) and return the result.
    :meth:`PackLayout.unpack` (done by ``TraceRunner.read``) gives ``{name: array}`` back in the original shapes.

    Device layout: while the packed total is at most :data:`SINGLE_ROW_MAX_ELEMS`, one ``[1, 1, 1, total]`` row (each
    tensor flattened and zero-padded to ``align`` elements, then concatenated): a ``[1, 1, N, 1]`` RM readback is N
    pages and costs milliseconds, one row is one page (RP section 1.4). Above it (or with ``row_elems=R``) a
    ``[1, 1, rows, R]`` tensor (default R = :data:`PACK_ROW_ELEMS`): each tensor is zero-padded to whole rows of R
    elements and the segments are concatenated on the row dim, so no RM page exceeds :data:`ROW_PAGE_MAX_BYTES`
    and outputs of tens of MB pack (a single row of more than ~0.6 MB does not fit the RM reshape's L1 staging:
    YOLOX PORT_LOG Q9). A tensor is padded with zero rows (contiguous), or, when that wastes more than 1/8 of it
    (an awkward last dim with few rows: e.g. ``[1, 1, 2, 20001]`` would take 1024 rows), flattened (up to the page
    budget) or stored with its rows zero-padded to a wider pitch (:class:`PackEntry` ``pitch``). A last dim wider
    than ``ROW_PAGE_MAX_BYTES`` is accepted for a flat vector (``[1, 1, 1, N]``, cut into chunks with
    ``ttnn.slice``); any other tensor needs a narrower last dim (``ValueError``)."""
    import ttnn

    if not tensors:
        raise ValueError("pack_outputs needs at least one tensor")
    if align <= 0:
        raise ValueError("align must be positive")
    if row_elems is not None and (int(row_elems) <= 0 or int(row_elems) % TILE):
        raise ValueError(f"row_elems={row_elems}: expected a positive multiple of {TILE}")
    dtype = ttnn_dtype(dtype)
    elem = _elem_bytes(dtype)
    items = [(str(name), t, tuple(int(s) for s in t.shape)) for name, t in tensors.items()]
    empty = [name for name, _, shape in items if math.prod(shape) == 0]
    if empty:
        raise ValueError(f"pack_outputs: {empty} have no elements")
    single_total = sum(round_up(math.prod(shape), align) for _, _, shape in items)
    if row_elems is None and single_total <= SINGLE_ROW_MAX_ELEMS:
        segments, entries, offset = [], [], 0
        for name, t, shape in items:                       # the single-row layout of ttaw 0.1.0 - 0.14.0
            numel = math.prod(shape)
            t = _row_major(ttnn, t, dtype)
            if shape != (1, 1, 1, numel):
                t = ttnn.reshape(t, (1, 1, 1, numel))
            padded = round_up(numel, align)
            if padded != numel:
                t = ttnn.pad(t, [(0, 0), (0, 0), (0, 0), (0, padded - numel)], 0.0)
            segments.append(t)
            entries.append(PackEntry(name, offset, numel, shape))
            offset += padded
        packed = segments[0] if len(segments) == 1 else ttnn.concat(segments, dim=-1)
        return Packed(packed, PackLayout(tuple(entries), offset))
    r = int(row_elems or PACK_ROW_ELEMS)
    segments, entries, offset = [], [], 0
    for name, t, shape in items:
        numel = math.prod(shape)
        t, padded, pitch = _pack_rows_segment(ttnn, name, _row_major(ttnn, t, dtype), shape, numel, r, elem)
        segments.append(t)
        entries.append(PackEntry(name, offset, numel, shape, pitch))
        offset += padded
    packed = segments[0] if len(segments) == 1 else ttnn.concat(segments, dim=2)
    return Packed(packed, PackLayout(tuple(entries), offset, rows=offset // r, row_elems=r))


# --------------------------------------------------------------------------------------------- internals

@dataclass
class _Slot:
    """A persistent device tensor (input, parameter or state) allocated before any capture."""

    name: str
    kind: str                     # "input" | "param" | "state" | "output"
    shape: Tuple[int, ...]
    dtype: Any
    layout: Any
    memory_config: Any
    buffers: List[Any]            # 1 buffer, or 2 for a ping-pong state
    init: Any                     # host tensor holding the initial value (states, params) / warm-up value (inputs)
    staging: Any = None           # DRAM staging copy (2CQ stage_inputs mode)
    stage_fn: Optional[Callable[[Any, Any], Any]] = None
    last_value: Optional[bytes] = None   # params: bytes of the last uploaded value (skip unchanged uploads)
    banks: List[Any] = field(default_factory=list)   # states: save / load buffers (D16 per-stream state banks)

    @property
    def pingpong(self) -> bool:
        return len(self.buffers) == 2

    def tensors(self) -> List[Any]:
        """Every device tensor the slot owns."""
        return self.buffers + ([self.staging] if self.staging is not None else []) + self.banks


@dataclass
class _Variant:
    name: str
    fn: Callable[["TraceContext"], Any]
    warmup_runs: int


@dataclass
class _Trace:
    variant: str
    phase: int
    trace_id: Any
    outputs: Any
    leaves: List[Tuple[Tuple, Any, Optional[PackLayout]]]   # (path, device tensor, pack layout or None)
    steps: bool                    # writes the ping-pong states (flips the phase)
    capture_ms: float
    host_buffers: Optional[List[Any]] = None
    done_event: Any = None


def _flatten(obj: Any, path: Tuple = ()) -> Iterator[Tuple[Tuple, Any]]:
    """Leaves of a tensor / Packed / list / tuple / dict structure with their paths."""
    if isinstance(obj, Mapping):
        for k, v in obj.items():
            yield from _flatten(v, path + (k,))
    elif isinstance(obj, (list, tuple)):
        for i, v in enumerate(obj):
            yield from _flatten(v, path + (i,))
    else:
        yield path, obj


def _rebuild(obj: Any, values: Dict[Tuple, Any], path: Tuple = ()) -> Any:
    if isinstance(obj, Mapping):
        return {k: _rebuild(v, values, path + (k,)) for k, v in obj.items()}
    if isinstance(obj, (list, tuple)):
        return type(obj)(_rebuild(v, values, path + (i,)) for i, v in enumerate(obj))
    return values[path]


def _buffer_address(t: Any) -> Optional[int]:
    try:
        return int(t.buffer_address())
    except Exception:  # noqa: BLE001 -- host tensor, deallocated tensor or a fake
        return None


def _same_buffer(a: Any, b: Any) -> bool:
    """``a`` is ``b`` or a tensor over the same device buffer."""
    if a is b:
        return True
    addr = _buffer_address(a)
    return addr is not None and addr == _buffer_address(b)


def _check_copy(what: str, value: Any, slot: "_Slot") -> None:
    """The preconditions of ``ttnn.copy(value, <slot buffer>)``, checked in Python so a mistake is a clear error
    instead of a TT_FATAL inside an open capture: same logical shape and layout; a dtype change only in TILE."""
    import ttnn

    shape = tuple(int(s) for s in value.shape)
    if shape != slot.shape:
        raise ValueError(f"{what}: value shape {shape} != {slot.shape}")
    if value.layout != slot.layout:
        raise ValueError(f"{what}: value layout {value.layout} != {slot.layout} (ttnn.copy keeps the layout)")
    if value.dtype != slot.dtype and slot.layout != ttnn.TILE_LAYOUT:
        raise ValueError(f"{what}: dtype {value.dtype} -> {slot.dtype} needs TILE layout (ttnn.copy)")


class TraceContext(Mapping):
    """What a variant function receives: persistent tensors by name (``ctx["x"]``), and state writes.

    ``ctx[name]`` returns an input, a parameter, a persistent output buffer, or the buffer a state is *read* from in
    this phase. ``ctx.write_state(name, value)`` copies ``value`` into the buffer the state is *written* to (in-place
    states: the same buffer, so read everything you need from it before writing); ``ctx.write_output(name, value)``
    copies into a persistent output and returns it."""

    def __init__(self, runner: "TraceRunner", variant: str, phase: int, capturing: bool):
        self._runner = runner
        self.variant = variant
        self.phase = phase
        self.capturing = capturing
        self.writes: set = set()

    @property
    def device(self):
        return self._runner.device

    def __getitem__(self, name: str):
        slot = self._runner._slot(name)
        if slot.kind == "state":
            return self.state(name)
        return slot.buffers[0]

    def __iter__(self) -> Iterator[str]:
        return iter(self._runner._slots)

    def __len__(self) -> int:
        return len(self._runner._slots)

    def state(self, name: str):
        """The buffer state ``name`` is read from in this phase."""
        slot = self._runner._slot(name, "state")
        return slot.buffers[self.phase] if slot.pingpong else slot.buffers[0]

    def write_target(self, name: str):
        """The buffer state ``name`` is written to in this phase (ping-pong: the other buffer; in place: the same
        one). Pass it as ``output_tensor=`` of the op that produces the new state and then call
        ``write_state(name, it)``: no copy program is traced (probe P13: saves one program per state per frame)."""
        slot = self._runner._slot(name, "state")
        return slot.buffers[1 - self.phase] if slot.pingpong else slot.buffers[0]

    def write_state(self, name: str, value) -> None:
        """``ttnn.copy(value, <write buffer>)``: same logical shape and layout as the state (dtype may differ in
        TILE layout). The copy is part of the trace; it is skipped when ``value`` already is the write buffer
        (:meth:`write_target`)."""
        import ttnn

        slot = self._runner._slot(name, "state")
        target = self.write_target(name)
        _check_copy(f"state {name!r}", value, slot)
        if not _same_buffer(value, target):
            ttnn.copy(value, target)
        self.writes.add(name)

    def write_output(self, name: str, value):
        """``ttnn.copy(value, <persistent output>)`` (part of the trace); returns the output buffer, which the
        variant can return as (part of) its outputs. Skipped when ``value`` already is that buffer."""
        import ttnn

        slot = self._runner._slot(name, "output")
        _check_copy(f"output {name!r}", value, slot)
        if not _same_buffer(value, slot.buffers[0]):
            ttnn.copy(value, slot.buffers[0])
        return slot.buffers[0]


class TraceRunner:
    """Persistent device I/O + warm-up + capture + replay of one model's traced stages (see the module docstring).

    Args:
        device: an open ttnn device.
        num_command_queues: 1 or 2. ``None`` uses what :func:`.device.open_device` recorded (1 if unknown).
        warmup_runs: eager runs of each variant (and phase) before any capture.
        stage_inputs: 2CQ only: upload into DRAM staging buffers and copy them into the trace inputs with an eager
            op on CQ0 (the upload of frame k+1 then overlaps the replay of frame k).
        forbid_cache_misses: call ``device.set_program_cache_misses_allowed(False)`` during capture.
        alloc_tracking: ``True`` raises unless the process tracks trace allocations
            (``TT_METAL_TRACE_ALLOC_TRACKING=1`` set before ``import ttnn``). Whenever tracking is active, the outputs
            of traces captured after the first are acknowledged as corruptible (see the module docstring).
        name: label used in messages and ``describe()``.
    """

    def __init__(self, device, *, num_command_queues: Optional[int] = None, warmup_runs: int = 1,
                 stage_inputs: bool = False, forbid_cache_misses: bool = True, alloc_tracking: Optional[bool] = None,
                 name: str = "model"):
        from .device import open_info

        opened = open_info(device).get("num_command_queues")
        if num_command_queues is None:
            num_command_queues = opened or 1
        if num_command_queues not in (1, 2):
            raise ValueError(f"num_command_queues={num_command_queues}: expected 1 or 2")
        if opened is not None and num_command_queues > opened:
            raise ValueError(f"the device was opened with {opened} command queue(s); cannot run {num_command_queues}")
        if stage_inputs and num_command_queues != 2:
            raise ValueError("stage_inputs=True needs num_command_queues=2")
        if warmup_runs < 1:
            raise ValueError("warmup_runs must be >= 1 (capture needs a warm program cache)")
        tracking = alloc_tracking_enabled()
        if alloc_tracking and not tracking:
            raise RuntimeError("alloc_tracking=True but trace allocation tracking is off: export "
                               "TT_METAL_TRACE_ALLOC_TRACKING=1 (optionally TT_METAL_TRACE_ALLOC_TRACEBACKS=1) "
                               "before Python imports ttnn")
        self.device = device
        self.name = name
        self.num_command_queues = int(num_command_queues)
        self.warmup_runs = int(warmup_runs)
        self.stage_inputs = bool(stage_inputs)
        self.forbid_cache_misses = bool(forbid_cache_misses)
        self.alloc_tracking = bool(tracking)
        self._slots: Dict[str, _Slot] = {}
        self._variants: Dict[str, _Variant] = {}
        self._traces: Dict[Tuple[str, int], _Trace] = {}
        self._phase = 0
        self._last_phase: Dict[str, int] = {}
        self._last_variant: Optional[str] = None
        self._pending_params: Dict[str, Any] = {}
        self._op_event = None
        self._stage_free_event = None
        self._captures = 0
        self._capturing = False
        self._eager_warmups: List[Callable[[], Any]] = []
        self._eager_warmed = False
        self._closed = False
        self.timings_ms: Dict[str, Dict[str, float]] = {}
        if hasattr(device, "enable_program_cache"):
            device.enable_program_cache()

    # ------------------------------------------------------------------------------------- registration
    def _check_registration(self, what: str) -> None:
        self._check_open()
        if self._capturing:
            raise RuntimeError(f"cannot add {what} while a variant is being warmed up or captured: register "
                               "inputs, params, states, outputs and variants before capture()")

    def _new_slot(self, name: str, kind: str, init: Any, shape: Optional[Sequence[int]], dtype: Any, layout: Any,
                  memory_config: Any, n_buffers: int, stage_fn: Optional[Callable], n_banks: int = 0) -> _Slot:
        import ttnn

        self._check_registration(f"{kind} {name!r}")
        if name in self._slots:
            raise ValueError(f"{name!r} is already registered (as {self._slots[name].kind})")
        if self._traces:
            raise RuntimeError(f"cannot add {kind} {name!r} after capture: persistent tensors must exist before the "
                               "first capture (release() and rebuild)")
        layout = ttnn.ROW_MAJOR_LAYOUT if layout is None else layout
        memory_config = ttnn.DRAM_MEMORY_CONFIG if memory_config is None else memory_config
        dtype = ttnn_dtype(dtype)
        if init is None:
            if shape is None:
                raise ValueError(f"{kind} {name!r}: give init= or shape=")
            init = np.zeros(tuple(shape), np.float32 if dtype_name(dtype) in ("float32", "bfloat16", "bfloat8_b",
                                                                              "bfloat4_b") else np.int64)
        if isinstance(init, ttnn.Tensor):
            host = to_host_tensor(init, dtype, layout, shape=shape)
        else:
            arr = to_numpy(init) if hasattr(init, "detach") else np.asarray(init)
            if shape is not None:
                arr = np.broadcast_to(arr, tuple(shape))
            host = to_host_tensor(arr, dtype, layout)
        shape_t = tuple(int(s) for s in host.shape)

        def allocate(config: Any):
            buf = ttnn.allocate_tensor_on_device(ttnn.Shape(list(shape_t)), dtype, layout, self.device, config)
            ttnn.copy_host_to_device_tensor(host, buf, cq_id=CQ_COMPUTE)   # defined contents, never garbage
            return buf

        buffers = [allocate(memory_config) for _ in range(n_buffers)]
        staging = allocate(ttnn.DRAM_MEMORY_CONFIG) if self.stage_inputs and kind in ("input", "param") else None
        banks = [allocate(ttnn.DRAM_MEMORY_CONFIG) for _ in range(n_banks)]
        slot = _Slot(name, kind, shape_t, dtype, layout, memory_config, buffers, host, staging,
                     stage_fn or (lambda src, dst: ttnn.copy(src, dst)), banks=banks)
        self._slots[name] = slot
        return slot

    def add_input(self, name: str, init: Any = None, *, shape: Optional[Sequence[int]] = None,
                  dtype: Any = "bfloat16", layout: Any = None, memory_config: Any = None,
                  stage_fn: Optional[Callable[[Any, Any], Any]] = None):
        """A persistent trace input (default ROW_MAJOR in DRAM). ``init`` (array / torch / ttnn host tensor) is the
        initial content and the warm-up input; zeros of ``shape`` otherwise -- give real data when the graph
        gathers with these values (garbage indices can hang the chip). ``stage_fn(staging, persistent)`` replaces
        the eager ``ttnn.copy`` in ``stage_inputs`` mode (e.g. a reshard into a sharded L1 input). Returns the
        device tensor."""
        return self._new_slot(name, "input", init, shape, dtype, layout, memory_config, 1, stage_fn).buffers[0]

    def add_param(self, name: str, value: Any = 0.0, *, shape: Sequence[int] = (1, 1, 1, 1), dtype: Any = "float32",
                  layout: Any = None, memory_config: Any = None):
        """An RT-dev parameter: a persistent device tensor (default fp32 ``[1, 1, 1, 1]`` TILE, which broadcasts in
        ttnn binary ops) refreshed before a replay when ``run(params=...)`` / :meth:`set_params` changes it."""
        import ttnn

        layout = ttnn.TILE_LAYOUT if layout is None else layout
        slot = self._new_slot(name, "param", np.broadcast_to(np.asarray(value), tuple(shape)), tuple(shape), dtype,
                              layout, memory_config, 1, None)
        slot.last_value = self._param_bytes(slot, value)
        return slot.buffers[0]

    def add_state(self, name: str, init: Any = None, *, shape: Optional[Sequence[int]] = None, dtype: Any = "float32",
                  layout: Any = None, memory_config: Any = None, pingpong: bool = False, banks: int = 0):
        """Temporal state kept on the device across replays (memory queues, previous BEV, ring buffers).

        In-place (default): one buffer, read via ``ctx[name]`` and written via ``ctx.write_state``. ``pingpong``:
        two buffers and two traces per variant. ``init`` is restored by :meth:`reset_state`. ``banks``: extra DRAM
        buffers of the same spec for :meth:`save_state` / :meth:`load_state` (one per stream id of a
        :class:`StreamBanks`, PLAN.md D16), allocated now because nothing may be allocated after a capture.
        Returns the buffer(s)."""
        import ttnn

        if banks < 0:
            raise ValueError("banks must be >= 0")
        layout = ttnn.TILE_LAYOUT if layout is None else layout
        slot = self._new_slot(name, "state", init, shape, dtype, layout, memory_config, 2 if pingpong else 1, None,
                              n_banks=int(banks))
        return tuple(slot.buffers) if pingpong else slot.buffers[0]

    def add_output(self, name: str, *, shape: Sequence[int], dtype: Any = "float32", layout: Any = None,
                   memory_config: Any = None):
        """A persistent output buffer (default TILE in DRAM, zeros) allocated before any capture. Variants write it
        with ``ctx.write_output(name, value)`` (one traced ``ttnn.copy``) and return it; its address survives
        recaptures and is shared by every variant (e.g. shape buckets with one readback). Returns the buffer."""
        import ttnn

        layout = ttnn.TILE_LAYOUT if layout is None else layout
        return self._new_slot(name, "output", None, shape, dtype, layout, memory_config, 1, None).buffers[0]

    def add_variant(self, name: str, fn: Callable[[TraceContext], Any], *, warmup_runs: Optional[int] = None) -> None:
        """Register a traced function. ``fn(ctx)`` runs ttnn ops on ``ctx[...]`` tensors and returns a device tensor,
        a :class:`Packed`, or a list / tuple / dict of them (the persistent trace outputs). No host I/O, no
        ``synchronize``, no torch ops inside ``fn``: it is called for warm-up and then recorded."""
        self._check_registration(f"variant {name!r}")
        if name in self._variants:
            raise ValueError(f"variant {name!r} already exists")
        self._variants[name] = _Variant(name, fn, self.warmup_runs if warmup_runs is None else int(warmup_runs))

    def add_eager_warmup(self, fn: Callable[[], Any]) -> None:
        """Register eager device work that the model runs *between* replays (host-fallback glue, an eager layout
        change of a read-back tensor, ...): ``fn()`` runs once in the first :meth:`capture`, before any trace is
        captured, so its programs are compiled -- and their kernel binaries allocated in DRAM -- before the first
        capture. A program compiled after a capture shares the address space of the traces' freed intermediates and
        a replay can overwrite its binaries (tt-metal ``tech_reports/.../TraceCorrectness.md``: corruption or a
        hang); ``TT_METAL_TRACE_ALLOC_TRACKING=1`` reports it."""
        self._check_registration("an eager warm-up")
        if self._traces or self._eager_warmed:
            raise RuntimeError("add eager warm-ups before the first capture (release() and rebuild)")
        self._eager_warmups.append(fn)

    # ------------------------------------------------------------------------------------------ capture
    @property
    def phases(self) -> int:
        """2 when a ping-pong state exists (two traces per variant), else 1."""
        return 2 if any(s.pingpong for s in self._slots.values()) else 1

    @property
    def captured(self) -> bool:
        return bool(self._variants) and all((v, p) in self._traces for v in self._variants for p in range(self.phases))

    def capture(self) -> None:
        """Warm up and capture every registered variant that has no trace yet (idempotent). If traces already exist
        and a new variant was added, all traces are released, the new variants warmed up, and all recaptured."""
        import ttnn

        self._check_open()
        if not self._variants:
            raise RuntimeError("no variant registered (add_variant)")
        pending = [v for v in self._variants if any((v, p) not in self._traces for p in range(self.phases))]
        if not pending:
            return
        if self._traces:
            ttnn.synchronize_device(self.device)        # no replay of a trace being released is still in flight
            self._release_traces()
            warm, pending = pending, list(self._variants)
        else:
            warm = pending
        self._capturing = True
        try:
            if not self._eager_warmed:
                self._warm_eager_programs()
            for vname in warm:
                t0 = time.perf_counter()
                variant = self._variants[vname]
                for phase in range(self.phases):
                    for _ in range(variant.warmup_runs):
                        ctx = TraceContext(self, vname, phase, capturing=False)
                        outputs = variant.fn(ctx)
                        ttnn.synchronize_device(self.device)
                        self._free_transient(outputs)
                self.timings_ms.setdefault(vname, {})["warmup"] = (time.perf_counter() - t0) * 1e3
            self._phase = 0
            self.reset_state()
            for vname in pending:
                self.timings_ms.setdefault(vname, {})["capture"] = 0.0
                for phase in range(self.phases):
                    self._capture_one(vname, phase)
        finally:
            self._capturing = False
        ttnn.synchronize_device(self.device)
        if self.num_command_queues == 2:
            self._op_event = ttnn.record_event(self.device, CQ_COMPUTE)
            self._stage_free_event = self._op_event

    def _warm_eager_programs(self) -> None:
        """Compile the runner's own eager programs before the first capture: the staging copies (``stage_inputs``),
        the state <-> bank copies (``banks``) and the registered eager warm-ups. Contents are kept: a staging buffer
        mirrors its input (:meth:`write_input` writes both), and a state and its first bank hold the same value
        before the states are reset ahead of the captures."""
        import ttnn

        for slot in self._slots.values():
            if slot.staging is not None:
                slot.stage_fn(slot.staging, slot.buffers[0])
            if slot.banks:
                ttnn.copy(slot.buffers[0], slot.banks[0])
                ttnn.copy(slot.banks[0], slot.buffers[0])
        for fn in self._eager_warmups:
            fn()
        ttnn.synchronize_device(self.device)
        self._eager_warmed = True

    @contextlib.contextmanager
    def _eager(self, what: str) -> Iterator[None]:
        """Eager runner work after a capture must not compile a program (see :meth:`add_eager_warmup`)."""
        dev = self.device
        count = getattr(dev, "num_program_cache_entries", None)
        before = count() if (self._traces and count is not None) else None
        yield
        if before is not None and count() != before:
            raise RuntimeError(f"{self.name}: {what} compiled a new program after capture; its kernel binaries share "
                               "DRAM with the traces' freed intermediates, so a replay can overwrite them. Run it "
                               "before the first capture (add_eager_warmup) or keep the tensor specs of the warmed "
                               "program")

    def _capture_one(self, vname: str, phase: int) -> None:
        import ttnn

        dev = self.device
        variant = self._variants[vname]
        ctx = TraceContext(self, vname, phase, capturing=True)
        entries_before = dev.num_program_cache_entries() if hasattr(dev, "num_program_cache_entries") else None
        forbid = self.forbid_cache_misses and hasattr(dev, "set_program_cache_misses_allowed")
        t0 = time.perf_counter()
        if forbid:
            dev.set_program_cache_misses_allowed(False)
        trace_id = ttnn.begin_trace_capture(dev, cq_id=CQ_COMPUTE)
        ok = ended = False
        try:
            outputs = variant.fn(ctx)
            ok = True
        finally:
            try:
                ttnn.end_trace_capture(dev, trace_id, cq_id=CQ_COMPUTE)
                ended = True
            finally:
                if forbid:
                    dev.set_program_cache_misses_allowed(True)
                if not (ok and ended):
                    self._safe_release(trace_id)
        try:
            entries_after = dev.num_program_cache_entries() if entries_before is not None else None
            if entries_before is not None and entries_after != entries_before:
                raise RuntimeError(f"{self.name}/{vname}: the program cache grew during capture ({entries_before} -> "
                                   f"{entries_after}); warm-up does not cover the traced graph")
            leaves = []
            for path, leaf in _flatten(outputs):
                if isinstance(leaf, Packed):
                    leaves.append((path, leaf.tensor, leaf.layout))
                elif isinstance(leaf, ttnn.Tensor):
                    leaves.append((path, leaf, None))
                else:
                    raise TypeError(f"{self.name}/{vname}: output {path} is {type(leaf).__name__}, expected a device "
                                    "tensor or Packed")
            if not leaves:
                raise ValueError(f"{self.name}/{vname}: the variant returned no output tensor")
            pingpong = {s.name for s in self._slots.values() if s.pingpong}
            written = ctx.writes & pingpong
            if written and written != pingpong:
                raise RuntimeError(f"{self.name}/{vname}: writes ping-pong states {sorted(written)} but not "
                                   f"{sorted(pingpong - written)}; a stepping variant must write all of them")
        except BaseException:
            self._safe_release(trace_id)
            raise
        if self.alloc_tracking and self._captures > 0:
            from ttnn.tools import trace_allocation_tracker as tracker

            keep = self._persistent_addresses()
            for _, tensor, _ in leaves:
                if _buffer_address(tensor) not in keep:
                    tracker.acknowledge_corruptible(tensor)
        self._captures += 1
        ms = (time.perf_counter() - t0) * 1e3
        self._traces[(vname, phase)] = _Trace(vname, phase, trace_id, outputs, leaves, bool(written), ms)
        self.timings_ms[vname]["capture"] += ms

    def _persistent_addresses(self) -> set:
        addrs = set()
        for slot in self._slots.values():
            for t in slot.tensors():
                a = _buffer_address(t)
                if a is not None:
                    addrs.add(a)
        return addrs

    def _deallocate_outputs(self, tensors: Iterator[Any]) -> None:
        """Deallocate op-produced output tensors, never a persistent buffer or a view of one. ``force=False``: a
        tensor sharing its device memory with another owner (a view of a model weight returned as an output) is
        left to its owners -- ``ttnn.deallocate`` forces by default and would free the weight."""
        import ttnn

        keep = self._persistent_addresses()
        seen = set()
        for tensor in tensors:
            if not isinstance(tensor, ttnn.Tensor) or id(tensor) in seen:
                continue
            seen.add(id(tensor))
            if tensor.is_allocated() and _buffer_address(tensor) not in keep:
                ttnn.deallocate(tensor, False)

    def _free_transient(self, outputs: Any) -> None:
        """Deallocate warm-up / eager outputs (see :meth:`_deallocate_outputs`)."""
        self._deallocate_outputs(leaf.tensor if isinstance(leaf, Packed) else leaf for _, leaf in _flatten(outputs))

    def _safe_release(self, trace_id) -> None:
        import ttnn

        try:
            ttnn.release_trace(self.device, trace_id)
        except Exception:  # noqa: BLE001 -- best effort on an error path; the original error is re-raised
            pass

    # -------------------------------------------------------------------------------------------- inputs
    def _slot(self, name: str, kind: Optional[str] = None) -> _Slot:
        slot = self._slots.get(name)
        if slot is None:
            raise KeyError(f"{self.name}: no input / param / state named {name!r}")
        if kind is not None and slot.kind != kind:
            raise KeyError(f"{self.name}: {name!r} is a {slot.kind}, not a {kind}")
        return slot

    @staticmethod
    def _param_array(slot: _Slot, value: Any) -> np.ndarray:
        dt = np.float32 if dtype_name(slot.dtype) in ("float32", "bfloat16", "bfloat8_b", "bfloat4_b") else np.int64
        return np.ascontiguousarray(np.broadcast_to(np.asarray(value, dtype=dt), slot.shape))

    def _param_bytes(self, slot: _Slot, value: Any) -> bytes:
        return self._param_array(slot, value).tobytes()

    def set_params(self, **values: Any) -> None:
        """Queue RT-dev parameter values for the next run (uploaded only if they changed)."""
        for name in values:
            self._slot(name, "param")
        self._pending_params.update(values)

    def _collect_uploads(self, inputs: Optional[Mapping[str, Any]],
                         params: Optional[Mapping[str, Any]]) -> List[Tuple[_Slot, Any, Optional[bytes]]]:
        uploads = []
        for name, value in (inputs or {}).items():
            slot = self._slot(name, "input")
            uploads.append((slot, to_host_tensor(value, slot.dtype, slot.layout, shape=slot.shape), None))
        merged = dict(self._pending_params)
        merged.update(params or {})
        for name, value in merged.items():
            slot = self._slot(name, "param")
            arr = self._param_array(slot, value)
            key = arr.tobytes()
            if key != slot.last_value:
                uploads.append((slot, to_host_tensor(arr, slot.dtype, slot.layout), key))
        return uploads

    def _ensure_events(self) -> None:
        """2CQ: the CQ0 events CQ1 waits for exist (``capture()`` records them; a partially failed capture, which
        leaves earlier variants runnable, does not get that far)."""
        import ttnn

        if self._op_event is None:
            self._op_event = ttnn.record_event(self.device, CQ_COMPUTE)
        if self._stage_free_event is None:
            self._stage_free_event = self._op_event

    def _enqueue_uploads(self, uploads: List[Tuple[_Slot, Any, Optional[bytes]]]) -> None:
        import ttnn

        if uploads:
            if self.num_command_queues == 1:
                for slot, host, _ in uploads:
                    ttnn.copy_host_to_device_tensor(host, slot.buffers[0], cq_id=CQ_COMPUTE)
            elif self.stage_inputs:
                self._ensure_events()
                ttnn.wait_for_event(CQ_INPUT, self._stage_free_event)
                for slot, host, _ in uploads:
                    ttnn.copy_host_to_device_tensor(host, slot.staging, cq_id=CQ_INPUT)
                written = ttnn.record_event(self.device, CQ_INPUT)
                ttnn.wait_for_event(CQ_COMPUTE, written)
                with self._eager("a stage_fn copy"):
                    for slot, _, _ in uploads:
                        slot.stage_fn(slot.staging, slot.buffers[0])
                self._stage_free_event = ttnn.record_event(self.device, CQ_COMPUTE)
            else:
                self._ensure_events()
                ttnn.wait_for_event(CQ_INPUT, self._op_event)
                for slot, host, _ in uploads:
                    ttnn.copy_host_to_device_tensor(host, slot.buffers[0], cq_id=CQ_INPUT)
                written = ttnn.record_event(self.device, CQ_INPUT)
                ttnn.wait_for_event(CQ_COMPUTE, written)
        for slot, _, key in uploads:
            if key is not None:
                slot.last_value = key
        self._pending_params.clear()

    def write_input(self, name: str, value: Any) -> None:
        """Upload ``value`` into input ``name`` now (CQ0, outside any trace): e.g. a realistic warm-up sample.
        With 2 CQs the event that later CQ1 uploads wait for is re-recorded after this write, so an upload of the
        next frame can never land before it (both write the same buffer from different queues)."""
        import ttnn

        self._check_open()
        slot = self._slot(name, "input")
        host = to_host_tensor(value, slot.dtype, slot.layout, shape=slot.shape)
        ttnn.copy_host_to_device_tensor(host, slot.buffers[0], cq_id=CQ_COMPUTE)
        if slot.staging is not None:   # the staging buffer mirrors the input (the stage-copy warm-up keeps it)
            ttnn.copy_host_to_device_tensor(host, slot.staging, cq_id=CQ_COMPUTE)
        if self.num_command_queues == 2 and self._op_event is not None:
            self._op_event = ttnn.record_event(self.device, CQ_COMPUTE)
            if self.stage_inputs:
                self._stage_free_event = self._op_event

    # --------------------------------------------------------------------------------------------- run
    def _trace_for(self, variant: Optional[str], phase: Optional[int] = None) -> _Trace:
        self._check_open()
        name = variant or self._last_variant or (next(iter(self._variants)) if len(self._variants) == 1 else None)
        if name is None:
            raise ValueError("variant name required")
        if name not in self._variants:
            raise KeyError(f"{self.name}: unknown variant {name!r}; have {sorted(self._variants)}")
        p = self._phase if phase is None else phase
        trace = self._traces.get((name, p))
        if trace is None:
            raise RuntimeError(f"{self.name}: variant {name!r} is not captured; call capture() first")
        return trace

    def _execute(self, trace: _Trace) -> None:
        import ttnn

        ttnn.execute_trace(self.device, trace.trace_id, cq_id=CQ_COMPUTE, blocking=False)
        if self.num_command_queues == 2:
            self._op_event = ttnn.record_event(self.device, CQ_COMPUTE)
            trace.done_event = self._op_event
        self._last_variant = trace.variant
        self._last_phase[trace.variant] = trace.phase
        if trace.steps:
            self._phase ^= 1

    def upload(self, inputs: Optional[Mapping[str, Any]] = None, params: Optional[Mapping[str, Any]] = None) -> int:
        """Enqueue the uploads of ``inputs`` and changed ``params`` (CQ0, or CQ1 + events with 2 CQs) without
        replaying anything; the next :meth:`run` / :meth:`replay` consumes them. Returns the number of tensors
        uploaded."""
        self._check_open()
        if not self.captured:
            raise RuntimeError(f"{self.name}: capture() before uploading")
        uploads = self._collect_uploads(inputs, params)
        self._enqueue_uploads(uploads)
        return len(uploads)

    def run(self, variant: Optional[str] = None, inputs: Optional[Mapping[str, Any]] = None,
            params: Optional[Mapping[str, Any]] = None):
        """Upload ``inputs`` (name -> numpy / torch / ttnn host tensor) and changed ``params``, then replay the
        variant's trace without blocking. Returns its device outputs (valid once the replay finished)."""
        trace = self._trace_for(variant)
        self._enqueue_uploads(self._collect_uploads(inputs, params))
        self._execute(trace)
        return trace.outputs

    def replay(self, variant: Optional[str] = None, n: int = 1) -> None:
        """Replay ``n`` times with no uploads (back-to-back device timing); honours ping-pong phases."""
        for _ in range(n):
            self._execute(self._trace_for(variant))

    def outputs(self, variant: Optional[str] = None, phase: Optional[int] = None):
        """Device outputs of a variant's trace (default: the phase that ran last for it, else phase 0)."""
        name = variant or self._last_variant
        if phase is None and name is not None:
            phase = self._last_phase.get(name, 0)
        return self._trace_for(name, phase if phase is not None else 0).outputs

    def read(self, variant: Optional[str] = None, *, cq_id: int = CQ_COMPUTE, as_torch: bool = False):
        """Read the outputs of the last run of ``variant`` into preallocated host tensors (blocking).

        Returns the output structure with numpy arrays (``as_torch=True``: torch tensors); a :class:`Packed` leaf
        becomes ``{name: array}``. ``cq_id=1`` (2CQ) waits on the host for the trace's completion event and reads on
        CQ1, so CQ0 can already replay the next segment (segmented D2H)."""
        import ttnn

        name = variant or self._last_variant
        if name is None or name not in self._last_phase:
            raise RuntimeError(f"{self.name}: variant {name!r} has not run yet")
        trace = self._trace_for(name, self._last_phase[name])
        if cq_id not in (CQ_COMPUTE, CQ_INPUT):
            raise ValueError(f"cq_id={cq_id}: expected 0 or 1")
        if cq_id == CQ_INPUT:
            if self.num_command_queues != 2:
                raise ValueError("cq_id=1 needs a device opened with 2 command queues")
            ttnn.event_synchronize(trace.done_event)
        if trace.host_buffers is None:
            trace.host_buffers = [ttnn.allocate_tensor_on_host(t.spec, self.device) for _, t, _ in trace.leaves]
        values: Dict[Tuple, Any] = {}
        for (path, tensor, layout), host in zip(trace.leaves, trace.host_buffers):
            ttnn.copy_device_to_host_tensor(tensor, host, blocking=True, cq_id=cq_id)
            arr = to_numpy(host)
            value: Any = layout.unpack(arr) if layout is not None else arr
            if as_torch:
                import torch

                value = ({k: torch.from_numpy(np.ascontiguousarray(v)) for k, v in value.items()}
                         if isinstance(value, dict) else torch.from_numpy(np.ascontiguousarray(value)))
            values[path] = value
        structure = trace.outputs
        if isinstance(structure, Packed) or not isinstance(structure, (Mapping, list, tuple)):
            return values[()]
        return _rebuild(structure, values)

    def __call__(self, variant: Optional[str] = None, inputs: Optional[Mapping[str, Any]] = None,
                 params: Optional[Mapping[str, Any]] = None, *, as_torch: bool = False):
        """``run`` + ``read`` (the common synchronous path)."""
        self.run(variant, inputs, params)
        return self.read(variant, as_torch=as_torch)

    def run_eager(self, variant: Optional[str] = None, inputs: Optional[Mapping[str, Any]] = None,
                  params: Optional[Mapping[str, Any]] = None):
        """Upload, run the variant's function *eagerly* (no trace) and return its outputs read to numpy, freeing
        every eager buffer before returning (safe next to captured traces). State writes happen as in a replay.
        Use it for replay-vs-eager bit checks."""
        import ttnn

        self._check_open()
        name = variant or self._last_variant or (next(iter(self._variants)) if self._variants else None)
        if name not in self._variants:
            raise KeyError(f"{self.name}: unknown variant {name!r}; have {sorted(self._variants)}")
        variant_obj = self._variants[name]
        uploads = self._collect_uploads(inputs, params)
        for slot, host, key in uploads:
            ttnn.copy_host_to_device_tensor(host, slot.buffers[0], cq_id=CQ_COMPUTE)
            if key is not None:
                slot.last_value = key
        self._pending_params.clear()
        ctx = TraceContext(self, name, self._phase, capturing=False)
        with self._eager(f"run_eager({name!r})"):
            outputs = variant_obj.fn(ctx)
        values = {}
        for path, leaf in _flatten(outputs):
            if isinstance(leaf, Packed):
                values[path] = leaf.layout.unpack(to_numpy(leaf.tensor))
            else:
                values[path] = to_numpy(leaf)
        self._free_transient(outputs)
        pingpong = {s.name for s in self._slots.values() if s.pingpong}
        if pingpong and (ctx.writes & pingpong) == pingpong:
            self._phase ^= 1
        if isinstance(outputs, Packed) or not isinstance(outputs, (Mapping, list, tuple)):
            return values[()]
        return _rebuild(outputs, values)

    # ------------------------------------------------------------------------------------------- state
    @property
    def phase(self) -> int:
        """Which buffer of each ping-pong state the next run reads (0 = the first buffer)."""
        return self._phase

    def state_buffer(self, name: str):
        """The device buffer the next run reads for state ``name`` (ping-pong: the buffer of the current
        :attr:`phase`)."""
        slot = self._slot(name, "state")
        return slot.buffers[self._phase] if slot.pingpong else slot.buffers[0]

    def reset_state(self, name: Optional[str] = None, value: Any = None) -> None:
        """Write the initial value (or ``value``) of state ``name`` (all states when ``None``) into the buffer the
        next run reads. ``value``: numpy / scalar / torch / ttnn host tensor (uploaded) or a ttnn *device* tensor of
        the same shape (``ttnn.copy`` on the device). Enqueued on CQ0, so it lands after any replay already
        enqueued."""
        import ttnn

        self._check_open()
        names = [name] if name is not None else [s.name for s in self._slots.values() if s.kind == "state"]
        for n in names:
            slot = self._slot(n, "state")
            target = self.state_buffer(n)
            if isinstance(value, ttnn.Tensor) and value.storage_type() == ttnn.StorageType.DEVICE:
                _check_copy(f"reset_state({n!r})", value, slot)
                with self._eager(f"reset_state({n!r}) from a device tensor"):
                    ttnn.copy(value, target)
                continue
            if value is None:
                host = slot.init
            elif isinstance(value, ttnn.Tensor):
                host = to_host_tensor(value, slot.dtype, slot.layout, shape=slot.shape)
            else:
                arr = to_numpy(value) if hasattr(value, "detach") else np.asarray(value)
                host = to_host_tensor(np.broadcast_to(arr, slot.shape), slot.dtype, slot.layout)
            ttnn.copy_host_to_device_tensor(host, target, cq_id=CQ_COMPUTE)

    def read_state(self, name: str) -> np.ndarray:
        """The value the next run will read for state ``name`` (blocking read on CQ0)."""
        return to_numpy(self.state_buffer(name))

    def _bank(self, name: str, bank: int):
        slot = self._slot(name, "state")
        if not 0 <= bank < len(slot.banks):
            raise IndexError(f"{self.name}: state {name!r} has {len(slot.banks)} bank(s) (add_state(banks=...)), "
                             f"not bank {bank}")
        return slot.banks[bank]

    def save_state(self, name: str, bank: int) -> None:
        """Copy the current value of state ``name`` into its ``bank`` (``ttnn.copy`` on CQ0, eager, after any
        replay already enqueued): the first half of a stream switch (PLAN.md D16; :class:`StreamBanks`)."""
        import ttnn

        self._check_open()
        with self._eager(f"save_state({name!r})"):
            ttnn.copy(self.state_buffer(name), self._bank(name, bank))

    def load_state(self, name: str, bank: int) -> None:
        """Copy ``bank`` back into the buffer the next run reads for state ``name`` (``ttnn.copy`` on CQ0)."""
        import ttnn

        self._check_open()
        with self._eager(f"load_state({name!r})"):
            ttnn.copy(self._bank(name, bank), self.state_buffer(name))

    # ------------------------------------------------------------------------------------------- misc
    def trace_ids(self) -> Dict[Tuple[str, int], Any]:
        """``{(variant, phase): trace_id}`` for profiling scripts."""
        return {k: t.trace_id for k, t in self._traces.items()}

    def describe(self) -> Dict[str, Any]:
        """A JSON-able summary for ``model.info`` / OPT_BASELINE."""
        def spec(s: _Slot) -> Dict[str, Any]:
            return {"shape": list(s.shape), "dtype": dtype_name(s.dtype), "layout": str(s.layout).rsplit(".", 1)[-1],
                    **({"pingpong": s.pingpong, "banks": len(s.banks)} if s.kind == "state" else {})}

        entries = None
        if hasattr(self.device, "num_program_cache_entries"):
            entries = int(self.device.num_program_cache_entries())
        return {
            "name": self.name, "num_command_queues": self.num_command_queues, "stage_inputs": self.stage_inputs,
            "warmup_runs": self.warmup_runs, "phases": self.phases, "alloc_tracking": self.alloc_tracking,
            "variants": sorted(self._variants), "traces": len(self._traces),
            "inputs": {s.name: spec(s) for s in self._slots.values() if s.kind == "input"},
            "params": {s.name: spec(s) for s in self._slots.values() if s.kind == "param"},
            "states": {s.name: spec(s) for s in self._slots.values() if s.kind == "state"},
            "outputs": {s.name: spec(s) for s in self._slots.values() if s.kind == "output"},
            "timings_ms": {k: {kk: round(vv, 3) for kk, vv in v.items()} for k, v in self.timings_ms.items()},
            "program_cache_entries": entries,
        }

    def _release_traces(self) -> None:
        traces, self._traces = self._traces, {}
        for trace in traces.values():
            self._safe_release(trace.trace_id)
        self._last_phase.clear()
        self._last_variant = None
        self._captures = 0
        self._deallocate_outputs(tensor for trace in traces.values() for _, tensor, _ in trace.leaves)

    def release(self) -> None:
        """Release every trace and deallocate the persistent tensors. Idempotent; the persistent tensors are freed
        even when the device sync or a trace release raises (the error propagates afterwards)."""
        import ttnn

        if self._closed:
            return
        self._closed = True
        try:
            try:
                ttnn.synchronize_device(self.device)
            finally:
                self._release_traces()
        finally:
            slots, self._slots = list(self._slots.values()), {}
            for slot in slots:
                for t in slot.tensors():
                    if t.is_allocated():
                        ttnn.deallocate(t)

    def _check_open(self) -> None:
        if self._closed:
            raise RuntimeError(f"{self.name}: the TraceRunner was released")

    def __enter__(self) -> "TraceRunner":
        return self

    def __exit__(self, *exc) -> None:
        self.release()


class StreamBanks:
    """Per-stream device state of a temporal model (PLAN.md D16) on top of the states of a :class:`TraceRunner`.

    The traces read and write one set of state buffers: the *active* stream's. With ``max_streams > 1`` every known
    stream id owns one bank of each state (``add_state(..., banks=max_streams)``) and a switch saves the active
    stream into its bank and loads the selected one (``ttnn.copy`` on CQ0, after the replays already enqueued;
    probe P13 measured ~80 us eager for a StreamPETR-size state). A new id beyond ``max_streams`` is refused
    (``on_full="reject"``: :class:`~.io.InputError`, HTTP 400) or takes over the least recently used stream
    (``"evict"``; with ``max_streams=1`` that simply restarts the one state). A stream starts fresh -- its states
    reset to their ``init`` values -- when it is new, when ``reset=True``, or when its timestamp goes backwards or
    jumps by more than ``max_gap_s``; :meth:`select` returns True then, so the model can also set its first-frame
    RT-dev params (BEVFormer ``use_prev_bev=0``, BEVDet ``flag``, ...). Create it after the last ``capture()`` (a
    capture resets every state) and call :meth:`select` under the model lock, before the run of each frame::

        S_MAX = 1                                            # first publish (D16)
        runner.add_state("prev_bev", shape=(1, 1, 22500, 256), dtype="bfloat16", pingpong=True,
                         banks=S_MAX if S_MAX > 1 else 0)
        streams = StreamBanks(runner, ["prev_bev"], max_streams=S_MAX, on_full="evict", max_gap_s=2.0)
        ...
        fresh = streams.select(stream.get("id", "default"), reset=stream.get("reset", False),
                               timestamp_s=stream.get("timestamp_s"))
        out = runner("frame", inputs=..., params={"use_prev_bev": 0.0 if fresh else 1.0})
    """

    def __init__(self, runner: TraceRunner, states: Sequence[str], *, max_streams: int = 1, on_full: str = "reject",
                 max_gap_s: Optional[float] = None):
        if int(max_streams) < 1:
            raise ValueError("max_streams must be >= 1")
        if on_full not in ("reject", "evict"):
            raise ValueError(f"on_full={on_full!r}: expected 'reject' or 'evict'")
        self.runner = runner
        self.states = tuple(states)
        if not self.states:
            raise ValueError("StreamBanks needs at least one state")
        self.max_streams = int(max_streams)
        for name in self.states:
            banks = len(runner._slot(name, "state").banks)
            if self.max_streams > 1 and banks < self.max_streams:
                raise ValueError(f"state {name!r} has {banks} bank(s); StreamBanks(max_streams={self.max_streams}) "
                                 f"needs add_state(..., banks={self.max_streams})")
        self.on_full = on_full
        self.max_gap_s = None if max_gap_s is None else float(max_gap_s)
        self.active: Optional[str] = None
        self._bank: Dict[str, int] = {}       # stream id -> bank index (max_streams > 1)
        self._last_t: Dict[str, float] = {}   # stream id -> timestamp of its last frame
        self._used: Dict[str, int] = {}       # stream id -> last use (LRU clock)
        self._clock = 0

    @property
    def streams(self) -> List[str]:
        """Known stream ids, most recently used first."""
        return sorted(self._used, key=lambda s: -self._used[s])

    def _save_active(self) -> None:
        if self.active is not None and self.active in self._bank:
            for name in self.states:
                self.runner.save_state(name, self._bank[self.active])

    def forget(self, stream_id: str) -> None:
        """Drop a stream (its bank becomes free; if it was active, the next :meth:`select` starts fresh)."""
        sid = str(stream_id)
        self._used.pop(sid, None)
        self._bank.pop(sid, None)
        self._last_t.pop(sid, None)
        if self.active == sid:
            self.active = None

    def select(self, stream_id: Any = "default", *, reset: bool = False, timestamp_s: Optional[float] = None) -> bool:
        """Make ``stream_id`` the active stream for the next run; returns True when its state starts fresh."""
        sid = str(stream_id)
        fresh = bool(reset)
        if sid != self.active:
            if sid in self._used:                                  # a known stream parked in its bank
                self._save_active()
                if not fresh:
                    for name in self.states:
                        self.runner.load_state(name, self._bank[sid])
            else:                                                  # a new stream
                if len(self._used) >= self.max_streams:
                    if self.on_full == "reject":
                        raise InputError(f"stream {sid!r}: this model keeps device state for {self.max_streams} "
                                         f"stream(s), in use by {self.streams}; reuse an id")
                    self.forget(min(self._used, key=self._used.__getitem__))
                self._save_active()
                if self.max_streams > 1:
                    self._bank[sid] = min(set(range(self.max_streams)) - set(self._bank.values()))
                fresh = True
            self.active = sid
        if not fresh and timestamp_s is not None and self.max_gap_s is not None and sid in self._last_t:
            dt = float(timestamp_s) - self._last_t[sid]
            fresh = dt < 0 or dt > self.max_gap_s
        if fresh:
            for name in self.states:
                self.runner.reset_state(name)
        if timestamp_s is not None:
            self._last_t[sid] = float(timestamp_s)
        elif fresh:
            self._last_t.pop(sid, None)
        self._clock += 1
        self._used[sid] = self._clock
        return fresh

    def describe(self) -> Dict[str, Any]:
        """JSON-able summary for ``model.info``."""
        return {"max_streams": self.max_streams, "on_full": self.on_full, "max_gap_s": self.max_gap_s,
                "states": list(self.states), "active": self.active, "streams": self.streams}