File size: 41,806 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
# SPDX-FileCopyrightText: © 2024 Tenstorrent USA, Inc.

# SPDX-License-Identifier: Apache-2.0

import math
import os
import re
from enum import Enum
from types import SimpleNamespace
from typing import List, Optional, Union

import torch
from loguru import logger
from PIL import Image as PIL_Image
from pydantic import AliasChoices, BaseModel, ConfigDict, Field

import ttnn
from models.common.tensor_utils import get_rot_transformation_mat as get_rot_transformation_mat_v2


class URL(BaseModel):
    uri: str

    def __str__(self) -> str:
        return self.uri


class ImageMedia(BaseModel):
    image: Union[PIL_Image.Image, URL]

    model_config = ConfigDict(arbitrary_types_allowed=True)


class Role(Enum):
    system = "system"
    user = "user"
    assistant = "assistant"
    ipython = "ipython"


InterleavedTextMedia = Union[
    str,
    # Specific modalities can be placed here, but not generic attachments
    # since models don't consume them in a generic way
    ImageMedia,
    List[Union[str, ImageMedia]],
]


class Mode(Enum):
    DECODE = "decode"
    PREFILL = "prefill"


class HostEmbedding(torch.nn.Module):
    def __init__(self, model_args):
        super().__init__()
        self.emb = torch.nn.Embedding(model_args.vocab_size, model_args.dim)

    def forward(self, x):
        return self.emb(x)


class HostScaledEmbedding(HostEmbedding):
    def __init__(self, model_args):
        super().__init__(model_args)
        self.embed_scale = model_args.embed_scale

    def forward(self, x):
        return self.emb(x) * self.embed_scale


# Default configuration for Paged Attention
class PagedAttentionConfig:
    def __init__(self, block_size=32, max_num_blocks=1024):
        self.block_size = block_size
        self.max_num_blocks = max_num_blocks


class RopeScalingType(str, Enum):
    """Types of RoPE scaling."""

    # DYNAMIC = "dynamic"
    LINEAR = "linear"
    YARN = "yarn"
    LLAMA3 = "llama3"
    PHI3 = "longrope"
    DEFAULT = "default"


class RopeScaling(BaseModel):
    """RoPE scaling configuration."""

    rope_type: RopeScalingType = Field(
        validation_alias=AliasChoices("rope_type", "type"), exclude=True, description="RoPE scaling type"
    )
    factor: Optional[float] = None
    original_max_position_embeddings: Optional[int] = None


class RopeScalingLinear(RopeScaling):
    """RoPE scaling configuration for linear."""


class RopeScalingLlama3(RopeScaling):
    """RoPE scaling configuration for Llama-3.x."""

    # Llama-3.x specific parameters
    low_freq_factor: Optional[float] = 1.0
    high_freq_factor: Optional[float] = 4.0


class RopeScalingYarn(RopeScaling):
    """RoPE scaling configuration for Yarn."""

    # Yarn-specific parameters
    beta_fast: Optional[float] = 32.0
    beta_slow: Optional[float] = 1.0
    mscale: Optional[float] = 1.0
    mscale_all_dim: Optional[float] = 0.0
    truncate: Optional[bool] = True  # Whether to truncate the correction range (floor/ceil)


class RopeScalingPhi3(RopeScaling):
    """RoPE scaling configuration for Phi3."""

    # Phi3-specific parameters
    long_factor: Optional[list]
    short_factor: Optional[list]


def rope_scaling_model_factory(
    rope_scaling_params: dict, original_max_context_len: Optional[int] = None
) -> RopeScaling:
    rope_scaling_type = rope_scaling_params.get("rope_type") or rope_scaling_params.get("type")
    if rope_scaling_type == RopeScalingType.LINEAR:
        return RopeScalingLinear(**rope_scaling_params)
    elif rope_scaling_type == RopeScalingType.LLAMA3:
        return RopeScalingLlama3(**rope_scaling_params)
    elif rope_scaling_type == RopeScalingType.YARN:
        return RopeScalingYarn(**rope_scaling_params)
    elif rope_scaling_type == RopeScalingType.PHI3:
        # transformers 5.x includes original_max_position_embeddings in the rope dict,
        # which collides with the explicit kwarg; merge so the caller value wins and the
        # key is only passed once.
        phi3_params = dict(rope_scaling_params)
        if original_max_context_len is not None:
            phi3_params["original_max_position_embeddings"] = original_max_context_len
        return RopeScalingPhi3(**phi3_params)
    elif rope_scaling_type in ["default", "mrope"]:
        logger.warning(
            f"Rope scaling type was set to {rope_scaling_type}, defaulting to no rope scaling as this rope type is not supported yet by TTT"
        )
        return None
    else:
        raise ValueError(f"Unexpected RoPE scaling type: {rope_scaling_type}")


# transformers 5.x consolidated the RoPE config: the top-level `rope_theta` /
# `rope_local_base_freq` / `rope_scaling` keys were replaced by a single nested
# `rope_parameters` dict (flat for Qwen/Llama; per-attention-type sub-dicts —
# `full_attention` / `sliding_attention` — for Gemma-style models). The helpers
# below read from either layout so configs from transformers <5 and >=5 work.
def get_rope_theta(config: dict, default=None):
    """RoPE base period (global / full-attention)."""
    if config.get("rope_theta") is not None:
        return config["rope_theta"]
    rope_parameters = config.get("rope_parameters") or {}
    if rope_parameters.get("rope_theta") is not None:  # flat (Qwen/Llama)
        return rope_parameters["rope_theta"]
    return (rope_parameters.get("full_attention") or {}).get("rope_theta", default)  # Gemma-style


