File size: 53,931 Bytes
b025706
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# SPDX-FileCopyrightText: © 2024 Tenstorrent USA, Inc.

# SPDX-License-Identifier: Apache-2.0


import math

import torch
from tqdm import tqdm

import ttnn
from models.common.lightweightmodule import LightweightModule
from models.common.rmsnorm import RMSNorm
from models.common.sampling.generator import SamplingGenerator
from models.common.sampling.tt_sampling import TOPK_MAX_WIDTH, TTSampling
from models.tt_transformers.tt.ccl import TT_CCL
from models.tt_transformers.tt.common import Mode, copy_host_to_device
from models.tt_transformers.tt.decoder import TransformerBlock
from models.tt_transformers.tt.distributed_norm import DistributedNorm
from models.tt_transformers.tt.embedding import Embedding, ScaledEmbedding
from models.tt_transformers.tt.lm_head import LMHead
from models.tt_transformers.tt.model_config import TensorGroup
from models.tt_transformers.tt.rope import HfRotarySetup, RotarySetup


def _get_trace_rope_table_len(max_seq_len, trace_prefill_seq_lens):
    if not trace_prefill_seq_lens:
        return max_seq_len

    slice_alignment = math.lcm(*trace_prefill_seq_lens)
    min_table_len = max_seq_len + max(trace_prefill_seq_lens)
    return ((min_table_len + slice_alignment - 1) // slice_alignment) * slice_alignment


def _prefill_rope_setups_to_pad(rope_setup, rope_local_setup, rope_setup_class):
    """Select the RoPE setups whose shared prefill tables Transformer slices.

    A caller-supplied rope_setup_class builds its prefill cosine and sine mats
    per request on the host, inside its own prepare_inputs_prefill, and its
    forward never calls Transformer._slice_prefill_rot_mats. Such a setup owns
    no cos_matrix_prefill to pad. rope_local_setup always comes from the
    built-in classes, so it always owns one.
    """
    rope_setups = [] if rope_setup_class is not None else [rope_setup]
    if rope_local_setup is not None:
        rope_setups.append(rope_local_setup)
    return rope_setups


def _pad_prefill_rope_tables(rope_setups, max_seq_len, trace_prefill_seq_lens):
    table_len = _get_trace_rope_table_len(max_seq_len, trace_prefill_seq_lens)
    pad_len = table_len - max_seq_len
    if pad_len == 0:
        return

    padding = [(0, 0), (0, 0), (0, pad_len), (0, 0)]
    for rope_setup in rope_setups:
        rope_setup.cos_matrix_prefill = ttnn.pad(
            rope_setup.cos_matrix_prefill,
            padding=padding,
            value=0.0,
        )
        rope_setup.sin_matrix_prefill = ttnn.pad(
            rope_setup.sin_matrix_prefill,
            padding=padding,
            value=0.0,
        )


class Transformer(LightweightModule):
    def __init__(
        self,
        args,
        dtype,
        mesh_device,
        state_dict,
        weight_cache_path,
        paged_attention_config=None,
        use_paged_kv_cache=False,
        attention_class=None,
        rope_setup_class=None,
        block_class=None,
        lm_head_cls=None,
        final_norm_builder=None,
        prefetcher=None,
    ):
        super().__init__()
        self.args = args
        self.vocab_size = args.vocab_size
        assert self.vocab_size > 0
        self.n_layers = args.n_layers
        self.mesh_device = mesh_device
        self.dtype = dtype
        self.model_config = args.get_model_config()
        self.grid_size = self.args.max_grid_size
        state_dict_prefix = args.get_state_dict_prefix("", None)
        self.decoders_optimizations = args.decoders_optimizations
        self.prefetcher = prefetcher
        self.tt_ccl = TT_CCL(self.mesh_device)
        # Runtime bounds for the post-prefill tail's slice. Allocated here, before any trace exists,
        # and rewritten in place per call - see process_logits_after_prefill_trace.
        # These buffers belong to this model/DP lane. Calls on a lane enqueue the
        # copy and slice on the same command queue, in order. Concurrent host
        # calls on the same model instance are not supported.
        self._tail_slice_start = ttnn.from_torch(
            torch.zeros(4, dtype=torch.int32),
            device=self.mesh_device,
            mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
        )
        self._tail_slice_end = ttnn.from_torch(
            torch.zeros(4, dtype=torch.int32),
            device=self.mesh_device,
            mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
        )

        embd_kwargs = {
            "mesh_device": mesh_device,
            "args": args,
            "weight_cache_path": args.weight_cache_path(dtype),
            "state_dict": state_dict,
            "dtype": ttnn.bfloat16,  # Row major layout requires bfloat16
        }
        if self.args.embed_scale is not None:
            embd_cls = ScaledEmbedding
            embd_kwargs["embed_scale"] = self.args.embed_scale
        else:
            embd_cls = Embedding
        self.embd = embd_cls(**embd_kwargs)

        DefaultRopeSetup = HfRotarySetup if self.args.use_hf_rope else RotarySetup
        ActualRopeSetupClass = rope_setup_class if rope_setup_class is not None else DefaultRopeSetup
        # NoPE global layers (EXAONE-4.x): full-attention layers apply no rotary at
        # all, so the global setup's cos/sin are neutralized to the identity. Only
        # the Meta-style RotarySetup implements this.
        use_global_nope = getattr(args, "use_global_nope", False)
        if use_global_nope and self.args.use_hf_rope:
            raise NotImplementedError("use_global_nope (NoPE global layers) requires the Meta-style RotarySetup")
        global_rope_kwargs = {"nope": True} if use_global_nope else {}
        self.rope_setup = ActualRopeSetupClass(
            device=mesh_device,
            batch_size=args.max_batch_size,
            head_dim=args.head_dim,
            max_seq_len=args.max_seq_len,
            rope_theta=args.rope_theta,
            rope_scaling=args.rope_scaling,
            use_qk_fused=args.use_qk_fused,
            prefetcher=prefetcher,
            **global_rope_kwargs,
        )

        if args.rope_theta_local:
            self.rope_local_setup = DefaultRopeSetup(
                mesh_device,
                args.max_batch_size,
                args.head_dim,
                args.max_seq_len,
                args.rope_theta_local,
                # Most hybrid models (Gemma-3) use unscaled rope on sliding layers;
                # EXAONE-4.x instead applies its llama3-scaled rope there and sets
                # rope_scaling_local (the global layers being NoPE).
                rope_scaling=getattr(args, "rope_scaling_local", None),
                use_qk_fused=args.use_qk_fused,
                prefetcher=None,
            )

        # Dynamic starts share one table across fixed-width trace buckets. The
        # tail prevents out-of-range reads and the common multiple preserves
        # the tensor-bound slice partition geometry for every traced length.
        _pad_prefill_rope_tables(
            _prefill_rope_setups_to_pad(
                self.rope_setup,
                getattr(self, "rope_local_setup", None),
                rope_setup_class,
            ),
            args.max_seq_len,
            args.trace_prefill_supported_seq_lens,
        )

        self.trans_mats_dict = self.rope_setup.get_both_trans_mats()

        # Device tensors used to build dynamic slice params for prefill RoPE slicing.
        # Keeps chunk_start_idx-driven slicing inside the traced graph.
        self._tt_seq_len_buffer = ttnn.from_torch(
            torch.tensor([1, 1, self.args.max_seq_len, self.args.head_dim], dtype=torch.int32),
            device=self.mesh_device,
            mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
        )
        self._tt_slice_start_zeros_4 = ttnn.from_torch(
            torch.tensor([0, 0, 0, 0], dtype=torch.int32),
            device=self.mesh_device,
            mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
        )

        # Model-family dispatch (Command-R / cohere): swap the decoder block, final
        # norm and LM head. Lazy imports keep the experimental cohere module out of the default path.
        ActualBlockClass = block_class
        ActualLMHeadCls = lm_head_cls
        final_norm_builder_resolved = final_norm_builder
        if str(getattr(self.args, "model_type", None) or "").lower() == "cohere":
            from models.experimental.cohere.tt.cohere_decoder import CohereDecoderLayer
            from models.experimental.cohere.tt.cohere_lm_head import CohereLMHead
            from models.experimental.cohere.tt.cohere_norm import build_cohere_final_norm

            ActualBlockClass = ActualBlockClass or CohereDecoderLayer
            ActualLMHeadCls = ActualLMHeadCls or CohereLMHead
            final_norm_builder_resolved = final_norm_builder_resolved or build_cohere_final_norm
        ActualBlockClass = ActualBlockClass or TransformerBlock
        ActualLMHeadCls = ActualLMHeadCls or LMHead

        self.layers = [
            ActualBlockClass(
                args=args,
                mesh_device=mesh_device,
                tt_ccl=self.tt_ccl,
                dtype=dtype,
                state_dict=state_dict,
                weight_cache_path=weight_cache_path,
                layer_num=i,
                transformation_mats=self.trans_mats_dict,
                paged_attention_config=paged_attention_config,
                use_paged_kv_cache=use_paged_kv_cache,
                attention_class=attention_class,
                prefetcher=prefetcher,
            )
            for i in tqdm(range(self.n_layers))
        ]
        self.norm = (
            final_norm_builder_resolved(
                args=args,
                mesh_device=mesh_device,
                state_dict=state_dict,
                weight_cache_path=weight_cache_path,
                dtype=dtype,
                tt_ccl=self.tt_ccl,
            )
            if final_norm_builder_resolved is not None
            else DistributedNorm(
                RMSNorm(
                    device=mesh_device,
                    dim=args.dim,
                    eps=args.norm_eps,
                    state_dict=state_dict,
                    state_dict_prefix=args.get_state_dict_prefix("", None),
                    weight_cache_path=None if args.dummy_weights else weight_cache_path,
                    weight_dtype=ttnn.bfloat16,
                    weight_key="norm",
                    add_unit_offset=self.args.rms_norm_add_unit_offset,
                    is_distributed=self.args.is_distributed_norm,
                    ccl_topology=self.args.ccl_topology(),
                    tt_ccl=self.tt_ccl,
                ),
                args,
                tt_ccl=self.tt_ccl,
                prefetcher=prefetcher,
                TG=args.is_galaxy,
            )
        )  # close the final_norm_builder_resolved conditional-expression paren

        self.lm_head = ActualLMHeadCls(
            args=args,
            mesh_device=mesh_device,
            tt_ccl=self.tt_ccl,
            dtype=dtype,
            state_dict=state_dict,
            state_dict_prefix=state_dict_prefix,
            weight_cache_path=weight_cache_path,
            max_columns_per_device=self.args.max_columns_per_device_lm_head,
            prefetcher=prefetcher,
        )

        # Initialize on-device sampling if supported
        # Sampling on device is supported only if each device holds at most TOPK_MAX_WIDTH logits.
        # On a single device TTSampling cuts the padded vocab into as many same-device chunks as
        # needed (power-of-two, each <= TOPK_MAX_WIDTH), so any vocab it can cut tile-aligned is
        # supported (#53064); anything it cannot falls back to host sampling.
        padded_vocab_size = getattr(self.args, "padded_vocab_size", None) or self.args.vocab_size
        if list(self.mesh_device.shape) != [1, 1]:
            vocab_fits_on_device = padded_vocab_size // self.args.num_devices <= TOPK_MAX_WIDTH
        else:
            vocab_fits_on_device = TTSampling.num_single_device_vocab_splits(padded_vocab_size) is not None
        self._supports_on_device_sampling = prefetcher is None and vocab_fits_on_device
        if self._supports_on_device_sampling:
            self.sampling = SamplingGenerator(
                args=args,
                mesh_device=mesh_device,
                tt_ccl=self.tt_ccl,
            )
        else:
            self.sampling = None

    def update_weights(
        self,
        hf_state_dict: dict[str, ttnn.Tensor],
        *,
        hf_rope: bool = False,
    ) -> None:
        """In-place replace every weight from an HF-keyed dict of on-device 4D
        ttnn tensors (replicated, DRAM-interleaved, TILE, bf16). Keys follow HF
        safetensors naming; shapes are HF Linear/gamma/embedding wrapped in two
        leading unit dims.

        Strict by construction: every required key must be present (missing ->
        ``KeyError``) and every provided key consumed by exactly one leaf
        ``.update()`` (extras -> ``ValueError``). No "loose" mode -- silent
        partial updates are an expensive class of bug.

        ``hf_rope=False`` (default): caller has already permuted Q/K rows into
        this model's convention (right for the ttml -> TTT transfer, both store
        Meta-permuted rows). ``hf_rope=True`` defers HF -> Meta permutation to
        ``Attention.update`` (currently raises -- kernel not wired up).

        Tied embeddings: the protocol still requires both
        ``model.embed_tokens.weight`` and ``lm_head.weight`` (typically the same
        source tensor), keeping dispatch one-to-one with device buffers.

        Every existing buffer keeps its device allocation, so captured traces
        and the prefetcher's recorded addresses stay valid.
        """
        unconsumed = set(hf_state_dict.keys())

        def consume(key: str) -> ttnn.Tensor:
            if key not in hf_state_dict:
                raise KeyError(f"Transformer.update_weights: missing required HF key {key!r}")
            unconsumed.discard(key)
            return hf_state_dict[key]

        # Top-level (always required).
        self.embd.update(embed_tokens=consume("model.embed_tokens.weight"))
        self.norm.update(weight=consume("model.norm.weight"))
        self.lm_head.update(weight=consume("lm_head.weight"))

        # Per-layer: prefix-strip into a layer-local dict, dispatch.
        for i, block in enumerate(self.layers):
            prefix = f"model.layers.{i}."
            layer_dict = {}
            for key in list(hf_state_dict.keys()):
                if key.startswith(prefix):
                    layer_dict[key[len(prefix) :]] = hf_state_dict[key]
                    unconsumed.discard(key)
            block.update_weights(layer_dict, hf_rope=hf_rope)

        if unconsumed:
            sample = sorted(unconsumed)[:10]
            raise ValueError(
                f"Transformer.update_weights: {len(unconsumed)} HF key(s) not "
                f"consumed by any leaf .update(). This usually means a typo, "
                f"a stray weight, or a layer-index off-by-one. "
                f"Showing up to 10: {sample}"
            )

    def process_logits_after_prefill_trace(self, logits, last_token_idx):
        get_last_token = (last_token_idx // 32) * 32
        seq_len = int(logits.shape[-2])
        # Pass the offset as a runtime argument rather than a compile-time attribute. With literal
        # bounds every distinct prompt offset compiles its own slice program, and since this runs
        # after the prefill traces are captured - once per data-parallel group - that was the single
        # largest source of buffers left live across trace replays on a DP run. Warmup cannot cover
        # it either: it only ever sees bucket-length mock prompts, and real prompts are shorter.
        #
        # The tensor-args path needs the slice tile-aligned, which this one already is: it takes 32
        # rows starting at a multiple of 32. num_devices splits the sequence into equal parts, so
        # seq_len // 32 gives exactly the 32-row window, and the program then keys on the padded
        # prefill bucket instead of the offset.
        if seq_len % 32 == 0:
            for device_tensor, values in (
                (self._tail_slice_start, [0, 0, get_last_token, 0]),
                (self._tail_slice_end, [1, 1, get_last_token + 32, int(logits.shape[-1])]),
            ):
                ttnn.copy_host_to_device_tensor(
                    ttnn.from_torch(
                        torch.tensor(values, dtype=torch.int32),
                        mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
                    ),
                    device_tensor,
                )
            logits = ttnn.slice(
                input_tensor=logits,
                starts=self._tail_slice_start,
                ends=self._tail_slice_end,
                slice_dim=2,
                num_devices=seq_len // 32,
            )
        else:
            logits = ttnn.slice(
                logits,
                (0, 0, get_last_token, 0),
                (1, 1, get_last_token + 32, logits.shape[-1]),
            )
        logits = self._apply_norm_and_lm_head(logits)
        return logits

    def extract_last_tokens_batched_prefill(
        self, hidden_states, last_token_idx_list, padded_batch, prefill_seq_len, target_batch=None
    ):
        """Extract each user's last-token hidden state from batched prefill output.

        Reads hidden states to host, extracts the relevant row for each user,
        and sends the combined tensor back to device with the correct column-sharded
        mesh mapping (ShardTensorToMesh dim=-1) so the DistributedNorm all-gather
        produces the correct full hidden dim.

        Args:
            hidden_states: [padded_batch, 1, prefill_seq_len, dim_per_device] on device (column-sharded, TILE_LAYOUT)
            last_token_idx_list: list of length padded_batch with per-user last token positions
            padded_batch: number of slots (typically 32)
            prefill_seq_len: padded sequence length per user

        Returns:
            user_tokens: [1, 1, target_batch or padded_batch, dim_per_device] per device,
            column-sharded, TILE_LAYOUT
        """
        active_indices = [lt for lt in last_token_idx_list if lt > 0]
        all_same = len(set(active_indices)) <= 1

        if all_same and active_indices:
            common_last = active_indices[0]
            get_last = (common_last // 32) * 32
            R = common_last % 32
            block = ttnn.slice(
                hidden_states,
                (0, 0, get_last, 0),
                (padded_batch, 1, get_last + 32, hidden_states.shape[-1]),
            )
        else:
            block = hidden_states
            R = None

        host_tensors = [ttnn.to_torch(dt) for dt in ttnn.get_device_tensors(block)]
        host_full = torch.cat(host_tensors, dim=-1)

        if R is not None:
            combined = host_full[:, :, R : R + 1, :].reshape(1, 1, padded_batch, -1).contiguous()
        else:
            rows = []
            for slot in range(padded_batch):
                lt_idx = last_token_idx_list[slot]
                rows.append(host_full[slot : slot + 1, :, lt_idx : lt_idx + 1, :])
            combined = torch.cat(rows, dim=0).reshape(1, 1, padded_batch, -1).contiguous()

        target_batch = padded_batch if target_batch is None else target_batch
        if target_batch < padded_batch:
            raise ValueError(f"target_batch {target_batch} must be >= padded_batch {padded_batch}")
        if target_batch > padded_batch:
            padded_combined = torch.zeros(
                1,
                1,
                target_batch,
                combined.shape[-1],
                dtype=combined.dtype,
            )
            padded_combined[:, :, :padded_batch, :] = combined
            combined = padded_combined

        user_tokens = ttnn.from_torch(
            combined,
            device=self.mesh_device,
            dtype=ttnn.bfloat16,
            layout=ttnn.TILE_LAYOUT,
            mesh_mapper=ttnn.ShardTensorToMesh(self.mesh_device, dim=-1),
        )
        return user_tokens

    def process_logits_after_batched_prefill(self, hidden_states, last_token_idx_list, padded_batch, prefill_seq_len):
        """Extract last tokens and run norm + lm_head once for all users."""
        user_tokens = self.extract_last_tokens_batched_prefill(
            hidden_states, last_token_idx_list, padded_batch, prefill_seq_len
        )
        return self._apply_norm_and_lm_head(user_tokens)

    def _apply_final_logit_softcapping(self, logits):
        """Gemma-2 final logit soft-capping: logits -> tanh(logits / cap) * cap.

        No-op unless args.final_logit_softcapping is set (only Gemma-2 sets it), so
        this leaves every other model's output path unchanged.
        """
        cap = self.args.final_logit_softcapping
        if cap is None or cap <= 0:
            return logits
        logits = ttnn.multiply(logits, 1.0 / cap)
        logits = ttnn.tanh(logits)
        logits = ttnn.multiply(logits, cap)
        return logits

    def _apply_norm_and_lm_head(self, x):
        """Shared norm + lm_head for prefill logit processing. Input: [1, 1, 32, hidden_dim]."""
        x = self.norm(
            x, mode=Mode.PREFILL, norm_config=self.args.get_norm_config("lm_head", Mode.PREFILL, self.prefetcher)
        )
        lm_head_input_mem_cfg = self.args.get_lm_head_input_mem_config(Mode.PREFILL, None)
        if lm_head_input_mem_cfg.is_sharded():
            x = ttnn.interleaved_to_sharded(x, lm_head_input_mem_cfg)
        logits = self.lm_head(x)
        logits = self._apply_final_logit_softcapping(logits)
        logits = ttnn.to_memory_config(logits, memory_config=ttnn.DRAM_MEMORY_CONFIG)
        return logits

    def process_hidden_states_after_prefill_trace(self, hidden_states, last_token_idx):
        """
        Process hidden states after prefill trace, stopping before LM head.
        Returns hidden states (after norm) instead of logits.
        Used for embedding models that need hidden states rather than logits.
        """
        get_last_token = (last_token_idx // 32) * 32
        hidden_states = ttnn.slice(
            hidden_states,
            (0, 0, get_last_token, 0),
            (1, 1, get_last_token + 32, hidden_states.shape[-1]),
        )
        # Apply norm (this is the final layer norm before LM head)
        hidden_states = self.norm(hidden_states, mode="prefill")
        # Convert to row major layout for output (but don't apply LM head)
        hidden_states = ttnn.to_layout(
            hidden_states, layout=ttnn.ROW_MAJOR_LAYOUT, memory_config=ttnn.DRAM_MEMORY_CONFIG
        )
        return hidden_states

    def prepare_prefill_inputs_trace(
        self,
        tokens,
        page_table=None,
        chunk_page_table=None,
        chunk_start_idx=0,
        batch_size=1,
        user_id=0,
        **kwargs,
    ):
        """
        Inputs are torch tensors or python types. This function returns ttnn
        tensors on host.
        """
        host_inputs = self.prepare_inputs_prefill(
            tokens,
            page_table=page_table,
            chunk_page_table=chunk_page_table,
            chunk_start_idx=chunk_start_idx,
            trace_enabled=True,
            batch_size=batch_size,
            user_id=user_id,
        )
        return host_inputs

    def transform_and_embed_prefill_inputs_device(
        self,
        tokens,
        tt_page_table,
        tt_chunk_page_table,
        tt_chunk_start_idx,
    ):
        tt_tokens = self.embd(tokens)
        tt_tokens = ttnn.unsqueeze_to_4D(tt_tokens)
        return tt_tokens, tt_page_table, tt_chunk_page_table, tt_chunk_start_idx

    def prepare_inputs_prefill(
        self,
        tokens,
        start_pos=0,
        page_table=None,
        chunk_page_table=None,
        chunk_start_idx=None,
        trace_enabled=False,
        last_token_idx=None,
        global_user_id=None,
        batch_size=1,
        user_id=0,
        **kwargs,
    ):
        """
        Inputs are torch tensors or python types. This function returns ttnn
        tensors on device if trace is disabled or on host if trace is enabled.
        TODO: Debate whether this function is responsible for padding
        """

        # We set the device to None if trace is enabled so we keep the tensors on host instead of sending it to the device (None - keeps on host, device - sends to specified device)
        # We will send them to device later (copy_host_to_device)
        device = None if trace_enabled else self.mesh_device

        assert tokens.dim() == 2, "tokens must be a 2D tensor"
        # For batched prefill, tokens come in as [padded_batch, S]
        # Each user's tokens are at their slot index in dimension 0
        # Reshape to [1, 1, 1, padded_batch * S] for embedding
        if batch_size > 1:
            # Tokens are in slot-based format [padded_batch, S_per_user]
            S = tokens.shape[-1]  # Per-user sequence length
            tokens = tokens.reshape(1, 1, 1, -1)  # Flatten to [1, 1, 1, padded_batch * S]
        else:
            tokens = tokens.reshape(1, 1, 1, -1)
            S = tokens.shape[-1]
        tokens = ttnn.from_torch(
            tokens,
            device=device,
            dtype=ttnn.uint32,
            layout=ttnn.ROW_MAJOR_LAYOUT,
            mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
        )

        # self.embd expects that tokens are on device ; if trace is enabled, the tensors will be later on device, so we will do these 2 steps when we copy the tokens to the device
        if not trace_enabled:
            tokens_embd = self.embd(tokens)
            tokens_embd = ttnn.unsqueeze_to_4D(tokens_embd)

        # Slice the rot mats to the prefill seqlen
        trace_mat_len = self.rope_setup.cos_matrix_prefill.shape[2]
        mat_len = self.args.max_seq_len
        seq_len = last_token_idx + 1 if last_token_idx is not None else S
        assert mat_len >= seq_len, f"Sequence length {seq_len} exceeds max seq len {mat_len}"

        required_end = start_pos + S
        pad_len = max(0, required_end - mat_len)

        # We set the end_pos to max_seq_len so that we don't create a new tensor for the whole cos_matrix and sin_matrix
        # In case of trace, we will use the whole matrix for all seq_lens supported by trace
        prefill_start_pos = 0 if trace_enabled else start_pos
        slice_end = trace_mat_len if trace_enabled else min(mat_len, required_end)

        cos_slice = self.rope_setup.cos_matrix_prefill[:, :, prefill_start_pos:slice_end, :]
        sin_slice = self.rope_setup.sin_matrix_prefill[:, :, prefill_start_pos:slice_end, :]

        if pad_len > 0:
            # Padding: [(before, after), ...] for each dim; pad at end of 3rd dim (dim=2) by pad_len
            padding = [(0, 0)] * 4
            padding[2] = (0, pad_len)
            cos_slice = ttnn.pad(cos_slice, padding=padding, value=0.0)
            sin_slice = ttnn.pad(sin_slice, padding=padding, value=0.0)

        tt_rot_mats_prefill_global = [cos_slice, sin_slice]

        if hasattr(self, "rope_local_setup"):
            local_trace_mat_len = self.rope_local_setup.cos_matrix_prefill.shape[2]
            local_mat_len = self.args.max_seq_len
            local_required_end = start_pos + S
            local_pad_len = max(0, local_required_end - local_mat_len)
            local_slice_end = local_trace_mat_len if trace_enabled else min(local_mat_len, local_required_end)

            local_cos_slice = self.rope_local_setup.cos_matrix_prefill[:, :, prefill_start_pos:local_slice_end, :]
            local_sin_slice = self.rope_local_setup.sin_matrix_prefill[:, :, prefill_start_pos:local_slice_end, :]

            if local_pad_len > 0:
                # Pad at end of 3rd dim (dim=2) by local_pad_len
                local_padding = [(0, 0)] * 4
                local_padding[2] = (0, local_pad_len)
                local_cos_slice = ttnn.pad(local_cos_slice, padding=local_padding, value=0.0)
                local_sin_slice = ttnn.pad(local_sin_slice, padding=local_padding, value=0.0)

            tt_rot_mats_prefill_local = [local_cos_slice, local_sin_slice]
        else:
            tt_rot_mats_prefill_local = None

        if page_table is not None:
            # For batched prefill, replicate page_table to all devices (same as single-user path)
            # The KV cache fill will loop over users and use batch_idx=user_id for each
            tt_page_table = ttnn.from_torch(
                page_table,
                device=device,
                dtype=ttnn.int32,
                layout=ttnn.ROW_MAJOR_LAYOUT,
                mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
            )
        else:
            tt_page_table = None

        if chunk_page_table is not None:
            tt_chunk_page_table = ttnn.from_torch(
                chunk_page_table,
                device=device,
                dtype=ttnn.int32,
                layout=ttnn.ROW_MAJOR_LAYOUT,
                mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
            )
        else:
            tt_chunk_page_table = None

        if chunk_start_idx is not None and int(chunk_start_idx) > 0:
            chunk_start_idx_tensor = torch.tensor([chunk_start_idx], dtype=torch.int32)
            tt_chunk_start_idx = ttnn.from_torch(
                chunk_start_idx_tensor,
                device=device,
                dtype=ttnn.int32,
                mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
            )
        else:
            tt_chunk_start_idx = None

        return (
            tokens if trace_enabled else tokens_embd,
            tt_rot_mats_prefill_global,
            tt_rot_mats_prefill_local,
            tt_page_table,
            tt_chunk_page_table,
            tt_chunk_start_idx,
        )

    def prepare_inputs_decode(self, *inputs):
        """
        Inputs are torch tensors or python types. This function returns ttnn
        tensors on device.
        Its implementation can take advantage of a few other functions which the
        model must implement.
        """
        host_inputs = self.prepare_decode_inputs_host(*inputs)
        device_inputs = copy_host_to_device(host_inputs, mesh_device=self.mesh_device)  # Helper function
        return device_inputs

    def prepare_decode_inputs_host(self, tokens, current_pos, page_table=None):
        """
        Inputs are torch tensors or python types. Outputs are ttnn tensors on host.
        NOTE: Tokens and current_pos are padded to batch
        """
        B = tokens.shape[0]
        assert current_pos.shape[0] == B, "Batch size mismatch"
        assert (
            B == self.args.max_batch_size
        ), f"Batch size {B} must be equal to max_batch_size {self.args.max_batch_size}"

        # Necessary padding to be full tile sized when on device
        tokens = torch.nn.functional.pad(tokens.view(-1), (0, 32 - len(tokens)), "constant", 0)
        tokens = ttnn.from_torch(
            tokens,
            device=None,
            dtype=ttnn.uint32,
            mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
        )
        tokens = ttnn.unsqueeze_to_4D(tokens)

        rot_current_pos = torch.maximum(
            current_pos, torch.tensor(0, dtype=torch.int64)
        )  # Ensure position indices are non-negative
        rope_idxs = self.rope_setup.get_rot_idxs(rot_current_pos, on_host=True)

        current_pos_tt = ttnn.from_torch(
            current_pos,
            device=None,
            dtype=ttnn.int32,
            mesh_mapper=ttnn.ShardTensor2dMesh(
                self.mesh_device,
                dims=(None, 0) if (self.args.is_galaxy and B > 1) else (None, None),
                mesh_shape=self.args.cluster_shape,
            ),
        )

        if page_table is not None:
            page_table = ttnn.from_torch(
                page_table,
                device=None,
                dtype=ttnn.int32,
                mesh_mapper=ttnn.ShardTensor2dMesh(
                    self.mesh_device,
                    dims=(None, -2) if (self.args.is_galaxy and B > 1) else (None, None),
                    mesh_shape=self.args.cluster_shape,
                ),
            )
        return tokens, current_pos_tt, rope_idxs, page_table

    def _transform_decode_inputs_device(
        self,
        tokens,
    ):
        """
        Inputs are ttnn tensors on device. This function applies any on-device
        transformations which should happen before forward decode.
        For example: tilize, reshape, shard.
        Return transformed device tensors

        Embed tokens
        """
        decode_residual_mem_cfg = self.args.get_residual_mem_config(Mode.DECODE, self.prefetcher)
        tt_tokens = self.embd(
            tokens,
            memory_config=ttnn.DRAM_MEMORY_CONFIG if self.prefetcher is None else decode_residual_mem_cfg,
        )
        tt_tokens = ttnn.unsqueeze_to_4D(tt_tokens)
        tt_tokens = ttnn.to_memory_config(tt_tokens, decode_residual_mem_cfg)
        return tt_tokens

    def concat_host_output(self, tt_out, is_log_probs=False):
        """
        Concatenate the output of the devices into a single host tensor.
        """
        torch_out_tensors = [ttnn.to_torch(x) for x in ttnn.get_device_tensors(tt_out)]
        if self.args.is_galaxy:
            row_dim, col_dim = (3, 1)
        else:
            row_dim, col_dim = (1, -1)

        rows, cols = self.args.cluster_shape
        mesh_shape = [torch_out_tensors[i : i + cols] for i in range(0, len(torch_out_tensors), cols)]
        if is_log_probs:
            row_concatenated = []
            for row in mesh_shape:
                row_reshaped = [tensor.reshape(1, 1, -1, 1) for tensor in row]
                row_concatenated.append(torch.cat(row_reshaped, dim=col_dim))
        else:
            row_concatenated = [torch.cat(row, dim=col_dim) for row in mesh_shape]

        return torch.cat(row_concatenated, dim=row_dim)

    def process_output_prefill(self, tt_out, last_token_idx):
        """
        Input is ttnn host tensor of logits. Output is torch logits tensor.
        NOTE: In this model, prefill always uses get_last_token
        """
        assert tt_out.storage_type() == ttnn.StorageType.HOST, "Expected host tensor"
        return self.concat_host_output(tt_out)[0, 0, last_token_idx, : self.vocab_size]

    def process_output_prefill_hidden_states(self, tt_out, last_token_idx):
        """
        Input is ttnn host tensor of hidden states (after norm, before LM head).
        Output is torch hidden states tensor of shape [hidden_size].
        Used for embedding models.
        """
        assert tt_out.storage_type() == ttnn.StorageType.HOST, "Expected host tensor"
        # Extract the last token's hidden state
        # Shape: [batch=1, head=1, seq, hidden_dim] -> [hidden_dim]
        # For hidden states, if they're replicated across devices (not sharded),
        # we should take just the first device's output to avoid incorrect concatenation.
        # If sharded, concat_host_output will properly concatenate them.
        concatenated = self.concat_host_output(tt_out)
        # Check if concatenation resulted in oversized tensor (replicated case)
        # If so, take only the first device's portion (first self.args.dim elements)
        if concatenated.shape[-1] > self.args.dim:
            # Hidden states are replicated, take first device's output
            return concatenated[0, 0, last_token_idx, : self.args.dim]
        else:
            # Hidden states are sharded, concatenation is correct
            return concatenated[0, 0, last_token_idx, :]

    def process_output_decode(self, tt_out, B, S=1, is_tokens=False, is_log_probs=False):
        """
        Input is ttnn host tensor of logits if is_tokens=False, otherwise tokens. Output is the corresponding torch tensor.
        """
        if is_tokens or is_log_probs:
            # Pad to 32 to match the expected batch size for decode operations (tiles are 32x32)
            padded_batch_size = 32
            if not is_log_probs:
                tt_out = ttnn.reshape(tt_out, ttnn.Shape([1, 1, padded_batch_size, 1]))
            return self.concat_host_output(tt_out, is_log_probs)[0, 0, :B, 0]
        if self.args.num_devices > 1:
            tt_out = ttnn.to_torch(ttnn.get_device_tensors(tt_out)[0]).float()
        else:
            tt_out = ttnn.to_torch(tt_out).float()
        tt_out = tt_out[:, :, :B, : self.vocab_size].view(B, S, -1)
        return tt_out

    def ttnn_prefill_forward(
        self,
        x,
        rot_mats_global=None,
        rot_mats_local=None,
        user_id=0,
        page_table=None,
        chunk_page_table=None,
        chunk_start_idx=None,
        get_last_token=-1,
        kv_cache=None,
        batch_size=1,
        page_tables_per_layer=None,
    ):
        """
        This method will take device tensors and any other args to run forward.
        It returns ttnn device tensors.
        """
        if page_tables_per_layer is None:
            # vLLM hybrid bridges (HybridAttentionForCausalLM subclasses) stash
            # the per-layer list on the model handle for the duration of a
            # forward call rather than threading the kwarg through Generator's
            # many ttnn_prefill_forward call sites. Pick it up here when set.
            page_tables_per_layer = getattr(self, "_active_page_tables_per_layer", None)
        page_tables_per_layer = self._page_tables_to_ttnn(page_tables_per_layer)
        return self.forward(
            x,
            current_pos=None,
            rot_mats_global=rot_mats_global,
            rot_mats_local=rot_mats_local,
            user_id=user_id,
            mode=Mode.PREFILL,
            page_table=page_table,
            chunk_page_table=chunk_page_table,
            chunk_start_idx=chunk_start_idx,
            get_last_token=get_last_token,
            kv_cache=kv_cache,
            batch_size=batch_size,
            page_tables_per_layer=page_tables_per_layer,
        )

    def _page_table_mesh_mapper(self, B):
        """Mesh mapper for per-layer page tables, matching the layout that
        :meth:`prepare_decode_inputs_host` uses for the legacy single
        ``page_table`` kwarg: shard the batch dim across mesh axis 1 on
        Galaxy when ``B>1``, replicate otherwise. The hybrid bridge
        chunks the global page table per-DP before calling into a
        submesh, so ``B`` here is the per-DP batch — same value the
        legacy path sees on entry to ``prepare_decode_inputs_host``.
        """
        return ttnn.ShardTensor2dMesh(
            self.mesh_device,
            dims=(None, -2) if (self.args.is_galaxy and B > 1) else (None, None),
            mesh_shape=self.args.cluster_shape,
        )

    def _page_tables_to_ttnn(self, page_tables_per_layer):
        """Resolve a per-layer list of ``torch.Tensor`` page tables to a
        list of *persistent* ttnn device tensors (allocate-only).

        Tracing bakes each input tensor's device address into the captured
        graph; replaying the trace reads from those exact addresses
        regardless of any new ttnn objects created on the Python side.
        Allocating fresh device tensors on every call would therefore
        make traced inference read stale memory at the original
        addresses, so we lazily allocate one persistent device tensor per
        layer on first use and *only* update contents from outside the
        traced ``ttnn_*_forward`` calls (writes are forbidden during trace
        capture). The hybrid bridge calls
        :meth:`update_persistent_per_layer_page_tables` *before* invoking
        ``Generator``'s decode/prefill which executes traces — that's
        where content updates happen.

        First call (warmup compile) populates the persistent buffers from
        the input torch tensors; subsequent calls return the existing
        buffers unchanged. ``None`` entries propagate; already-ttnn
        entries pass through.
        """
        if page_tables_per_layer is None:
            return None
        persistent = getattr(self, "_persistent_per_layer_page_tables", None)
        n = len(page_tables_per_layer)
        if persistent is None or len(persistent) != n:
            persistent = []
            for pt in page_tables_per_layer:
                if pt is None:
                    persistent.append(None)
                    continue
                if isinstance(pt, ttnn.Tensor):
                    persistent.append(pt)
                    continue
                persistent.append(
                    ttnn.from_torch(
                        pt,
                        device=self.mesh_device,
                        dtype=ttnn.int32,
                        layout=ttnn.ROW_MAJOR_LAYOUT,
                        mesh_mapper=self._page_table_mesh_mapper(pt.shape[0]),
                    )
                )
            self._persistent_per_layer_page_tables = persistent
        return persistent

    def update_persistent_per_layer_page_tables(self, page_tables_per_layer):
        """Update content of persistent per-layer page_table device
        tensors in place. Called by the hybrid bridge *before* invoking
        ``Generator``'s decode/prefill so traced replay observes the new
        block IDs at the captured addresses. Must be called outside trace
        capture (writes forbidden inside).

        No-op if persistent tensors haven't been allocated yet (first
        call goes through :meth:`_page_tables_to_ttnn`'s allocation).
        """
        if page_tables_per_layer is None:
            return
        persistent = getattr(self, "_persistent_per_layer_page_tables", None)
        if persistent is None or len(persistent) != len(page_tables_per_layer):
            return
        for i, pt in enumerate(page_tables_per_layer):
            if pt is None or persistent[i] is None or isinstance(pt, ttnn.Tensor):
                continue
            host_pt = ttnn.from_torch(
                pt,
                device=None,
                dtype=ttnn.int32,
                layout=ttnn.ROW_MAJOR_LAYOUT,
                mesh_mapper=self._page_table_mesh_mapper(pt.shape[0]),
            )
            ttnn.copy_host_to_device_tensor(host_pt, persistent[i])

    def _increment_decode_positions_device(self, current_pos, rot_mat_idxs):
        ttnn.plus_one(current_pos, skip_negative_entries=True)
        ttnn.plus_one(rot_mat_idxs)

    def _slice_prefill_rot_mats(self, rot_mats, chunk_start_idx, prefill_seq_len):
        """Slice full prefill RoPE mats to the traced prefill sequence length."""
        if rot_mats is None or chunk_start_idx is None or not isinstance(chunk_start_idx, ttnn.Tensor):
            return rot_mats

        full_rot_cos, full_rot_sin = rot_mats[0], rot_mats[1]
        full_seq_len = full_rot_cos.shape[2]
        if prefill_seq_len <= 0:
            raise ValueError(f"Prefill sequence length must be positive, got {prefill_seq_len}")
        if full_rot_sin.shape[2] != full_seq_len:
            raise ValueError(
                f"Prefill RoPE cosine and sine sequence lengths must match, got "
                f"{full_seq_len} and {full_rot_sin.shape[2]}"
            )
        if full_seq_len == prefill_seq_len:
            return rot_mats
        if full_seq_len % prefill_seq_len != 0:
            raise ValueError(
                f"Full RoPE sequence length {full_seq_len} must be evenly divisible by "
                f"prefill sequence length {prefill_seq_len}"
            )
        # Tensor-bound slice fixes output geometry as input length divided by
        # num_devices; this argument is a partition count, not the mesh width.
        num_partitions = full_seq_len // prefill_seq_len

        z = self._tt_slice_start_zeros_4
        tt_slice_starts = ttnn.concat([z[0:2], chunk_start_idx, z[3:4]], dim=0)

        rot_cos_slice = ttnn.slice(
            input_tensor=full_rot_cos,
            starts=tt_slice_starts,
            ends=self._tt_seq_len_buffer,
            slice_dim=2,
            num_devices=num_partitions,
        )
        rot_sin_slice = ttnn.slice(
            input_tensor=full_rot_sin,
            starts=tt_slice_starts,
            ends=self._tt_seq_len_buffer,
            slice_dim=2,
            num_devices=num_partitions,
        )
        return (rot_cos_slice, rot_sin_slice)

    def ttnn_decode_forward(
        self,
        x,
        current_pos,
        rot_mat_idxs=None,
        page_table=None,
        kv_cache=None,
        on_device_logits=False,
        page_tables_per_layer=None,
    ):
        """
        This method will take device tensors and any other args to run forward.
        It returns ttnn device tensors.
        """
        rot_mats_global = self.rope_setup.get_rot_mats(rot_mat_idxs)
        rot_mats_local = self.rope_local_setup.get_rot_mats(rot_mat_idxs) if hasattr(self, "rope_local_setup") else None

        x_embed = self._transform_decode_inputs_device(x)

        if page_tables_per_layer is None:
            # See ttnn_prefill_forward: hybrid bridges stash the per-layer list
            # on the model when active, since Generator doesn't thread the kwarg.
            page_tables_per_layer = getattr(self, "_active_page_tables_per_layer", None)
        page_tables_per_layer = self._page_tables_to_ttnn(page_tables_per_layer)

        tt_logits = self.forward(
            x_embed,
            current_pos,
            rot_mats_global=rot_mats_global,
            rot_mats_local=rot_mats_local,
            mode=Mode.DECODE,
            page_table=page_table,
            kv_cache=kv_cache,
            page_tables_per_layer=page_tables_per_layer,
        )

        if on_device_logits:
            assert self.sampling is not None, (
                "ttnn_decode_forward got on_device_logits=True but no on-device sampling "
                "module exists (self.sampling is None)."
            )
            self._increment_decode_positions_device(current_pos, rot_mat_idxs)
            return tt_logits

        # Gather the output across all devices and untilize the tensor (for argmax)
        if self.args.num_devices > 1:
            cluster_axis = 0 if self.args.is_galaxy else None
            num_links = 2 if self.args.is_galaxy else 1
            tt_logits = ttnn.experimental.all_gather_async(
                tt_logits,
                persistent_output_buffer=None,
                dim=3,
                multi_device_global_semaphore=self.tt_ccl.get_and_cycle_ag_semaphore_handles(cluster_axis),
                num_links=num_links,
                memory_config=tt_logits.memory_config() if self.prefetcher is None else ttnn.DRAM_MEMORY_CONFIG,
                cluster_axis=cluster_axis,
                topology=self.args.ccl_topology(),
                barrier_semaphore=self.tt_ccl.get_and_cycle_barrier_semaphore_handle(cluster_axis),
                chunks_per_sync=10,
                num_workers_per_link=2,
                num_buffers_per_channel=2,
                subdevice_id=self.prefetcher.worker_sub_device_id if self.prefetcher is not None else None,
            )

        tt_logits = ttnn.untilize(
            tt_logits,
            use_multicore=True,
            memory_config=ttnn.DRAM_MEMORY_CONFIG,
            sub_core_grids=self.prefetcher.all_worker_cores_range_set if self.prefetcher is not None else None,
        )

        return tt_logits, None

    def switch_mode(self, mode: Mode):
        if self.prefetcher is not None:
            self.prefetcher.init(mode)
            self.prefetcher.prefetch()

    def forward(
        self,
        x: ttnn.Tensor,
        current_pos,
        rot_mats_global=None,
        rot_mats_local=None,
        user_id=0,
        mode: Mode = Mode.DECODE,
        page_table=None,
        chunk_page_table=None,
        chunk_start_idx=None,
        get_last_token=-1,
        kv_cache=None,
        batch_size=1,
        page_tables_per_layer=None,
    ):
        if mode == Mode.DECODE:
            # Run prefetcher if it is enabled
            if self.prefetcher is not None:
                self.prefetcher.run()

        if mode == Mode.PREFILL:
            # For traced prefill, keep RoPE slicing in-graph and driven by the
            # on-device chunk_start_idx input. Batched prefill arrives flattened
            # to [1, 1, batch_size * S_per_user, dim] and each TransformerBlock
            # restores the batch dimension before attention, so the RoPE slice
            # width is the per-user length, not the flattened one.
            prefill_seq_len = x.shape[2] // batch_size
            rot_mats_global = self._slice_prefill_rot_mats(rot_mats_global, chunk_start_idx, prefill_seq_len)
            if rot_mats_local is not None:
                rot_mats_local = self._slice_prefill_rot_mats(rot_mats_local, chunk_start_idx, prefill_seq_len)

        if page_tables_per_layer is not None and len(page_tables_per_layer) != len(self.layers):
            raise ValueError(
                f"page_tables_per_layer has {len(page_tables_per_layer)} entries "
                f"but model has {len(self.layers)} layers"
            )

        for i, layer in enumerate(self.layers):
            # No-op if callers already provide the right memory config
            activation_dtype = self.args.decoders_optimizations.get_tensor_dtype(
                decoder_id=i, tensor=TensorGroup.ACTIVATION
            )

            if mode == Mode.DECODE and not self.args.is_galaxy:
                x = ttnn.to_memory_config(
                    x,
                    self.args.get_residual_mem_config(mode, self.prefetcher),
                    activation_dtype,
                )
            elif activation_dtype is not None and x.dtype != activation_dtype:
                x = ttnn.typecast(x, activation_dtype)

            # vLLM hybrid kv-cache-groups: each attention layer gets its own
            # paged pool (sliding-window vs full-attention have different
            # block counts). When ``page_tables_per_layer`` is None we fall
            # back to broadcasting the single ``page_table`` to every layer
            # — byte-equivalent to the pre-hybrid path used by every legacy
            # caller (demos, unit tests, non-hybrid vLLM bridges).
            layer_page_table = page_tables_per_layer[i] if page_tables_per_layer is not None else page_table

            x = layer(
                x,
                current_pos,
                rot_mats_global=rot_mats_global,
                rot_mats_local=rot_mats_local,
                user_id=user_id,
                mode=mode,
                page_table=layer_page_table,
                chunk_page_table=chunk_page_table,
                chunk_start_idx=chunk_start_idx,
                kv_cache=kv_cache[i] if kv_cache is not None else None,
                batch_size=batch_size,
            )

        if mode == Mode.DECODE:
            if self.prefetcher is not None:
                self.prefetcher.stop()

        if mode == Mode.PREFILL and get_last_token == -1:
            return x

        # Slicing the tensor to the nearest ceiling/floor multiples of 32 for the prefill_len, to get the last token
        if get_last_token != -1:
            seq_len = int(x.shape[2])
            if seq_len % 32 == 0:
                # Runtime bounds, as in process_logits_after_prefill_trace: with literal bounds every
                # distinct prompt offset compiles its own slice program. Untraced-prefill models
                # (Gemma-3) reach this slice on every real request, after warmup has recorded their
                # decode traces, and warmup only ever sees bucket-length mock prompts - measured on
                # Gemma-3-27B DP-4 as the last 8 buffers left live across trace replays.
                for device_tensor, values in (
                    (self._tail_slice_start, [0, 0, get_last_token, 0]),
                    (self._tail_slice_end, [1, 1, get_last_token + 32, int(x.shape[-1])]),
                ):
                    ttnn.copy_host_to_device_tensor(
                        ttnn.from_torch(
                            torch.tensor(values, dtype=torch.int32),
                            mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
                        ),
                        device_tensor,
                    )
                x = ttnn.slice(
                    input_tensor=x,
                    starts=self._tail_slice_start,
                    ends=self._tail_slice_end,
                    slice_dim=2,
                    num_devices=seq_len // 32,
                )
            else:
                x = ttnn.slice(x, (0, 0, get_last_token, 0), (1, 1, get_last_token + 32, x.shape[-1]))

        # Output norm
        x = self.norm(x, mode=mode, norm_config=self.args.get_norm_config("lm_head", mode, self.prefetcher))

        lm_head_input_mem_cfg = self.args.get_lm_head_input_mem_config(
            mode, None if mode == Mode.PREFILL else self.prefetcher
        )
        if mode == Mode.PREFILL and lm_head_input_mem_cfg.is_sharded():
            x = ttnn.interleaved_to_sharded(x, lm_head_input_mem_cfg)
        if mode == Mode.DECODE and self.prefetcher is not None:
            x = ttnn.to_memory_config(x, self.args.get_lm_head_input_mem_config(mode, self.prefetcher))

        x = self.lm_head(x)
        x = self._apply_final_logit_softcapping(x)
        if mode == Mode.PREFILL:
            x = ttnn.to_memory_config(x, memory_config=ttnn.DRAM_MEMORY_CONFIG)

        return x