def get_rope_local_base_freq(config: dict, default=None):
    """Gemma sliding-window local RoPE base (was top-level `rope_local_base_freq`)."""
    if config.get("rope_local_base_freq") is not None:
        return config["rope_local_base_freq"]
    rope_parameters = config.get("rope_parameters") or {}
    return (rope_parameters.get("sliding_attention") or {}).get("rope_theta", default)


def get_rope_scaling(config: dict):
    """RoPE scaling params (factor, original_max_position_embeddings, rope_type, ...).

    transformers <5 put these under `rope_scaling`; >=5 merges them into
    `rope_parameters` (flat, or `full_attention` for Gemma-style). Returns the
    holding dict, or None when no non-default scaling is configured.
    """
    rope_scaling = config.get("rope_scaling")
    if rope_scaling:
        return rope_scaling
    rope_parameters = config.get("rope_parameters") or {}
    if "full_attention" in rope_parameters:  # Gemma-style nesting
        rope_parameters = rope_parameters.get("full_attention") or {}
    # Only a non-default rope_type carries scaling (factor, etc.).
    if rope_parameters.get("rope_type") not in (None, "default"):
        return rope_parameters
    return None


# Minimal addition for Mistral vision support
def position_ids_in_meshgrid_tt(tt_patch_embeds_list, max_width, device):
    position_ids_tt = []
    for tt_patch in tt_patch_embeds_list:
        shape = tt_patch.shape
        height, width = shape[-2], shape[-1]
        mesh = torch.meshgrid(torch.arange(height), torch.arange(width), indexing="ij")
        h_grid, v_grid = torch.stack(mesh, dim=-1).reshape(-1, 2).chunk(2, -1)
        ids = h_grid * max_width + v_grid

        tt_ids = ttnn.from_torch(
            ids,
            device=device,
            dtype=ttnn.uint32,
            layout=ttnn.ROW_MAJOR_LAYOUT,
            memory_config=ttnn.DRAM_MEMORY_CONFIG,
        )
        position_ids_tt.append(tt_ids[:, 0])
    return ttnn.concat(position_ids_tt, dim=0)


def encode_prompt_instruct(tokenizer, prompt_text, system_prompt_text=None):
    """<|begin_of_text|><|start_header_id|>system<|end_header_id|>
    {{ system_prompt }}<|eot_id|><|start_header_id|>user<|end_header_id|>
    {{ user_msg_1 }}<|eot_id|><|start_header_id|>assistant<|end_header_id|>
    {{ model_answer_1 }}<|eot_id|>
    """
    begin_of_text = [tokenizer.special_tokens["<|begin_of_text|>"]]
    start_header = [tokenizer.special_tokens["<|start_header_id|>"]]
    end_header = [tokenizer.special_tokens["<|end_header_id|>"]]
    end_turn = [tokenizer.special_tokens["<|eot_id|>"]]
    system = tokenizer.encode("system", bos=False, eos=False)
    user = tokenizer.encode("user", bos=False, eos=False)
    assistant = tokenizer.encode("assistant", bos=False, eos=False)
    prompt = tokenizer.encode(prompt_text, bos=False, eos=False)

    system_prompt = start_header + system + end_header + system_prompt_text + end_turn if system_prompt_text else []
    user_prompt = start_header + user + end_header + prompt + end_turn
    assistant_reply = start_header + assistant + end_header
    return begin_of_text + system_prompt + user_prompt + assistant_reply


def preprocess_inputs_prefill(
    input_prompts,
    tokenizer,
    model_args,
    instruct,
    max_generated_tokens,
    max_prefill_len=128 * 1024,
):
    """
    Run tokenizer on inputs, and create embeddings for the first token of each input
    """
    # To avoid going out of memory, clip the max prefill length by the maximum number of tokens that will be generated

    for m_args in model_args:
        assert (
            max_prefill_len <= m_args.max_context_len
        ), f"max_prefill_len {max_prefill_len} cannot exceed max_context_len {m_args.max_context_len}"

    # we need to make room for the generated tokens in the total token budget
    max_prefill_len -= max_generated_tokens
    assert (
        max_prefill_len > 0
    ), f"max_prefill_len ({max_prefill_len + max_generated_tokens}) must be greater than max_generated_tokens ({max_generated_tokens})"

    encoded_prompts = [
        model_args[idx % len(model_args)].encode_prompt(prompt, instruct=instruct)
        for idx, prompt in enumerate(input_prompts)
    ]

    # Print the length of encoded prompts
    logger.info("Encoded prompt lengths:" + ", ".join(str(len(prompt)) for prompt in encoded_prompts))

    prompt_lens = [len(x) for x in encoded_prompts]
    min_prompt_len = min(prompt_lens)
    max_prompt_len = max(prompt_lens)

    # To avoid running out of memory when giving prompts larger than the maximum, clip to max_prefill_len
    if min_prompt_len > max_prefill_len:
        logger.info(f"Left-clipping prompts to {max_prefill_len}")
        if instruct:
            # We need to allow a few tokens for the system prompt and the special turn tokens for assistant and user;
            # to find out how big those will be, we will:
            # 1. Tokenize the entire prompt with non-instruct tokenization
            # 2. Calculate overhead = length of instruct tokenization - length of non-instruct tokenization
            # 3. Shorten the tokenized clipped prompt by the overhead and convert back to text
            # 4. Tokenize the result with instruct tokenization
            # 5. Assert that the length of this is equal to the max_prefill_len
            raw_prompts = [
                model_args[idx % len(model_args)].encode_prompt(prompt, instruct=False)
                for idx, prompt in enumerate(input_prompts)
            ]
            overhead = [len(e) - len(r) for e, r in zip(encoded_prompts, raw_prompts)]

            shortened = []
            for idx, (e, o) in enumerate(zip(raw_prompts, overhead)):
                if isinstance(tokenizer, list):
                    sp = tokenizer[idx % len(model_args)].decode(e[-(max_prefill_len - o) :])
                else:
                    sp = tokenizer.decode(e[-(max_prefill_len - o) :])
                shortened.append(sp)

            encoded_prompts = [
                model_args[idx % len(model_args)].encode_prompt(prompt, instruct=instruct)
                for idx, prompt in enumerate(shortened)
            ]
            # Instruct re-tokenization can drift by a few tokens vs the overhead
            # estimate (seen on Gemma4-26B-A4B: 65337 vs 65336). Re-trim / accept
            # slightly-short prompts rather than hard-failing the demo.
            trimmed = []
            for e in encoded_prompts:
                if len(e) > max_prefill_len:
                    e = e[-max_prefill_len:]
                trimmed.append(e)
            encoded_prompts = trimmed
            lens = [len(e) for e in encoded_prompts]
            assert all(
                0 < n <= max_prefill_len for n in lens
            ), f"Clipped prompts are not of the correct length, expected <= {max_prefill_len} but got {lens}"
            if any(n != max_prefill_len for n in lens):
                logger.warning(
                    f"Instruct re-clip lengths {lens} != target {max_prefill_len}; "
                    f"continuing with trimmed/short prompts"
                )
        else:
            encoded_prompts = [encod[-max_prefill_len:] for encod in encoded_prompts]

        # Update prompt lengths
        prompt_lens = [len(x) for x in encoded_prompts]
        min_prompt_len = min(prompt_lens)
        max_prompt_len = max(prompt_lens)
    for m in model_args:
        assert (
            max_prompt_len <= m.max_seq_len
        ), f"Max prompt length {max_prompt_len} exceeds model max seq len {m.max_seq_len}"
    assert min_prompt_len > 0, "Minimum prompt length must be greater than 0"
    assert min_prompt_len <= max_prompt_len, f"Minimum prompt length {min_prompt_len} exceeds max len {max_prompt_len}"

    logger.info(f"# of users: {len(encoded_prompts)}")
    input_tokens_prefill = []
    decoding_pos = []
    prefill_lens = []

    # Pad each prompt to the maximum length among all prompts.
    # To avoid issues, we keep track of the decoding position to decode correctly the user's prompt
    for i, encoded in enumerate(encoded_prompts):
        # Initial prefill tensors full of pad tokens
        input_tokens_prefill_i = torch.full((1, max_prompt_len), 0, dtype=torch.int32)
        input_tokens_prefill_i[0, : len(encoded[:])] = torch.tensor(encoded[:]).to(input_tokens_prefill_i)
        input_tokens_prefill.append(input_tokens_prefill_i)

        # Keep the correct decoding position of each user
        decoding_pos.append(len(encoded))
        prefill_lens.append(max_prompt_len)

    return (
        input_tokens_prefill,
        encoded_prompts,
        decoding_pos,
        prefill_lens,
    )


def _chat_template_ids(encoded):
    """Normalize apply_chat_template(tokenize=True) output to a flat List[int].

    transformers <5 returned a plain List[int]; transformers 5.x defaults
    apply_chat_template to ``return_dict=True`` and returns a ``BatchEncoding``
    (a ``UserDict`` — NOT a ``dict`` subclass, so ``isinstance(x, dict)`` is
    False), or a `tokenizers.Encoding` (exposes ``.ids``). Iterating a
    ``BatchEncoding``/``UserDict`` yields its *keys* ("input_ids", ...), so we
    must extract ``input_ids`` via mapping membership rather than ``isinstance``.
    """
    # dict / BatchEncoding / UserDict — use mapping membership, since BatchEncoding
    # is a UserDict and fails isinstance(x, dict).
    if hasattr(encoded, "keys") and "input_ids" in encoded:
        encoded = encoded["input_ids"]
    if hasattr(encoded, "ids"):  # tokenizers.Encoding
        return list(encoded.ids)
    if hasattr(encoded, "tolist"):  # torch tensor / np array
        encoded = encoded.tolist()
    # apply_chat_template(return_dict=True) on a single conversation can nest the
    # ids in a 1-element batch dim ([[ids]]); unwrap it.
    if isinstance(encoded, (list, tuple)) and len(encoded) == 1 and isinstance(encoded[0], (list, tuple)):
        encoded = encoded[0]
    return list(encoded)  # already a List[int]


def encode_prompt_hf(tokenizer, prompt_text, system_prompt_text=None):
    """See https://huggingface.co/docs/transformers/main/en/chat_templating"""
    chat = []
    if isinstance(prompt_text, str):
        if system_prompt_text:
            chat.append({"role": "system", "content": system_prompt_text})
        if prompt_text:
            chat.append({"role": "user", "content": prompt_text})
        encoded = tokenizer.apply_chat_template(chat, add_generation_prompt=True, tokenize=True)
    else:
        encoded = tokenizer.apply_chat_template(prompt_text, add_generation_prompt=True, tokenize=True)
    return _chat_template_ids(encoded)


def compute_llama3_parameters(freqs: torch.Tensor, scale_factor: float, orig_context_len: int):
    """Llama-3.x specific scaling for rotary embeddings."""
    low_freq_factor = 1
    high_freq_factor = 4

    low_freq_wavelen = orig_context_len / low_freq_factor
    high_freq_wavelen = orig_context_len / high_freq_factor
    new_freqs = []
    for freq in freqs:
        wavelen = 2 * math.pi / freq
        if wavelen < high_freq_wavelen:
            new_freqs.append(freq)
        elif wavelen > low_freq_wavelen:
            new_freqs.append(freq / scale_factor)
        else:
            assert low_freq_wavelen != high_freq_wavelen
            smooth = (orig_context_len / wavelen - low_freq_factor) / (high_freq_factor - low_freq_factor)
            new_freqs.append((1 - smooth) * freq / scale_factor + smooth * freq)
    return torch.tensor(new_freqs, dtype=freqs.dtype, device=freqs.device)


def compute_linear_parameters(freqs: torch.Tensor, scale_factor: float, orig_context_len: int):
    """Linear scaling for rotary embeddings."""
    freqs /= scale_factor
    return freqs


def compute_default_parameters(freqs: torch.Tensor, scale_factor: float, orig_context_len: int):
    """Default scaling for rotary embeddings."""
    return freqs


def apply_scaling(freqs: torch.Tensor, scale_factor: float, orig_context_len: int, rope_type="llama3"):
    # FIXME: Llama-3.x specific scaling - we need to support yarn for Qwen2.5 models

    if rope_type == "default":
        freqs = compute_default_parameters(freqs, scale_factor, orig_context_len)
    elif rope_type == "linear":
        freqs = compute_linear_parameters(freqs, scale_factor, orig_context_len)
    elif rope_type == "llama3":
        freqs = compute_llama3_parameters(freqs, scale_factor, orig_context_len)

    return freqs


# Minimal addition for Mistral vision RoPE support
def apply_scaling_vision(freqs: torch.Tensor, scale_factor: float, orig_context_len: int):
    return freqs / scale_factor


# Minimal addition for Mistral vision RoPE support
def precompute_mistral_vision_freqs(
    dim: int, max_patches_per_side: int, theta: float, scale_factor=None, orig_context_len=None
):
    # Compute base frequencies
    base_freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim))
    if scale_factor is not None:
        base_freqs = apply_scaling_vision(base_freqs, scale_factor, orig_context_len)

    # Get height and width indices
    h_idx = torch.arange(max_patches_per_side)
    w_idx = torch.arange(max_patches_per_side)

    # Compute 2D frequency matrices
    freqs_h = torch.outer(h_idx, base_freqs[::2])
    freqs_w = torch.outer(w_idx, base_freqs[1::2])

    # Broadcast + merge
    inv_freq = torch.cat(
        [
            freqs_h[:, None, :].repeat(1, max_patches_per_side, 1),
            freqs_w[None, :, :].repeat(max_patches_per_side, 1, 1),
        ],
        dim=-1,
    ).reshape(
        -1, dim // 2
    )  # Shape: [H*W, dim//2]

    full_freqs = torch.cat([inv_freq, inv_freq], dim=-1)
    cos = full_freqs.cos()
    sin = full_freqs.sin()
    return cos, sin  # Shape: [H*W, dim]


def precompute_freqs(dim: int, end: int, theta, scale_factor, orig_context_len, rope_type="llama3"):
    """
    Precompute the frequency tensor for sine and cosine values with given dimensions.

    Args:
        dim (int): Dimension of the frequency tensor.
        end (int): End index for precomputing frequencies.
        theta (float, optional): Scaling factor for frequency computation. Defaults to 500000.0.

    Returns:
        Tuple[torch.Tensor, torch.Tensor]: Tensors containing cosine and sine values.
    """
    freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))
    t = torch.arange(end)
    if scale_factor is not None:
        freqs = apply_scaling(freqs, scale_factor, orig_context_len, rope_type=rope_type)
    freqs = torch.outer(t, freqs).float()
    return torch.cos(freqs), torch.sin(freqs)


def freqs_to_rotation_matrix(cos_freqs, sin_freqs):
    """
    Transform cos/sin frequencies to a rotation matrix.
    """
    emb_size, emb_dim = cos_freqs.shape
    dhead = emb_dim * 2
    rot_emb_matrix = torch.zeros(emb_size, dhead, dhead)
    rot_emb_matrix[..., torch.arange(0, dhead, 2), torch.arange(0, dhead, 2)] = cos_freqs.clone()
    rot_emb_matrix[..., torch.arange(1, dhead, 2), torch.arange(1, dhead, 2)] = cos_freqs.clone()
    rot_emb_matrix[..., torch.arange(0, dhead, 2), torch.arange(1, dhead, 2)] = -sin_freqs.clone()
    rot_emb_matrix[..., torch.arange(1, dhead, 2), torch.arange(0, dhead, 2)] = sin_freqs.clone()

    rot_emb_matrix = rot_emb_matrix.transpose(-1, -2)  # Necessary for correct rotation when applied as (x @ R)
    return rot_emb_matrix


def gather_cos_sin(position_ids, cos, sin):
    position_id_expanded = position_ids.unsqueeze(1).expand(-1, cos.shape[-1])
    cos = cos.gather(0, position_id_expanded)
    sin = sin.gather(0, position_id_expanded)
    cos = torch.stack([cos, cos], dim=-1).flatten(-2).unsqueeze(0).unsqueeze(0)
    sin = torch.stack([sin, sin], dim=-1).flatten(-2).unsqueeze(0).unsqueeze(0)
    return cos, sin


def get_prefill_rot_mat(head_dim, mesh_device, seq_len, theta, scale_factor, orig_context_len, start_pos=0):
    cos, sin = precompute_freqs(
        head_dim, seq_len * 2, theta=theta, scale_factor=scale_factor, orig_context_len=orig_context_len
    )
    cos_gathered, sin_gathered = gather_cos_sin(torch.arange(start_pos, start_pos + seq_len), cos, sin)
    assert cos_gathered.size() == (1, 1, seq_len, head_dim)
    assert sin_gathered.size() == (1, 1, seq_len, head_dim)

    cos_gathereds = ttnn.from_torch(
        cos_gathered,
        dtype=ttnn.bfloat16,
        layout=ttnn.TILE_LAYOUT,
        device=mesh_device,
        mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device),
    )
    sin_gathereds = ttnn.from_torch(
        sin_gathered,
        dtype=ttnn.bfloat16,
        layout=ttnn.TILE_LAYOUT,
        device=mesh_device,
        mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device),
    )

    rot_mats = [cos_gathereds, sin_gathereds]
    return rot_mats


#  Add-Multiply method of rotary embeddings for prefill
def get_rot_transformation_mat(dhead=32):
    # ROPE op uses a single tile
    dhead = 32
    # Delegate to TTTv2 implementation for consistency
    return get_rot_transformation_mat_v2(dhead)


def get_single_rot_mat(
    dhead,
    mesh_device,
    num_devices,
    start_pos,
    theta,
    scale_factor,
    orig_context_len,
    on_host=False,
):
    freqs_unscaled = 1.0 / (theta ** (torch.arange(0, dhead, 2)[: (dhead // 2)].float() / dhead))
    if scale_factor is not None:
        freqs = apply_scaling(freqs_unscaled, scale_factor, orig_context_len, rope_type="llama3")
    rot_matrix = torch.zeros(dhead, dhead)
    # [INFO] freqs_unscaled and freqs are forced to float dtype above and it should be converted back to match dtype of rot_matrix
    sin_freqs, cos_freqs = torch.sin(freqs).to(rot_matrix.dtype), torch.cos(freqs).to(rot_matrix.dtype)
    rot_matrix[torch.arange(0, dhead, 2), torch.arange(0, dhead, 2)] = cos_freqs.clone()
    rot_matrix[torch.arange(1, dhead, 2), torch.arange(1, dhead, 2)] = cos_freqs.clone()
    rot_matrix[torch.arange(0, dhead, 2), torch.arange(1, dhead, 2)] = -sin_freqs.clone()
    rot_matrix[torch.arange(1, dhead, 2), torch.arange(0, dhead, 2)] = sin_freqs.clone()
    rot_matrix = rot_matrix.transpose(-1, -2)

    # Support for start_pos different than 0
    freqs = start_pos * freqs_unscaled
    if scale_factor is not None:
        freqs = apply_scaling(freqs, scale_factor, orig_context_len, rope_type="llama3")
    current_rot_mat = torch.zeros(dhead, dhead)
    # [INFO] freqs_unscaled and freqs are forced to float dtype above and it should be converted back to match dtype of current_rot_mat
    sin_freqs, cos_freqs = torch.sin(freqs).to(current_rot_mat.dtype), torch.cos(freqs).to(current_rot_mat.dtype)
    current_rot_mat[torch.arange(0, dhead, 2), torch.arange(0, dhead, 2)] = cos_freqs.clone()
    current_rot_mat[torch.arange(1, dhead, 2), torch.arange(1, dhead, 2)] = cos_freqs.clone()
    current_rot_mat[torch.arange(0, dhead, 2), torch.arange(1, dhead, 2)] = -sin_freqs.clone()
    current_rot_mat[torch.arange(1, dhead, 2), torch.arange(0, dhead, 2)] = sin_freqs.clone()

    return ttnn.from_torch(
        current_rot_mat.T.unsqueeze(0).unsqueeze(0),  # 1,1,head_dim,head_dim
        device=mesh_device if not on_host else None,
        dtype=ttnn.bfloat16,
        layout=ttnn.TILE_LAYOUT,
        mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device) if num_devices > 1 or not on_host else None,
    ), ttnn.from_torch(
        rot_matrix.unsqueeze(0).unsqueeze(0),  # 1,1,head_dim,head_dim
        device=mesh_device if not on_host else None,
        dtype=ttnn.bfloat16,
        layout=ttnn.TILE_LAYOUT,
        mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device) if num_devices > 1 or not on_host else None,
    )


def num_to_core_range_set(x):
    assert x < 8 or x % 8 == 0
    num_x = min(x, 8)
    num_y = x // num_x
    assert num_x * num_y == x
    return ttnn.CoreRangeSet(
        {
            ttnn.CoreRange(
                ttnn.CoreCoord(0, 0),
                ttnn.CoreCoord(num_x - 1, num_y - 1),
            ),
        }
    )


def copy_host_to_device(
    host_tensors,
    device_tensors=None,
    mesh_device=None,
    shard_specs=None,
):
    """
    Helper function which copies host tensors to device tensors.
    If no device_tensors are provided, it creates new device tensors and returns them.
    """
    if device_tensors is None:
        assert mesh_device is not None, "mesh_device is required when device_tensors is None"
        ret = []
        for i in range(len(host_tensors)):
            if shard_specs and shard_specs[i] is not None:
                on_device = host_tensors[i].to(mesh_device, shard_specs[i]) if host_tensors[i] else None
            else:
                on_device = ttnn.to_device(host_tensors[i], device=mesh_device) if host_tensors[i] else None
            ret.append(on_device)
        return ret
    else:
        for i in range(len(host_tensors)):
            if host_tensors[i] is None:
                assert device_tensors[i] is None
                continue
            ttnn.copy_host_to_device_tensor(host_tensors[i], device_tensors[i])
        return device_tensors


def calculate_hidden_dim(dim, ffn_dim_multiplier, multiple_of):
    """Helper function based on logic used in reference model:
    https://github.com/meta-llama/llama-models/blob/e4a6ed52a142bb9b5106dcbf48e41f97f8e7378e/models/llama3/reference_impl/model.py#L227C7-L231C83
    """
    hidden_dim = int(2 * (4 * dim) / 3)
    if ffn_dim_multiplier is not None:
        hidden_dim = int(ffn_dim_multiplier * hidden_dim)
    hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
    return hidden_dim


def get_out_subblock_w(per_core_N, out_subblock_h):
    """
    Helper function to calculate the out_subblock_w based on the per_core_N and out_subblock_h
    """
    out_subblock_w = 4  # TODO: Check with LLK team if this is the true bound, might be 8 now
    while out_subblock_w > 1:
        if out_subblock_w * out_subblock_h <= 4 and per_core_N % out_subblock_w == 0:
            break
        out_subblock_w -= 1
    return out_subblock_w


def first_five(tensor, mesh_device, start=0, end=5):
    """
    Helper function to return the first 5 elements of a tensor via torch, or optionally another slice
    """
    return torch.Tensor(ttnn.to_torch(tensor, mesh_composer=ttnn.ConcatMeshToTensor(mesh_device, dim=-1)))[
        0, 0, 0, start:end
    ]


def last_five(tensor, mesh_device):
    """
    Helper function to return the last 5 elements of a tensor via torch
    """
    return torch.Tensor(ttnn.to_torch(tensor, mesh_composer=ttnn.ConcatMeshToTensor(mesh_device, dim=-1)))[0, 0, 0, -5:]


# Sample logits from a distribution
def sample_top_p(probs: torch.Tensor, p: float):
    assert 0 <= p <= 1

    probs_sort, probs_idx = torch.sort(probs, dim=-1, descending=True)
    probs_sum = torch.cumsum(probs_sort, dim=-1)
    mask = probs_sum - probs_sort > p
    probs_sort[mask] = 0.0
    probs_sort.div_(probs_sort.sum(dim=-1, keepdim=True))

    next_token = torch.multinomial(probs_sort, num_samples=1)
    return torch.gather(probs_idx, -1, next_token)


def sample_host(tt_input, temperature=0.6, top_p=0.08, on_host=True):
    vocab_size = tt_input.shape[-1]
    pt_input = tt_input[..., :vocab_size]

    if temperature > 0:
        probs = torch.softmax(pt_input / temperature, dim=-1)
        pt_out = sample_top_p(probs.squeeze(), top_p)
    else:
        pt_out = torch.argmax(pt_input, dim=-1)

    if pt_out.dim() == 1:  # if sampling a single token re-add the batch dim to the tensor
        pt_out = pt_out.unsqueeze(0)
    return None, pt_out


def get_padded_prefill_len(seq_len: int) -> int:
    """
    Get the padded prefill length for a given sequence length.
    This is used to pad the sequence length to the nearest power of 2.
    """
    # TODO: https://github.com/tenstorrent/tt-metal/issues/34117
    if seq_len <= 128:
        return 128
    if seq_len <= 1024:
        return 1024
    else:
        # return next power of 2 greater than seq_len
        return 2 ** (seq_len - 1).bit_length()


def get_all_padded_prefill_lengths(max_len):
    lengths = [128]
    k = 0
    while (v := (1 << k) * 1024) <= max_len:
        lengths.append(v)
        k += 1
    return lengths


def calculate_prefill_warmup_seq_lens(max_seq_len_to_warmup, trace_supported_seq_lens):
    to_warmup_seq_lens = get_all_padded_prefill_lengths(max_seq_len_to_warmup)
    for trace_supported_seq_len in trace_supported_seq_lens:
        if trace_supported_seq_len not in to_warmup_seq_lens:
            to_warmup_seq_lens.append(trace_supported_seq_len)
    to_warmup_seq_lens.sort()

    return to_warmup_seq_lens


def cap_seq_lens_to_max_prefill_chunk_size(seq_lens, cap):
    for seq_len in seq_lens:
        if seq_len > cap:
            seq_lens = seq_lens[: seq_lens.index(seq_len)]
            break
    return seq_lens


def get_block_size(kv_cache):
    return kv_cache[0][0].shape[2]


def num_blocks_in_seq(seq_len, block_size):
    return math.ceil(seq_len / block_size)


def nearest_pow_2(x):
    return 2 ** math.ceil(math.log2(x))


def get_max_prefill_chunk_size(seq_len, max_prefill_seq_len):
    """
    Determine the largest multiple of 2048 that divides `seq_len` and is less than or equal to `max_prefill_seq_len`.

    **Assumptions**:
    - `seq_len` is a multiple of 2048.
    - `max_prefill_seq_len` is a multiple of 2048.
    """
    MIN_CHUNK_SIZE = 2048

    if not isinstance(seq_len, int) or not isinstance(max_prefill_seq_len, int):
        raise TypeError("Both seq_len and max_prefill_seq_len must be integers.")
    if seq_len <= 0 or max_prefill_seq_len <= 0:
        raise ValueError("Both seq_len and max_prefill_seq_len must be positive integers.")

    if seq_len % MIN_CHUNK_SIZE != 0:
        raise ValueError(f"seq_len ({seq_len}) must be a multiple of {MIN_CHUNK_SIZE}.")
    if max_prefill_seq_len % MIN_CHUNK_SIZE != 0:
        raise ValueError(f"max_prefill_seq_len ({max_prefill_seq_len}) must be a multiple of {MIN_CHUNK_SIZE}.")

    # Calculate the maximum possible chunk size
    # It cannot exceed either max_prefill_seq_len or seq_len
    max_possible_chunk = min(max_prefill_seq_len, seq_len)

    # Iterate from the largest possible multiple of MIN_CHUNK_SIZE down to MIN_CHUNK_SIZE
    for chunk_size in range(max_possible_chunk, 0, -MIN_CHUNK_SIZE):
        if seq_len % chunk_size == 0:
            return chunk_size

    raise ValueError("No valid chunk size found")


def nearest_multiple(x, multiple_of):
    return math.ceil(x / multiple_of) * multiple_of


def pad_to_size(x: torch.Tensor, dim: int, size: int) -> torch.Tensor:
    """
    Pads the specified dimension of the input tensor with zeros

    :param x: Input PyTorch Tensor
    :param dim: The dimension to pad
    :param size: The size to pad to
    :return: Padded PyTorch Tensor
    """
    # handle negative dim
    if dim < 0:
        dim = x.dim() + dim
    assert isinstance(x, torch.Tensor), "Input must be a torch.Tensor"
    assert -x.dim() <= dim < x.dim(), f"Dimension {dim} out of range (expected between {-x.dim()} and {x.dim() - 1})"
    dim = x.dim() + dim if dim < 0 else dim

    current_size = x.size(dim)
    pad_size = size - current_size

    if pad_size == 0:
        return x  # No padding needed

    # Prepare the padding configuration for F.pad
    # F.pad expects padding in the form (pad_last_dim_left, pad_last_dim_right, ..., pad_dim_left, pad_dim_right)
    # We only pad on the "end" side of the specified dimension
    pad = [0] * (2 * x.dim())  # Initialize padding for all dimensions
    pad_index = 2 * (x.dim() - dim - 1)
    pad[pad_index + 1] = pad_size  # Pad on the "right" side of the specified dimension

    padded_x = torch.nn.functional.pad(x, pad, mode="constant", value=0)
    return padded_x


def get_base_model_name(model_name: str) -> str:
    # Explicitly handle phi-4 which doesn't follow the <Size>B format
    if "phi-4" in model_name.lower():
        return "Phi-4"
    # Remove the suffix after B- (case insensitive), e.g. "Llama-3.1-70B-Instruct" -> "Llama-3.1-70B"
    match = re.search(r"(.*?\d+[bB])-", model_name)
    return match.group(1) if match else model_name


def get_hf_model_name(model_path: str) -> str:
    # HF model name
    if model_path.count("/") == 1:
        return model_path

    # HF cache path
    pattern = r".*/?models--(?P<model_provider>[^/]+?)--(?P<model_name>[^/]+)/?"
    match = pattern.search(pattern, model_path)
    if match:
        model_provider = match.group("model_provider")
        model_name = match.group("model_name")
        return f"{model_provider}/{model_name}"
    raise ValueError(
        f"Unsupported '{model_path}', please use HF model name or follow HF format with 'models--<model_provider>--<model_name>'"
    )


def get_hf_tt_cache_path(model_path: str) -> str:
    tt_cache_home = os.getenv("TT_CACHE_HOME", "/mnt/MLPerf/huggingface/tt_cache/")
    if not os.path.exists(tt_cache_home):
        tt_cache_home = "model_cache"

    model_name = get_hf_model_name(model_path)
    tt_cache_path = os.path.join(tt_cache_home, model_name)
    if not os.path.exists(tt_cache_path):
        os.makedirs(tt_cache_path, exist_ok=True)

    return tt_cache_path


def create_tt_model(
    mesh_device,
    instruct,
    max_batch_size,
    optimizations,
    max_seq_len,
    paged_attention_config: PagedAttentionConfig = None,
    dtype=ttnn.bfloat8_b,
    state_dict=None,
    num_layers=None,
    use_prefetcher=False,
    use_hf_rope=False,
):
    from models.tt_transformers.tt.model import Transformer
    from models.tt_transformers.tt.model_config import ModelArgs
    from models.tt_transformers.tt.prefetcher import Prefetcher

    num_tensors = 5 if use_prefetcher else 0
    prefetcher = Prefetcher(mesh_device, num_tensors, num_layers) if use_prefetcher else None

    tt_model_args = ModelArgs(
        mesh_device,
        instruct=instruct,
        max_batch_size=max_batch_size,
        optimizations=optimizations,
        max_seq_len=max_seq_len,
        prefetcher=prefetcher,
        use_hf_rope=use_hf_rope,
    )

    if num_layers is not None:
        tt_model_args.n_layers = num_layers

    if prefetcher is not None:
        prefetcher.num_layers = tt_model_args.n_layers

    # Decide whether the HF weights are still needed on host. When the ttnn weight cache for
    # this build was already fully built on a previous run, ttnn.as_tensor loads every weight from
    # disk and the state_dict is never read -- so skip the expensive from_pretrained host load
    # entirely (the load that OOMs/hangs in prefill, #48509). Generalizes GPT-OSS PR #48531 (whose
    # --skip-model-load pytest flag is gpt_oss-only; nothing equivalent exists for these models).
    #
    # state_dict is None  -> decide here (warm cache => placeholder, else cold load).
    # state_dict falsy/{}  -> caller already decided to skip (e.g. a prior DP submesh); build as-is.
    # state_dict populated -> reuse across DP models (avoid reloading for every submesh).
    loaded_real_weights = False
    if state_dict is None:
        if not tt_model_args.dummy_weights and tt_model_args.weight_cache_is_complete(dtype):
            logger.info("Warm ttnn weight cache detected -- skipping HF state_dict load.")
            # Dataless placeholder: every weight is loaded from its .tensorbin by ttnn.as_tensor;
            # the placeholder only satisfies the host-side reshape ops (see placeholder_state_dict).
            state_dict = tt_model_args.placeholder_state_dict(dtype)
        else:
            state_dict = tt_model_args.load_state_dict()
            loaded_real_weights = bool(state_dict) and not tt_model_args.dummy_weights

    # A populated state_dict handed in by the caller (DP submeshes after the first) bypasses
    # load_state_dict(), which is the only place the cold path sets is_mixture_of_experts. Without
    # this the later lanes build a dense MLP for an MoE checkpoint and fail on the missing
    # feed_forward.w1 key. Derive the flag from the keys, as load_state_dict does.
    # (The warm-cache placeholder mapping is deliberately falsy, so test for None, not truthiness.)
    if state_dict is not None and not getattr(tt_model_args, "is_mixture_of_experts", False):
        tt_model_args.is_mixture_of_experts = any(".experts." in k for k in state_dict.keys())
    if getattr(tt_model_args, "is_mixture_of_experts", False):
        # Reused weights must initialize the same MoE configuration as load_state_dict.
        tt_model_args.moe = True
        expert_indices = [
            int(k.split(".experts.")[1].split(".")[0]) + 1 for k in state_dict if "block_sparse_moe.experts." in k
        ]
        tt_model_args.num_experts = max(expert_indices) if expert_indices else tt_model_args.num_local_experts

    model = Transformer(
        args=tt_model_args,
        mesh_device=mesh_device,
        dtype=dtype,
        state_dict=state_dict,
        weight_cache_path=tt_model_args.weight_cache_path(dtype),
        paged_attention_config=paged_attention_config,
        prefetcher=prefetcher,
    )

    # If this run populated the cache from a cold host load, record completion so future runs
    # can skip the load. Only for full-model builds (a num_layers override produces a partial
    # cache that must not satisfy the completeness check).
    if loaded_real_weights and num_layers is None:
        tt_model_args.mark_weight_cache_complete(dtype, state_dict)

    tt_kv_cache = [l.attention.layer_past for l in model.layers] if paged_attention_config else None

    return tt_model_args, model, tt_kv_cache, state_dict


def hf_multimodal_encode(messages, processor):
    hf_messages = []

    for msg in messages:
        hf_content = []

        for item in msg.content:
            if isinstance(item, ImageMedia):
                hf_content.append(
                    {
                        "type": "image",
                        "image": item.image,
                    }
                )
            elif isinstance(item, str):
                hf_content.append(
                    {
                        "type": "text",
                        "text": item,
                    }
                )

        hf_messages.append(
            {
                "role": msg.role,
                "content": hf_content,
            }
        )

    encoded = processor.apply_chat_template(
        hf_messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt"
    ).to("cpu", dtype=torch.bfloat16)

    return SimpleNamespace(
        **encoded,
        tokens=encoded["input_ids"].squeeze(0),
        vision=SimpleNamespace(
            images=encoded.get("pixel_values", None),
            mask=None,
        ),
    )


def get_decode_mask(args, mesh_device, paged_attention_config=None):
    """Function to create a decoding mask for the attention mechanism."""
    if paged_attention_config is not None:
        max_seq_len = (paged_attention_config.max_num_blocks * paged_attention_config.block_size) // args.max_batch_size
    else:
        max_seq_len = args.max_seq_len
    mask = torch.triu(
        torch.full(
            (args.max_batch_size, args.n_heads // mesh_device.shape[1], max_seq_len, max_seq_len),
            -float("inf"),
            dtype=torch.bfloat16,
        ),
        diagonal=1,
    )
    if args.sliding_window > 0:
        mask += torch.tril(
            torch.full(
                (args.max_batch_size, args.n_heads // mesh_device.shape[1], max_seq_len, max_seq_len),
                -float("inf"),
                dtype=torch.bfloat16,
            ),
            diagonal=-args.sliding_window,
        )

    return mask


def build_encoder_attention_mask(
    x: torch.Tensor,
    ar: torch.Tensor,
    ntok: int,
    num_chunks: int,
    n_heads: int,
):
    """
    Build vision encoder attention mask that omits padding tokens.
    """

    def get_negative_inf_value(dtype):
        return torch.finfo(dtype).min

    masks = []
    for arx in ar:
        mask_i = torch.ones((num_chunks, x.shape[2], 1), dtype=x.dtype)
        mask_i[: arx[0] * arx[1], :ntok] = 0
        mask_i = mask_i.view(num_chunks * x.shape[2], -1)
        mask_i = mask_i @ mask_i.T * get_negative_inf_value(x.dtype)
        mask_i = mask_i.unsqueeze(0)
        masks.append(mask_i)
    masks = torch.stack(masks).to(x.device).expand(-1, n_heads, -1, -1)
    return masks