# SPDX-FileCopyrightText: © 2024 Tenstorrent USA, Inc. # SPDX-License-Identifier: Apache-2.0 import hashlib import inspect import json import math import os import re from enum import Enum, auto from functools import lru_cache from pathlib import Path from typing import Tuple import torch from loguru import logger import ttnn from models.common.utility_functions import hf_cache_to_legacy, is_blackhole, is_wormhole_b0, nearest_32 from models.common.weight_cache import WEIGHT_CACHE_FORMAT_VERSION as _WC_FORMAT_VERSION from models.common.weight_cache import WEIGHT_CACHE_MARKER as _WC_MARKER from models.common.weight_cache import mark_weight_cache_complete as _mark_weight_cache_complete from models.common.weight_cache import marker_path as _wc_marker_path from models.common.weight_cache import weight_cache_is_complete as _weight_cache_is_complete from models.tt_transformers.tt.common import ( Mode, calculate_hidden_dim, calculate_prefill_warmup_seq_lens, cap_seq_lens_to_max_prefill_chunk_size, encode_prompt_hf, get_base_model_name, get_out_subblock_w, get_rope_local_base_freq, get_rope_scaling, get_rope_theta, nearest_multiple, num_to_core_range_set, rope_scaling_model_factory, ) from models.tt_transformers.tt.load_checkpoints import convert_vision_meta_to_hf # Minimal addition for Mistral vision from models.tt_transformers.tt.load_checkpoints import ( convert_hf_to_meta, convert_hf_to_meta_mllama, convert_hf_to_meta_mllama_no_qkv_permute, convert_hf_to_meta_no_qkv_permute, convert_meta_to_hf, convert_meta_to_hf_no_qkv_permute, convert_vision_hf_to_meta, convert_vision_hf_to_meta_no_qkv_permute, load_hf_state_dict_for_layers, reverse_permute, standardize_hf_keys, standardize_hf_keys_multimodal, ) from models.tt_transformers.tt.prefetcher import Prefetcher _HF_LAYER_KEY_RE = re.compile(r"^model\.layers\.(\d+)\.") # file names for performance and accuracy mode override files PERFORMANCE_DECODER_CONFIG_FILENAME = "performance_decoder_config.json" ACCURACY_DECODER_CONFIG_FILENAME = "accuracy_decoder_config.json" # Repository root used to resolve bundled model parameter directories. Prefers # TT_METAL_RUNTIME_ROOT (the runtime-root convention set by tt-run / honored by # tt_metal/llrt/rtoptions.cpp), falling back to the traditional TT_METAL_HOME. # Resolving these paths against the repo root makes LOCAL_LLAMA_PARAMS / # LOCAL_HF_PARAMS independent of the caller's current working directory, which # matters when tt-run scripts cd into an example directory before launching. _REPO_ROOT = Path( os.environ.get("TT_METAL_RUNTIME_ROOT") or os.environ.get("TT_METAL_HOME") or str(Path(__file__).resolve().parents[3]) ) class TensorGroup(Enum): FF1_FF3 = "ff1_3" FF2 = "ff2" WQKV = "wqkv" WO = "wo" KV_CACHE = "kv_cache" ACTIVATION = "activation" class PrecisionSetting(Enum): BFP4 = "bfp4" BFP8 = "bfp8" BF16 = "bf16" class OpGroup(Enum): """ LI_* are linear operator groups SDPA_* are scaled_dot_product_attention operator groups """ LI_FF1_FF3 = "li_ff1_3" LI_FF2 = "li_ff2" LI_QKV_DECODE = "li_qkv_decode" LI_O_DECODE = "li_o_decode" SDPA_DECODE = "sdpa_decode" LI_QKV_PREFILL = "li_qkv_prefill" LI_O_PREFILL = "li_o_prefill" SDPA_PREFILL = "sdpa_prefill" ACCURACY = "accuracy" # This is a special group for accuracy mode, not an actual operator group def compute_padded_vocab_size(vocab_size: int, num_devices: int) -> int: """Pad total vocab so each device shard is tile-aligned.""" if num_devices < 1: raise ValueError(f"num_devices must be >= 1, got {num_devices}") return nearest_multiple(vocab_size, ttnn.TILE_SIZE * num_devices) def compute_galaxy_padded_vocab_size(vocab_size: int, num_devices: int) -> int: """Preserve Galaxy's 128K layout while accommodating larger vocabularies.""" return compute_padded_vocab_size(max(vocab_size, 128 * 1024), num_devices) def compute_galaxy_width_shard_cores(width: int, max_cores: int = 32) -> int: """Use the largest core count that leaves a tile-aligned width shard.""" if width <= 0 or width % ttnn.TILE_SIZE != 0: raise ValueError(f"width must be a positive multiple of {ttnn.TILE_SIZE}, got {width}") width_tiles = width // ttnn.TILE_SIZE num_cores = min(max_cores, width_tiles) while width_tiles % num_cores != 0: num_cores -= 1 return num_cores # Silicon-validated Llama-70B 7x4 MLP reduce-scatter grid. _GALAXY_FF1_LEGACY_CORES = 28 def create_galaxy_ff1_out_reduce_scatter_memcfg(hidden_dim: int, mesh_rows: int, mesh_cols: int) -> ttnn.MemoryConfig: """Create the Galaxy FF1 reduce-scatter *output* layout. ``ReduceScatterMinimalAsyncDeviceOperation::compute_output_specs`` derives the output shape itself as ``input_shape[dim] / ring_size``; this memory config only supplies the layout for that shape, so it must describe the scattered width, not the pre-scatter width. w1/w3 are 2D-sharded ``dims=(-1, -2)`` (mlp.py), so ``hidden_dim`` splits across ``mesh_rows`` and the per-device FF1 output is ``hidden_dim // mesh_rows``. The collective then scatters that over ``cluster_axis=1``, i.e. ``mesh_cols`` devices. The silicon-validated Galaxy demo agrees: its reduce-scatter input is 3840 wide (``SHARDED_FF12_OUT_RING_MEMCFG`` / ``SHARDED_FF12_PRE_MUL_RING_REDUCE_MEMCFG``, the 28672 // 8 = 3584 per-device width padded to a 30-core layout) and its output ``REDUCE_SCATTER_OUT_MEMCFG`` is ``[32, 32]`` over 30 cores = 960 = 3840 // 4. Keeps the legacy 7x4 grid wherever it tile-aligns the scattered width, which for Llama-70B it does exactly: 28672 // 8 // 4 = 896 = 28 * 32. """ per_device_width = hidden_dim // mesh_rows // mesh_cols legacy_grid = ttnn.CoreRangeSet({ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(6, 3))}) num_cores, core_grid = _GALAXY_FF1_LEGACY_CORES, legacy_grid if per_device_width % (_GALAXY_FF1_LEGACY_CORES * ttnn.TILE_SIZE) != 0: # The legacy 28-core grid cannot tile-align this width. Try a core count that # can, but fall back to the legacy layout if no supported core grid exists -- # this config is only consumed on the Galaxy decode path (mlp.py, dim == 8192), # so an unusable width here must stay harmless rather than raise during # ModelArgs construction on every SKU. if per_device_width % ttnn.TILE_SIZE == 0: candidate_cores = compute_galaxy_width_shard_cores(per_device_width) candidate_grid = num_to_coregrid(candidate_cores) if candidate_grid is not None: num_cores, core_grid = candidate_cores, candidate_grid if core_grid is legacy_grid: logger.warning( f"Galaxy FF1 per-device width {per_device_width} is not tile-shardable across " f"{_GALAXY_FF1_LEGACY_CORES} cores and has no supported alternative grid; keeping " "the legacy layout. Only consumed on the Galaxy decode path." ) # Round the shard up to a tile boundary. This is exact for every width a grid can # cover evenly (Llama-70B 896 = 28 * 32, Qwen-72B 1024 = 32 * 32) and mirrors what # the validated demo does otherwise -- its 30-core layout pads 896 up to 960. shard_width = math.ceil(per_device_width / num_cores / ttnn.TILE_SIZE) * ttnn.TILE_SIZE # LAYOUT HISTORY -- start here if Galaxy decode perf or L1 usage regresses (PR #53838). # # Llama-70B (hidden_dim 28672): [32, 128] over 28 cores -> [32, 32] over 28 cores # 3584 columns 896 columns # # 3584 is the pre-scatter per-device width (28672 // 8); 896 is what # reduce_scatter_minimal_async actually emits (28672 // 8 // 4). Same 7x4 grid, # 4x less L1 for this buffer. # # The collective tolerates an over-provisioned output shard spec -- the old value # was oversized, not wrong -- so a regression here would show up as perf or L1 # pressure, never as a failed assertion. If Galaxy decode slows down or starts # hitting L1 limits, revert this shard width to `per_device_width // num_cores` # with `per_device_width = hidden_dim // mesh_rows` and see if it recovers. return ttnn.create_sharded_memory_config( shape=(32, shard_width), core_grid=core_grid, strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) def should_pad_sampling_logits_to_power_of_2( base_model_name: str, padded_vocab_size: int, sampling_splits: int ) -> bool: # Enable optional sampling padding for models that regress to single-core TopK. More info at issue #40399 if sampling_splits < 1: raise ValueError(f"sampling_splits must be >= 1, got {sampling_splits}") per_device_vocab = padded_vocab_size // sampling_splits return per_device_vocab > 0 and (per_device_vocab & (per_device_vocab - 1)) != 0 class MathFidelitySetting(Enum): LOFI = "lofi" HIFI2 = "hifi2" HIFI2_NA = "hifi2na" # na specified `packer_l1_acc=False` and `fp32_dest_acc_en=False` in compute kernel config HIFI2_FP16 = "hifi2fp16" # fp16 specified `fp32_dest_acc_en=False` in compute kernel config HIFI2_NOL1ACC = "hifi2nol1acc" # fp32_dest_acc_en=True but packer_l1_acc=False (issue #36378) HIFI4 = "hifi4" HIFI4_FP16 = "hifi4fp16" # fp16 specified `fp32_dest_acc_en=False` in compute kernel config HIFI4_FP32 = "hifi4fp32" class ModelOptimizations: @classmethod def accuracy(cls, model_name): """Configuration optimized for accuracy 70B+ models still use bfp4 MLPs and BFP8 attention in this configuration """ base_model_name = get_base_model_name(model_name) if base_model_name in ["Llama-3.1-70B", "Llama-3.2-90B", "DeepSeek-R1-Distill-Llama-70B"]: logger.info( f"{model_name} is >70B and large models test insensitive precision, using BFP4 MLPs and BFP8 attention even in accuracy mode" ) inst = cls( { "TensorPrecision": {TensorGroup.FF1_FF3: PrecisionSetting.BFP4}, "OpFidelity": {OpGroup.LI_FF1_FF3: MathFidelitySetting.LOFI}, } ) else: if ( base_model_name.startswith("Llama-3") or base_model_name.startswith("Mistral-7B") or base_model_name.startswith("Phi-3-mini") or base_model_name.startswith("phi-4") or base_model_name.startswith("Meta-Llama-3") ): if model_name.startswith("phi-4"): logger.info( f"Model {model_name} is running out of DRAM memory for weight fetching under standard accuracy settings, using BFP8 for WQKV" ) logger.info( f"Llama 3, Mistral 7B and Phi3-mini models test insensitive to attention precision, using BFP8 attention and kv-cache with FP16 MLP accumulation even in accuracy mode" ) settings = { "TensorPrecision": { TensorGroup.WQKV: PrecisionSetting.BFP8, TensorGroup.KV_CACHE: PrecisionSetting.BFP8, TensorGroup.WO: PrecisionSetting.BFP8, }, "OpFidelity": { OpGroup.LI_FF1_FF3: MathFidelitySetting.HIFI2_FP16, OpGroup.LI_FF2: MathFidelitySetting.HIFI2_FP16, }, } if model_name.startswith("Phi-3-mini"): # TODO: Only do this for N150 logger.info( f"Model {model_name} is running out of L1 memory under standard accuracy settings, using FP16 accumulate in attention prefill QKV Matmul" ) settings["OpFidelity"][OpGroup.LI_QKV_PREFILL] = MathFidelitySetting.HIFI2_FP16 inst = cls(settings) else: inst = cls( { "TensorPrecision": { TensorGroup.WQKV: PrecisionSetting.BF16, TensorGroup.KV_CACHE: PrecisionSetting.BF16, TensorGroup.WO: PrecisionSetting.BF16, }, "OpFidelity": { OpGroup.LI_QKV_DECODE: MathFidelitySetting.HIFI4, OpGroup.LI_QKV_PREFILL: MathFidelitySetting.HIFI4, OpGroup.SDPA_DECODE: MathFidelitySetting.HIFI4, OpGroup.SDPA_PREFILL: MathFidelitySetting.HIFI4, OpGroup.LI_O_DECODE: MathFidelitySetting.HIFI4, OpGroup.LI_O_PREFILL: MathFidelitySetting.HIFI4, }, } ) inst.__name__ = "accuracy" return inst @classmethod def performance(cls, model_name): """Configuration optimized for performance All models use bfp4 in FF1 and FF3 MLPs in this configuration """ base_model_name = get_base_model_name(model_name) if base_model_name in ["Qwen2.5-7B", "Qwen2.5-VL-7B"]: logger.info( f"Model {model_name} is degraded under standard high-performance settings, using BF16 attention and BFP8 MLP" ) inst = cls( { "TensorPrecision": { TensorGroup.WQKV: PrecisionSetting.BF16, TensorGroup.KV_CACHE: PrecisionSetting.BF16, TensorGroup.WO: PrecisionSetting.BF16, }, "OpFidelity": { OpGroup.LI_QKV_DECODE: MathFidelitySetting.HIFI4, OpGroup.LI_QKV_PREFILL: MathFidelitySetting.HIFI4, OpGroup.SDPA_DECODE: MathFidelitySetting.HIFI4, OpGroup.SDPA_PREFILL: MathFidelitySetting.HIFI4, OpGroup.LI_O_DECODE: MathFidelitySetting.HIFI4, OpGroup.LI_O_PREFILL: MathFidelitySetting.HIFI4, }, } ) else: settings = { "TensorPrecision": {TensorGroup.FF1_FF3: PrecisionSetting.BFP4}, "OpFidelity": {OpGroup.LI_FF1_FF3: MathFidelitySetting.LOFI}, } if model_name.startswith("Phi-3-mini"): # TODO: Only do this for N150 logger.info( f"Model {model_name} is running out of L1 memory under standard high-performance settings, using FP16 accumulate in attention prefill QKV Matmul" ) settings["OpFidelity"][OpGroup.LI_QKV_PREFILL] = MathFidelitySetting.HIFI2_FP16 inst = cls(settings) inst.__name__ = "performance" return inst def __init__(self, settings: dict = None): if settings: self._validate_settings(settings) self._opt_settings = self._default_settings() self._names = {} for key, enum_type in (("TensorPrecision", TensorGroup), ("OpFidelity", OpGroup)): self._opt_settings[key].update((settings or {}).get(key, {})) curr = self._opt_settings[key] self._names[key] = ", ".join( [f"{k.value}: {curr[k].value if curr[k] else 'mixed'}" for k in list(enum_type)] ) self._full_name = ( "precision_cfg = {" + self._names["TensorPrecision"] + "}, fidelity_cfg = {" + self._names["OpFidelity"] + "}" ) # NOTE: self.__name__ is used by test/demo flows to distinguish performance vs accuracy # mode when selecting centralized targets from models/model_targets.yaml. self.__name__ = self._full_name # TODO: maybe we could warn about some unwanted settings here def _validate_settings(self, settings: dict): # Check that only valid top-level keys are used valid_keys = {"TensorPrecision", "OpFidelity"} invalid_keys = set(settings.keys()) - valid_keys if invalid_keys: raise ValueError(f"Invalid settings keys: {invalid_keys}. Must be one of {valid_keys}") # Validate TensorPrecision settings if "TensorPrecision" in settings: for key, value in settings["TensorPrecision"].items(): if not isinstance(key, TensorGroup): raise ValueError(f"Invalid TensorPrecision key: {key}. Must be a TensorGroup enum value") if not isinstance(value, PrecisionSetting): raise ValueError(f"Invalid TensorPrecision value: {value}. Must be a PrecisionSetting enum value") # Validate OpFidelity settings if "OpFidelity" in settings: for key, value in settings["OpFidelity"].items(): if not isinstance(key, OpGroup): raise ValueError(f"Invalid OpFidelity key: {key}. Must be an OpGroup enum value") if not isinstance(value, MathFidelitySetting): raise ValueError(f"Invalid OpFidelity value: {value}. Must be a MathFidelitySetting enum value") def _default_settings(self): """Default is BFP8/HIFI2 everywhere, activation follows input type (usually BF16) Only exceptions: - SDPA runs in HIFI4 during prefill (still HIFI2 during decode) """ return { "TensorPrecision": { # MLP TensorGroup.FF1_FF3: PrecisionSetting.BFP8, TensorGroup.FF2: PrecisionSetting.BFP8, # Attention TensorGroup.WQKV: PrecisionSetting.BFP8, TensorGroup.WO: PrecisionSetting.BFP8, TensorGroup.KV_CACHE: PrecisionSetting.BFP8, # Activation across whole model TensorGroup.ACTIVATION: None, # this signals that original dtype should be used }, "OpFidelity": { # MLP linear operators - BFP8 with FP16 accumulation to save L1 OpGroup.LI_FF1_FF3: MathFidelitySetting.HIFI2_FP16, OpGroup.LI_FF2: MathFidelitySetting.HIFI2_FP16, # Attention operators -- linear and scaled_dot_product_attention, in decode and prefill modes OpGroup.LI_QKV_DECODE: MathFidelitySetting.HIFI2, OpGroup.SDPA_DECODE: MathFidelitySetting.HIFI2, OpGroup.LI_O_DECODE: MathFidelitySetting.HIFI2, OpGroup.LI_QKV_PREFILL: MathFidelitySetting.HIFI2, OpGroup.SDPA_PREFILL: MathFidelitySetting.HIFI4, OpGroup.LI_O_PREFILL: MathFidelitySetting.HIFI2, # FP32 accumulate is important here OpGroup.ACCURACY: MathFidelitySetting.HIFI4_FP32, }, } @property def tensor_dtype_settings(self): return self._opt_settings["TensorPrecision"] @property def op_fidelity_settings(self): return self._opt_settings["OpFidelity"] def parse_optimizations(string): """ Parse the optimizations full name and return a ModelOptimizations instance. """ # Find the precision and fidelity config sections precision_start = string.find("precision_cfg") fidelity_start = string.find("fidelity_cfg") if precision_start == -1 and fidelity_start == -1: raise ValueError("String must contain either precision_cfg or fidelity_cfg") # Extract the config dictionaries between { } def extract_config(start_idx, cfg_name): open_brace = string.find("{", start_idx) if open_brace == -1: raise ValueError(f"Missing opening brace for {cfg_name}") close_brace = string.find("}", open_brace) if close_brace == -1: raise ValueError(f"Missing closing brace for {cfg_name}") return string[open_brace + 1 : close_brace].strip() precision_dict = extract_config(precision_start, "precision_cfg") if precision_start != -1 else {} fidelity_dict = extract_config(fidelity_start, "fidelity_cfg") if fidelity_start != -1 else {} # Create ModelOptimizations instance with the parsed configs settings = {"TensorPrecision": {}, "OpFidelity": {}} # Parse precision config for pair in precision_dict.split(","): if ":" not in pair: raise ValueError("Invalid format - missing ':' separator") key, value = pair.split(":") key = TensorGroup(key.strip()) value = value.strip() if key == TensorGroup.ACTIVATION and value == "mixed": # special case for activation's mixed precision, which is the default configuration continue settings["TensorPrecision"][key] = PrecisionSetting(value) # Parse fidelity config for pair in fidelity_dict.split(","): if ":" not in pair: raise ValueError("Invalid format - missing ':' separator") key, value = pair.split(":") key = OpGroup(key.strip()) value = MathFidelitySetting(value.strip()) settings["OpFidelity"][key] = value model_opt = ModelOptimizations(settings) def apply_settings(model_args): return DecodersPrecision(model_args.n_layers, model_args.model_name, model_opt) apply_settings.__name__ = model_opt.__name__ return apply_settings def parse_decoder_json(json_file_path, default_optimization=ModelOptimizations.performance): """ Reads a JSON file and returns a DecodersPrecision instance. """ if not json_file_path: return None json_file_path = Path(json_file_path) if not json_file_path.exists(): raise FileNotFoundError(f"JSON configuration file not found: {json_file_path}") try: with open(json_file_path, "r") as f: config_data = json.load(f) if "decoders" not in config_data: raise ValueError("Invalid JSON format: Missing 'decoders' key") num_decoders = max(int(decoder_id) for decoder_id in config_data["decoders"].keys()) + 1 placeholder_model_name = "model" decoder_conf = default_optimization(placeholder_model_name) default_tensor_dtype_settings = decoder_conf.tensor_dtype_settings default_op_fidelity_settings = decoder_conf.op_fidelity_settings decoders_precision = DecodersPrecision(num_decoders, placeholder_model_name, decoder_conf) for decoder_id, settings in config_data["decoders"].items(): decoder_id = int(decoder_id) tensor_precision = ( {TensorGroup[key]: PrecisionSetting[value] for key, value in settings.get("precision_cfg").items()} if "precision_cfg" in settings else default_tensor_dtype_settings ) op_fidelity = ( {OpGroup[key]: MathFidelitySetting[value] for key, value in settings.get("fidelity_cfg").items()} if "fidelity_cfg" in settings else default_op_fidelity_settings ) custom_opt = ModelOptimizations({"TensorPrecision": tensor_precision, "OpFidelity": op_fidelity}) decoders_precision.set_decoder_conf(decoder_id, custom_opt) return decoders_precision except Exception as e: raise ValueError(f"Error loading JSON configuration: {e}") class CheckpointType(Enum): Meta = auto() HuggingFace = auto() class ModelArgs: OP_KEYS = ( # Embedding "EMB_WEIGHTS", # Feed forward "MLP_WEIGHTS", "FF1_OUTPUT", "FF3_OUTPUT", "FF2_OUTPUT", "MLP_W_LAYOUT", # Attention "ATTN_WEIGHTS", "XQKV_MM_OUTPUT", "QKV_HEADS_OUTPUT", "QV_ROT_EMB_OUTPUT", "KV_UNPAD_OUTPUT", "QK_MM_OUTPUT", "QKV_MM_OUTPUT", "CONCAT_HEADS_OUTPUT", "ATTN_OUTPUT", "ATTN_W_LAYOUT", # Decoder "DECODE_RESIDUAL", "OUTPUT_MM", # MoE "GATE_W_LAYOUT", "GATE_WEIGHTS", "GATE_MM_OUTPUT", ) LOCAL_LLAMA_PARAMS = { k: str(_REPO_ROOT / v) for k, v in { "LLAMA3_2_1B_PARAMS": "models/tt_transformers/model_params/Llama-3.2-1B-Instruct", "LLAMA3_2_3B_PARAMS": "models/tt_transformers/model_params/Llama-3.2-3B-Instruct", "LLAMA3_1_8B_PARAMS": "models/tt_transformers/model_params/Llama-3.1-8B-Instruct", "LLAMA3_2_11B_PARAMS": "models/tt_transformers/model_params/Llama-3.2-11B-Vision-Instruct", "LLAMA3_1_70B_PARAMS": "models/tt_transformers/model_params/Llama-3.1-70B-Instruct", "LLAMA3_2_90B_PARAMS": "models/tt_transformers/model_params/Llama-3.2-90B-Vision-Instruct", }.items() } LOCAL_HF_PARAMS = { k: str(_REPO_ROOT / v) for k, v in { "Llama-3.1-8B-Instruct": "models/tt_transformers/model_params/Llama-3.1-8B-Instruct", "Llama-3.1-70B-Instruct": "models/tt_transformers/model_params/Llama-3.1-70B-Instruct", "Llama-3.2-1B-Instruct": "models/tt_transformers/model_params/Llama-3.2-1B-Instruct", "Llama-3.2-3B-Instruct": "models/tt_transformers/model_params/Llama-3.2-3B-Instruct", "Llama-3.2-11B-Instruct": "models/tt_transformers/model_params/Llama-3.2-11B-Vision-Instruct", "Llama-3.2-11B-Vision-Instruct": "models/tt_transformers/model_params/Llama-3.2-11B-Vision-Instruct", "Llama-3.2-90B-Instruct": "models/tt_transformers/model_params/Llama-3.2-90B-Vision-Instruct", "Llama-3.2-90B-Vision-Instruct": "models/tt_transformers/model_params/Llama-3.2-90B-Vision-Instruct", "Mistral-7B-Instruct-v0.3": "models/tt_transformers/model_params/Mistral-7B-Instruct-v0.3", "Qwen2.5-VL-3B-Instruct": "models/tt_transformers/model_params/Qwen2.5-VL-3B-Instruct", "Qwen2.5-VL-32B-Instruct": "models/tt_transformers/model_params/Qwen2.5-VL-32B-Instruct", "Phi-4": "models/tt_transformers/model_params/phi-4", "Qwen2.5-VL-72B-Instruct": "models/tt_transformers/model_params/Qwen2.5-VL-72B-Instruct", "Qwen3-VL-32B-Instruct": "models/tt_transformers/model_params/Qwen3-VL-32B-Instruct", "Qwen3-32B": "models/tt_transformers/model_params/Qwen3-32B", "EXAONE-4.5-33B": "models/tt_transformers/model_params/EXAONE-4.5-33B", "Qwen2.5-72B-Instruct": "models/tt_transformers/model_params/Qwen2.5-72B-Instruct", "Qwen2.5-32B-Instruct": "models/tt_transformers/model_params/Qwen2.5-32B-Instruct", "Meta-Llama-3-8B": "models/tt_transformers/model_params/Meta-Llama-3-8B", "Meta-Llama-3-8B-Instruct": "models/tt_transformers/model_params/Meta-Llama-3-8B", "Qwen3.6-27B": "models/tt_transformers/model_params/Qwen3.6-27B", "LFM2.5-VL-1.6B": "models/tt_transformers/model_params/LFM2.5-VL-1.6B", }.items() } MAX_QKV_MM_SEQ_LEN = 2048 # Opt-in: False so TP > n_kv_heads stays a hard error unless a subclass implements KV replication. SUPPORTS_KV_REPLICATION = False def __init__( self, mesh_device, instruct=False, dummy_weights=False, max_batch_size=1, max_seq_len=1024 * 128, optimizations=None, cache_hf=False, # Set to False to reduce memory usage by not caching HF model prefetcher=None, use_hf_rope=False, # Choose HF or mllama RoPE (default: mllama, previously, only that one was used). mllama will be removed, only HF will remain (Issue #37605). ): self.num_devices = mesh_device.get_num_devices() if mesh_device else 0 self.mesh_device = mesh_device self.arch_name = ttnn.get_arch_name() self.dram_grid_size = mesh_device.dram_grid_size() if mesh_device else None # CoreCoord with (x, y) self.prefetcher = prefetcher self.device_name = determine_device_name(self.mesh_device) if mesh_device is not None else "CPU" logger.info(f"Inferring device name: {self.device_name}") self.cluster_shape = list(mesh_device.shape) if mesh_device is not None else None self.cluster_type = ttnn.cluster.get_cluster_type() if mesh_device is not None else None self.is_galaxy_cluster = self.cluster_type in [ ttnn.cluster.ClusterType.GALAXY, ttnn.cluster.ClusterType.TG, ttnn.cluster.ClusterType.BLACKHOLE_GALAXY, ] self.is_galaxy = self.num_devices == 32 self.model_name = "Unknown" # Llama model name will be dependent on the checkpoint directory self.max_seq_len = max_seq_len self.max_batch_size = max_batch_size if self.num_devices == 32: self.batch_size_per_device_group = max(self.max_batch_size // list(mesh_device.shape)[1], 1) else: self.batch_size_per_device_group = self.max_batch_size self.tile_size = ttnn.TILE_SIZE # Expose for downstream consumers (attention, etc.) self.is_70b = False self.is_90b = False self.fuse_qkv = False self.fuse_mlp = False self.trust_remote_code_hf = False self.prefill_len_cutoff = 512 if is_blackhole() else 1024 self.dummy_weights = dummy_weights self.cache_hf_flag = cache_hf # Whether to cache HF model to avoid multiple loads (uses extra memory) self.cached_hf_model = None # Save any HF model object to avoid loading it multiple times for reference methods self.rms_norm_add_unit_offset = False self.embed_scale = None # Post-norm decoder (EXAONE-4.x): no input_layernorm; norms are applied to the # attention/MLP outputs before the residual adds. Set in _set_model_specific_params(). self.use_post_norm = False # Hybrid-rope inversion (EXAONE-4.x): sliding layers use the (llama3-scaled) # rope while full-attention layers are NoPE. rope_scaling_local feeds the local # rope setup; use_global_nope neutralizes the global setup's cos/sin to identity. self.rope_scaling_local = None self.use_global_nope = False # Text-only port of a multimodal checkpoint: keeps is_multimodal False so the # text pipeline (AutoModel class choice aside) is used end-to-end. self.force_text_only = False # Final logit soft-capping (Gemma-2). None => disabled, so no effect on other models. # Attention-score softcapping (HF attn_logit_softcapping=50.0) is intentionally # not stored or applied: ttnn SDPA has no softcap hook, and HF documents that # omitting attn softcap at inference has only minor effect. See final-logit # application in Transformer._apply_final_logit_softcapping. self.final_logit_softcapping = None self.model_type = None # Decode-SDPA tuning. Architecture-specific overrides are applied once in # _set_model_specific_params(); these defaults preserve existing behaviour for # every model (q/k chunk 0 => op auto-selects, forced framework compute config). self.sdpa_decode_q_chunk_size = 0 self.sdpa_decode_k_chunk_size = 0 self.sdpa_decode_use_default_compute_config = False self.use_hf_rope = use_hf_rope assert not os.getenv( "FAKE_DEVICE" ), "FAKE_DEVICE has been renamed to MESH_DEVICE for consistency with vLLM, please update your environment variables and run again." # Remove trailing slashes so basename gets the right model name HF_MODEL = os.getenv("HF_MODEL") self.CACHE_PATH = os.getenv("TT_CACHE_PATH") if HF_MODEL: self.checkpoint_type = CheckpointType.HuggingFace self.CKPT_DIR = HF_MODEL self.TOKENIZER_PATH = HF_MODEL if not self.CACHE_PATH: self.CACHE_PATH = os.path.join("model_cache", HF_MODEL, self.device_name) else: # For HF models, always append the device name (e.g. N150/N300/T3K/TG) to the cache path self.CACHE_PATH = os.path.join(self.CACHE_PATH, self.device_name) if self.use_hf_rope: self.CACHE_PATH = os.path.join(self.CACHE_PATH, "hf_rope") self.model_name = HF_MODEL.strip("/").split("/")[ -1 ] # HF model names use / even on windows. May be overridden by config. if "phi-4" in self.model_name.lower(): self.model_name = "Phi-4" else: raise ValueError("Please set HF_MODEL to a HuggingFace name e.g. meta-llama/Llama-3.1-8B-Instruct") logger.info(f"Checkpoint directory: {self.CKPT_DIR}") logger.info(f"Tokenizer file: {self.TOKENIZER_PATH + '/tokenizer.model'}") logger.info(f"Cache directory: {self.CACHE_PATH}") logger.info(f"Model name: {self.model_name}") # Some consumers like SentencePiece only accept str not Path for files self.model_base_path = Path(self.CKPT_DIR) self.model_cache_path = Path(self.CACHE_PATH) # Load weights and tokenizer self.consolidated_weights_path = self.CKPT_DIR + "/consolidated.00.pth" self.tokenizer_path = self.TOKENIZER_PATH + "/tokenizer.model" self.instruct = instruct # If the weights file contain the keyword `instruct` also set self.instruct to true if any(keyword in self.CKPT_DIR.lower() for keyword in ("instruct", "it")): self.instruct = True # Check for supported batches since previous logic that contained the check was removed because it was unused supported_batches = {1, 2, 4, 8, 16, 32} if self.max_batch_size not in supported_batches: raise ValueError(f"Batch size {self.max_batch_size} not supported") # Load model params if self.base_model_name in ["Phi-3-mini-128k-instruct"]: self.trust_remote_code_hf = True self._set_hf_params(self.CKPT_DIR) # Set the max number of tokens for each prefill chunk based on the model and device self.max_prefill_chunk_size = self.get_max_prefill_chunk_size() # TODO: Enable batched_prefill once this is fixed: https://github.com/tenstorrent/tt-metal/issues/47238 # Prefill logits are batch-variant on multi-chip Blackhole: the # float-reduction order in the prefill matmul/attention depends on the # batched-token count, so identical same-seed requests that land in # different-sized prefill batches get slightly different logits. Seeded # sampling under concurrency then diverges (tt-inference-server#4004, # test_non_uniform_seeding). Forcing per-user batch-1 prefill removes the # variance. Workaround until the prefill kernels are batch-invariant. # Disabled for Qwen3-32B (P150x4) and Llama-3.1-8B (P300/P150x4/P150x8). self.disable_batched_prefill = (self.base_model_name == "Qwen3-32B" and self.device_name == "P150x4") or ( self.base_model_name == "Llama-3.1-8B" and self.device_name in ("P150", "P300", "P150x4", "P150x8") ) if ( self.base_model_name in ["Llama-3.1-8B", "Llama-3.2-11B", "Mistral-7B", "gemma-3-27b", "gemma-3-4b", "Phi-4"] and self.device_name == "N150" ) or (self.base_model_name in ["Qwen2.5-7B", "Qwen2.5-VL-7B", "Phi-4"] and self.device_name == "N300"): logger.info(f"Reducing prefill_len_cutoff to 512 for {self.model_name} on {self.device_name}") self.prefill_len_cutoff = 512 elif self.base_model_name in ["Mixtral-8x7B"] and self.device_name == "T3K": self.prefill_len_cutoff = 512 if callable(optimizations): self.optimizations = optimizations(self) else: self.optimizations = optimizations # Configure data precision and math fidelity for tensors and kernels if self.optimizations is None: self.optimizations = DecodersPrecision.accuracy(num_decoders=self.n_layers, model_name=self.model_name) self.dummy_weights = dummy_weights self.tile_padded_batch_rows = ttnn.TILE_SIZE * int(math.ceil(self.max_batch_size / ttnn.TILE_SIZE)) # Enable workarounds by default until di/dt issues are fixed self.di_dt_workaround = os.getenv("DISABLE_DI_DT_WORKAROUND") != "1" if not self.di_dt_workaround: logger.info("Disabling di/dt workaround, re-enable if you see hangs") self.model_config = {} # Update memory configs (weights->DRAM, activations->L1) self.model_config.update( { f"{key}_MEMCFG": ttnn.DRAM_MEMORY_CONFIG if "WEIGHTS" in key else ttnn.L1_MEMORY_CONFIG for key in self.OP_KEYS } ) self.model_config["DECODERS_OPTIMIZATIONS"] = self.optimizations # Update memory layouts (Tile, except MLP) self.model_config.update({f"{key}_TILE": ttnn.TILE_LAYOUT for key in self.OP_KEYS if "LAYOUT" in key}) self.tokenizer = None if dummy_weights else self.create_tokenizer() self.processor = None if dummy_weights else self.create_processor() # Flag to indicate whether we use fused version of QK ops (rotary embedding + page cached update) # We currently disable this fusion of ops for vision-capable or multimodal models # we also disable fused qk when using HF-style rotary embedding self.use_qk_fused = not self.is_multimodal and not self.use_hf_rope if self.prefetcher is not None: self.use_qk_fused = False if self.mesh_device is not None: # Avoid issue with test_torch.py not having a device # ============================================================================ # Parameter initialization # ============================================================================ # nlp_concat_heads_decode will shard the data across this number of cores assert ( self.n_heads % self.cluster_shape[1] == 0 ), f"n_heads must be divisible by num_devices: {self.n_heads} % {self.cluster_shape[1]}" # KV heads normally split one-or-more per device. SUPPORTS_KV_REPLICATION models # instead replicate a head across the devices holding its GQA query group, so they # only need TP to be a whole multiple of n_kv_heads (e.g. 4 KV heads on TP=8 -> each # head shared by a device pair). n_local_kv_heads is then 1, hence the max() below. if self.SUPPORTS_KV_REPLICATION and self.cluster_shape[1] > self.n_kv_heads: assert self.cluster_shape[1] % self.n_kv_heads == 0, ( f"with KV replication, num_devices must be a multiple of n_kv_heads: " f"{self.cluster_shape[1]} % {self.n_kv_heads}" ) else: assert self.n_kv_heads % self.cluster_shape[1] == 0, "n_kv_heads must be divisible by num_devices" self.n_local_heads = self.n_heads // self.cluster_shape[1] self.qkv_size = self.head_dim * (2 * self.n_kv_heads + self.n_heads) self.min_kv_prefill_shard_seqlen = (ttnn.TILE_SIZE * 8 * 8) / max( 1, self.n_kv_heads // self.cluster_shape[1] ) # All Gather Matmul for Dense Out (DO) - computed flag stored as instance attribute # NOTE: Fused all gather matmul only supports a core grid of size num_devices x 1 # Galaxy DP4 gives Llama 8B a routeable 1x8 row submesh, so it can # use the same fused AGMM path as T3K on Galaxy-class systems. use_galaxy_dp4_8b_submesh_agmm = ( self.is_galaxy_cluster and self.base_model_name == "Llama-3.1-8B" and self.num_devices == 8 and tuple(self.cluster_shape) == (1, 8) ) self._use_t3k_fused_agmm_config = not self.is_galaxy_cluster or use_galaxy_dp4_8b_submesh_agmm self._use_fused_all_gather_matmul = ( self.num_devices == 8 and self._use_t3k_fused_agmm_config and (self.dim // ttnn.TILE_SIZE // self.num_devices) % self.num_devices == 0 and self.num_devices > 1 and self.ccl_topology() == ttnn.Topology.Ring ) or self.prefetcher is not None # Using dram_shard_grid_width to ensure per_core_N matches DRAM shard width for P100, otherwise matmuls silently give bad PCC self.dram_shard_grid_width = 8 if is_wormhole_b0() else self.dram_grid_size.x # 7 for P100, 8 for P150 # ============================================================================ # Core Grid Configurations for DRAM weight sharding, LM Head and MLP # ============================================================================ # DRAM weight grid specs for dram sharding matmuls grid = self.mesh_device.compute_with_storage_grid_size() self.max_grid_size = ttnn.CoreGrid(x=grid.x, y=grid.y) self.dram_weight_grid = ttnn.CoreRangeSet( { ttnn.CoreRange( ttnn.CoreCoord(0, 0), ttnn.CoreCoord(self.dram_grid_size.x - 1, self.dram_grid_size.y - 1), ) } ) if self.num_devices == 32: lm_head_num_rows = 4 while self.dim % (self.num_devices * ttnn.TILE_SIZE * lm_head_num_rows) != 0: lm_head_num_rows -= 1 else: lm_head_num_rows = 8 lm_head_cores_per_row = 8 while self.dim % (ttnn.TILE_SIZE * lm_head_num_rows * lm_head_cores_per_row) != 0: lm_head_num_rows -= 1 if lm_head_num_rows == 0: lm_head_cores_per_row -= 1 if lm_head_cores_per_row == 0: raise ValueError( f"Could not find a lm_head_num_rows such that self.dim(={self.dim}) % (lm_head_num_rows * 8) == 0" ) lm_head_num_rows = 8 self.lm_head_core_grid = ttnn.CoreGrid(y=lm_head_num_rows, x=lm_head_cores_per_row) self.max_columns_per_device_lm_head = self.get_lm_head_max_columns_per_device( self.lm_head_core_grid, self.prefetcher ) # For maximum performance, set the prefill grid row to 8, even if it can fit in a smaller grid self.prefill_rows = 8 self.attn_input_grid = self.dram_shard_core_grid_for_k(self.dim) self.mlp1_3_grid = lambda seq_len: ( (8, min(min(seq_len, 1024) // 32, 4)) if self.is_galaxy else self.find_prefill_grid(self.prefill_rows, self.dim // ttnn.TILE_SIZE) ) self.mlp2_grid = lambda seq_len: ( (8, min(min(seq_len, 1024) // 32, 4)) if self.is_galaxy else self.find_prefill_grid(self.prefill_rows, self.hidden_dim // ttnn.TILE_SIZE) ) self.mlp_core_grid = ( self.dram_shard_core_grid_for_k(self.dim) if self.is_galaxy else self.dram_shard_core_grid_for_k_and_n(self.dim, self.hidden_dim // self.num_devices) ) self.mlp2_core_grid = ( ttnn.CoreGrid(y=1, x=8) if self.is_galaxy else self.dram_shard_core_grid_for_k_and_n(self.hidden_dim // self.num_devices, self.dim) ) # ============================================================================ # Compute kernels Configs # Note: FP32 acc does not appear to be needed for accuracy in model tests or demo runs. # ============================================================================ self.compute_kernel_config_lofi = ttnn.WormholeComputeKernelConfig( math_fidelity=ttnn.MathFidelity.LoFi, math_approx_mode=False, fp32_dest_acc_en=False, packer_l1_acc=True, ) self.compute_kernel_config_hifi2 = ttnn.WormholeComputeKernelConfig( math_fidelity=ttnn.MathFidelity.HiFi2, math_approx_mode=True, fp32_dest_acc_en=True, packer_l1_acc=True, ) self.compute_kernel_config_hifi2_fp16 = ttnn.WormholeComputeKernelConfig( math_fidelity=ttnn.MathFidelity.HiFi2, math_approx_mode=False, fp32_dest_acc_en=False, packer_l1_acc=True, ) self.compute_kernel_config_hifi4 = ttnn.WormholeComputeKernelConfig( math_fidelity=ttnn.MathFidelity.HiFi4, math_approx_mode=False, fp32_dest_acc_en=True, packer_l1_acc=True, ) self.compute_kernel_config_hifi4_fp16 = ttnn.WormholeComputeKernelConfig( math_fidelity=ttnn.MathFidelity.HiFi4, math_approx_mode=False, fp32_dest_acc_en=False, packer_l1_acc=True, ) self.compute_kernel_config_hifi4_fp32 = ttnn.WormholeComputeKernelConfig( math_fidelity=ttnn.MathFidelity.HiFi4, fp32_dest_acc_en=True, packer_l1_acc=True, dst_full_sync_en=False, ) self.compute_kernel_config_hifi2_na = ttnn.WormholeComputeKernelConfig( math_fidelity=ttnn.MathFidelity.HiFi2, math_approx_mode=False, fp32_dest_acc_en=False, packer_l1_acc=False, ) self.compute_kernel_config_hifi2_nol1acc = ttnn.WormholeComputeKernelConfig( math_fidelity=ttnn.MathFidelity.HiFi2, math_approx_mode=True, fp32_dest_acc_en=True, packer_l1_acc=False, ) self.compute_kernel_config_sdpa = ttnn.WormholeComputeKernelConfig( math_fidelity=ttnn.MathFidelity.HiFi4, math_approx_mode=False, fp32_dest_acc_en=True, packer_l1_acc=False, ) # ============================================================================ # Mixtral/MoE Model Configs (dictionary access only) # TODO: Migrate these to use getter methods after TTTv2 migration # These configs are used by mixtral_moe.py and mixtral_mlp.py # ============================================================================ n_w1_w3 = self.hidden_dim // self.cluster_shape[1] self.model_config["MIXTRAL_PREFILL_MLP_COMPUTE_CONFIG"] = self.compute_kernel_config_lofi self.model_config["MIXTRAL_GATE_MM_OUTPUT_KERNEL_CONFIG"] = self.compute_kernel_config_lofi self.model_config["DECODERS_OPTIMIZATIONS"] = self.optimizations # Mixtral prefill program configs self.model_config["PREFILL_MIXTRAL_MLP_W1_PRG_CONFIG"] = lambda seq_len: self.matmul_config( m=min(seq_len, self.prefill_len_cutoff), k=self.dim // self.cluster_shape[0], n=self.hidden_dim // self.cluster_shape[1], grid_size=self.mlp1_3_grid(min(seq_len, self.prefill_len_cutoff)), per_core_M=math.ceil(min(seq_len, self.prefill_len_cutoff) / ttnn.TILE_SIZE / self.cluster_shape[1]), per_core_N=math.ceil(n_w1_w3 / ttnn.TILE_SIZE / self.cluster_shape[0]), fused_activation=ttnn.UnaryOpType.SILU, ) self.model_config["PREFILL_MIXTRAL_MLP_W3_PRG_CONFIG"] = lambda seq_len: self.matmul_config( m=min(seq_len, self.prefill_len_cutoff), # 512 if BH, 1024 if WH k=self.dim // self.cluster_shape[0], n=n_w1_w3, grid_size=self.mlp1_3_grid(min(seq_len, self.prefill_len_cutoff)), per_core_M=math.ceil(min(seq_len, self.prefill_len_cutoff) / ttnn.TILE_SIZE / self.cluster_shape[1]), per_core_N=math.ceil(n_w1_w3 / ttnn.TILE_SIZE / self.cluster_shape[0]), ) self.model_config["PREFILL_MLP_W1_PRG_CONFIG_128"] = ttnn.MatmulMultiCoreReuseMultiCastProgramConfig( compute_with_storage_grid_size=(8, 8), in0_block_w=1, # how much inner dim you take each time out_subblock_h=1, # Must be divisible by per_core_M out_subblock_w=1, # Must be divisible by per_core_N, out_subblock_w * out_subblock_h <= 4 per_core_M=1, # 32, #16, # M / TILE_HEIGHT / Grid_Size (dynamic based on seqlen) per_core_N=56, # N / TILE_WIDTH / Grid_Size transpose_mcast=False, fused_activation=ttnn.UnaryOpType.SILU, fuse_batch=False, ) self.model_config["PREFILL_MLP_W3_PRG_CONFIG_128"] = ttnn.MatmulMultiCoreReuseMultiCastProgramConfig( compute_with_storage_grid_size=(8, 8), in0_block_w=1, # how much inner dim you take each time out_subblock_h=1, # Must be divisible by per_core_M out_subblock_w=1, # Must be divisible by per_core_N, out_subblock_w * out_subblock_h <= 4 per_core_M=1, # M / TILE_HEIGHT / Grid_Size (dynamic based on seqlen) per_core_N=56, # N / TILE_WIDTH / Grid_Size transpose_mcast=False, fused_activation=None, fuse_batch=False, ) self.model_config["PREFILL_MLP_W2_PRG_CONFIG_128"] = ttnn.MatmulMultiCoreReuseMultiCastProgramConfig( compute_with_storage_grid_size=(8, 8), in0_block_w=1, # how much inner dim you take each time out_subblock_h=1, # Must be divisible by per_core_M out_subblock_w=1, # Must be divisible by per_core_N, out_subblock_w * out_subblock_h <= 4 per_core_M=1, # M / TILE_HEIGHT / Grid_Size (dynamic based on seqlen) per_core_N=16, # N / TILE_WIDTH / Grid_Size transpose_mcast=False, fused_activation=None, fuse_batch=False, ) # ============================================================================ # TG (Galaxy) specific MLP memory configs (dictionary access only) # TODO: Migrate these to use getter methods after TTTv2 migration # These configs are used by mlp.py for TG (Galaxy) multi-device setups # ============================================================================ # Sized to the reduce-scatter output width, not the pre-scatter width. See the # LAYOUT HISTORY note in create_galaxy_ff1_out_reduce_scatter_memcfg if Galaxy # decode perf or L1 usage regresses -- Llama-70B went [32, 128] -> [32, 32]. self.model_config["FF1_OUT_REDUCE_SCATTER_MEMCFG"] = create_galaxy_ff1_out_reduce_scatter_memcfg( self.hidden_dim, self.cluster_shape[0], self.cluster_shape[1] ) self.model_config["FF1_OUT_GATHERED_MEMCFG"] = ttnn.create_sharded_memory_config( shape=(32 * 4, self.hidden_dim // 8 // 8), core_grid=ttnn.CoreGrid(y=1, x=8), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) self.model_config["SELF_OUT_REDUCE_SCATTER_MEMCFG"] = ( ttnn.create_sharded_memory_config( shape=(32, 2048 // 8 // 8), # mesh_rows = 8, num_cores=8 core_grid=ttnn.CoreGrid(y=1, x=8), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) if self.dim == 8192 else ttnn.create_sharded_memory_config( shape=(32 * 8, nearest_32(self.dim // 4 // 32)), # mesh_rows = 8 core_grid=ttnn.CoreGrid(y=4, x=8), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) ) self.model_config["FF2_OUT_GATHERED_MEMCFG"] = ttnn.create_sharded_memory_config( shape=(32 * 8, self.dim // 4 // 8), core_grid=ttnn.CoreGrid(y=1, x=8), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) # ============================================================================ # Vision/Multimodal Model Configs (dictionary access only) # TODO: Migrate these to use getter methods after TTTv2 migration # These configs are used by multimodal modules (llama_image_*, llama_cross_*) # ============================================================================ self.model_config["IMAGE_MLP_FC_PROGCFG"] = lambda seq_len, max_seq: self.matmul_config( m=min(seq_len, max_seq), k=self.vision_dim, n=self.vision_hidden_dim // self.num_devices, grid_size=(8, 8), in0_block_w=1, fuse_batch=seq_len <= max_seq, ) self.model_config["IMAGE_MLP_PROJ_PROGCFG"] = lambda seq_len, max_seq: self.matmul_config( m=min(seq_len, max_seq), k=self.vision_hidden_dim // self.num_devices, n=self.vision_dim, grid_size=(8, 8), in0_block_w=1, fuse_batch=seq_len <= max_seq, ) self.model_config["IMAGE_ATTN_QKV_PROGCFG"] = lambda seq_len, max_seq: self.matmul_config( m=min(seq_len, max_seq), k=self.vision_dim, n=(nearest_32(self.vision_head_dim) * self.vision_attn_n_heads * 3) // self.num_devices, # Head dim was padded to nearest 32 grid_size=(8, 8), in0_block_w=1, fuse_batch=seq_len <= max_seq, ) self.model_config["IMAGE_ATTN_OUT_PROGCFG"] = lambda seq_len, max_seq: self.matmul_config( m=min(seq_len, max_seq), k=(nearest_32(self.vision_head_dim) * self.vision_attn_n_heads * 3) // self.num_devices, n=self.vision_dim, grid_size=(8, 8), in0_block_w=1, fuse_batch=seq_len <= max_seq, ) self.model_config["VISION_XATTN_Q_PROGCFG"] = lambda seq_len: self.matmul_config( m=min(seq_len, 1024), k=self.dim, n=(self.head_dim * self.n_heads) // self.num_devices, grid_size=(8, 8), in0_block_w=1, fuse_batch=seq_len <= 1024, ) self.model_config["VISION_XATTN_KV_PROGCFG"] = lambda seq_len, max_seq: self.matmul_config( m=min(seq_len, max_seq), k=self.dim, n=(self.head_dim * self.n_kv_heads) // self.num_devices, grid_size=(8, 8), in0_block_w=1, fuse_batch=seq_len <= max_seq, ) self.model_config["VISION_XATTN_SCORE_PROGCFG"] = lambda seq_len, cache_seq_len: self.matmul_config( m=seq_len, k=self.head_dim, n=cache_seq_len, grid_size=(8, 8), in0_block_w=1, fuse_batch=False, ) self.model_config["VISION_XATTN_OUTPUT_PROGCFG"] = lambda seq_len, cache_seq_len: self.matmul_config( m=seq_len, k=cache_seq_len, n=self.head_dim, grid_size=(8, 8), # in0_block_w=1, # TODO: Remove this when we get non-causal FlashDecode fuse_batch=False, ) self.model_config["VISION_XATTN_DENSE_PROGCFG"] = lambda seq_len: self.matmul_config( m=min(seq_len, 1024), k=self.dim // self.num_devices, n=self.dim, grid_size=(8, 8), in0_block_w=1, fuse_batch=False, ) self.model_config["VISION_PROJ_PROGCFG"] = lambda seq_len: self.matmul_config( m=seq_len, k=self.vision_dim * 6, n=self.dim // self.num_devices, grid_size=(8, 8), in0_block_w=1, fuse_batch=False, ) self.model_config["CROSS_TRANSFORMER_TEXT_OUTPUT_PROGCFG"] = lambda seq_len, max_seq: self.matmul_config( m=min(seq_len, max_seq), k=self.dim, n=self.vocab_size // 8, # Magic number. LM Head always contains 8 splits grid_size=(8, 8), in0_block_w=1, fuse_batch=seq_len <= max_seq, ) def _get_xattn_kv_prefill_mem_cfg(seq_len): # max(1, ...): KV replication (n_kv_heads < num_devices) makes the integer # divide 0, which would build a 0-row shard config; clamp to one KV head/device. M = max(1, self.n_kv_heads // self.num_devices) * seq_len cores_x, cores_y = self.find_grid(M // ttnn.TILE_SIZE) return ttnn.create_sharded_memory_config( ( nearest_32(M // (cores_x * cores_y)), self.head_dim, ), ttnn.CoreGrid(y=cores_y, x=cores_x), ttnn.ShardStrategy.HEIGHT, ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) self.model_config["XATTN_KV_PREFILL_MEM_CFG"] = _get_xattn_kv_prefill_mem_cfg if self.is_multimodal: self.VISION_MAX_MM_SEQ = nearest_32(self.vision_chunk_ntok) self.set_tg_attention_config() self.is_multichip = self.num_devices > 1 self.num_reduce_scatter_links = 1 self.num_all_gather_links = ( 2 if self.is_galaxy else 1 ) # TODO: try out 3 for short axis and 4 for long axis (TG only) <- should work but untested in model self.ccl_dtype = ttnn.bfloat8_b # model specific CCL configs default_ln_ag = {"num_links": 1, "chunks_per_sync": 10, "num_workers_per_link": 2} default_agmm = {"num_links": 1, "chunks_per_sync": 10, "num_workers_per_link": 2} default_mlp_rs = { "num_links": self.num_reduce_scatter_links, "chunks_per_sync": 10, "num_workers_per_link": 2, "rs_memory_config": ttnn.DRAM_MEMORY_CONFIG, } default_sampling_force_argmax = { # Single-chip only: the argmax fast-path speeds up greedy decode on P150 # (~22 -> ~38 t/s/u), but multi-chip meshes are faster on the top-k path. "allow_force_argmax": self.num_devices == 1, "num_links": 1, "chunks_per_sync": 10, "num_workers_per_link": 2, "topology": ttnn.Topology.Linear, } model_specific_ccl_configs = { "Llama-3.1-8B": { "attn_ln_ag": {"num_links": 4, "chunks_per_sync": 10, "num_workers_per_link": 1}, "ffn_ln_ag": {"num_links": 4, "chunks_per_sync": 25, "num_workers_per_link": 1}, "attn_agmm": {"num_links": 4, "chunks_per_sync": 1, "num_workers_per_link": 1}, "mlp_rs": { "num_links": 4, "chunks_per_sync": 1, "num_workers_per_link": 1, "rs_memory_config": ttnn.L1_MEMORY_CONFIG, }, "sampling_force_argmax": { "allow_force_argmax": True, "num_links": 4, "chunks_per_sync": 10, "num_workers_per_link": 2, "topology": ttnn.Topology.Ring, }, } } # Model-specific CCL configs are tuned for Galaxy (TG) with 4 links # Only apply them on Galaxy, otherwise use defaults executed_on_galaxy = self.is_galaxy_cluster if executed_on_galaxy and self.base_model_name in model_specific_ccl_configs: self.model_config["ATTN_LN_AG_CONFIG"] = model_specific_ccl_configs[self.base_model_name]["attn_ln_ag"] self.model_config["FFN_LN_AG_CONFIG"] = model_specific_ccl_configs[self.base_model_name]["ffn_ln_ag"] self.model_config["ATTN_AGMM_CONFIG"] = model_specific_ccl_configs[self.base_model_name]["attn_agmm"] self.model_config["MLP_RS_CONFIG"] = model_specific_ccl_configs[self.base_model_name]["mlp_rs"] self.model_config["SAMPLING_AG_CONFIG"] = model_specific_ccl_configs[self.base_model_name][ "sampling_force_argmax" ] else: self.model_config["ATTN_LN_AG_CONFIG"] = default_ln_ag self.model_config["FFN_LN_AG_CONFIG"] = default_ln_ag self.model_config["ATTN_AGMM_CONFIG"] = default_agmm self.model_config["MLP_RS_CONFIG"] = default_mlp_rs self.model_config["SAMPLING_AG_CONFIG"] = default_sampling_force_argmax logger.info(f"Attention grid: {self.attn_input_grid}") logger.info(f"MLP grid: {self.mlp_core_grid}") logger.info(f"MLP prefill grids @ 32: w1/w3: {self.mlp1_3_grid(32)}, w2: {self.mlp2_grid(32)}") logger.info( f"MLP prefill grids @ max_seq_len({self.max_seq_len}): w1/w3: {self.mlp1_3_grid(self.max_seq_len)}, w2: {self.mlp2_grid(self.max_seq_len)}" ) logger.info(f"LM head grid: {self.lm_head_core_grid}") self.capped_warmup_seq_len = min(self.max_prefill_chunk_size, self.max_seq_len) self.trace_prefill_supported_seq_lens = self.get_trace_prefill_supported_seq_lens() @property def decoders_optimizations(self): """Get the decoders optimizations configuration.""" return self.model_config["DECODERS_OPTIMIZATIONS"] @property def use_fused_all_gather_matmul(self): """Get whether fused all-gather matmul should be used.""" return getattr(self, "_use_fused_all_gather_matmul", False) @property def is_galaxy_8_device_row_submesh(self): """True for the Galaxy DP4 submesh shape that behaves like a T3K row.""" return ( self.mesh_device is not None and self.num_devices == 8 and tuple(self.mesh_device.shape) == (1, 8) and self.is_galaxy_cluster ) def get_warmup_prefill_supported_seq_lens(self): assert ( self.capped_warmup_seq_len > 0 and (self.capped_warmup_seq_len & (self.capped_warmup_seq_len - 1)) == 0 ), f"capped_warmup_seq_len must be a power of 2, but got {self.capped_warmup_seq_len}" DEFAULT_VALUE = self.capped_warmup_seq_len # This dictionary is used to override the default ceil warmup prefill value model_specific_ceil_warmup_lengths = { # Qwen3-32B hangs at 8192, so we cap at 4096 "Qwen3-32B": 4096, } max_seq_len_to_warmup = model_specific_ceil_warmup_lengths.get(self.base_model_name, DEFAULT_VALUE) if max_seq_len_to_warmup > self.capped_warmup_seq_len: max_seq_len_to_warmup = self.capped_warmup_seq_len to_warmup_seq_lens = calculate_prefill_warmup_seq_lens( max_seq_len_to_warmup, self.trace_prefill_supported_seq_lens ) to_warmup_seq_lens = self.filter_warmup_seq_lens(to_warmup_seq_lens) return to_warmup_seq_lens def filter_warmup_seq_lens(self, to_warmup_seq_lens): # TODO: Add more model-specific filtering here # This filtering is based on the current PR's (https://github.com/tenstorrent/tt-metal/pull/33143) sequence lengths that are used for warmup # TODO: https://github.com/tenstorrent/tt-metal/issues/33991 - for P100 only, P150 has assert for ISL > 1K if self.base_model_name == "Llama-3.1-8B" and self.device_name == "P100": for seq_len in to_warmup_seq_lens: if seq_len > 1024: to_warmup_seq_lens = to_warmup_seq_lens[: to_warmup_seq_lens.index(seq_len)] break if self.base_model_name == "Mistral-Small-3.1-24B": to_warmup_seq_lens = [s for s in to_warmup_seq_lens if s <= self.max_seq_len] return to_warmup_seq_lens # ========================================================================= # RESIDUAL MEMORY CONFIGS # ========================================================================= @lru_cache(maxsize=None) def get_residual_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): """Get the memory config for decode residual tensors.""" if mode == Mode.DECODE: if prefetcher is not None: num_residual_worker_cores = 32 if self.num_devices == 4 else 16 return ttnn.create_sharded_memory_config( shape=( 32, self.dim // self.cluster_shape[1] // prefetcher.dynamic_worker_core_grid(num_residual_worker_cores).num_cores(), ), core_grid=prefetcher.dynamic_worker_core_grid(num_residual_worker_cores), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) elif self.is_galaxy: return ttnn.L1_MEMORY_CONFIG else: residual_grid = self.dram_shard_core_grid_for_k(self.dim // self.num_devices) return ttnn.create_sharded_memory_config( ( self.tile_padded_batch_rows, self.dim // residual_grid.num_cores // self.num_devices, ), residual_grid, ttnn.ShardStrategy.WIDTH, ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") # ========================================================================= # MLP PROGRAM AND MEMORY CONFIGS # ========================================================================= @lru_cache(maxsize=None) def get_mlp_input_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): """Get the sharded memory config for MLP input.""" if mode == Mode.DECODE: if self.is_galaxy: return self.get_mlp_act_mem_config(Mode.DECODE) elif prefetcher is not None: return ttnn.create_sharded_memory_config( shape=(32, self.dim // prefetcher.ring_size), core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return ttnn.create_sharded_memory_config( (self.tile_padded_batch_rows, self.dim // self.mlp_core_grid.num_cores), self.mlp_core_grid, ttnn.ShardStrategy.WIDTH, ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_mlp_ff1_3_prg_config(self, mode: Mode, seq_len: int = 1, prefetcher: Prefetcher = None): if mode == Mode.DECODE: if self.dim >= 4096 and self.is_galaxy: return self.matmul_1d_config_from_tensor_shapes( ( 1, 1, 32, self.dim // 4, ), ( 1, 1, self.dim // 4, self.hidden_dim // 8, ), grid=ttnn.CoreGrid(x=8, y=2), overwrite_subblock_h=1, overwrite_subblock_w=1, ) else: if prefetcher is not None: return self.matmul_1d_ring_config( 1, 32, self.dim, self.hidden_dim // self.cluster_shape[1], # Use padded N prefetcher.ring_size, num_global_cb_receivers=prefetcher.num_receiver_cores, ) else: return self.dram_matmul_config( m=self.tile_padded_batch_rows, k=self.dim, n=self.hidden_dim // self.cluster_shape[1], num_cores=self.mlp_core_grid.num_cores, num_workers_per_dram_bank=self.get_dram_sharded_matmul_num_workers( TensorGroup.FF1_FF3, self.hidden_dim // self.cluster_shape[1] ), ) elif mode == Mode.PREFILL: return self.matmul_config( m=min(seq_len, self.prefill_len_cutoff), # 512 if BH, 1024 if WH k=self.dim // self.cluster_shape[0], n=self.hidden_dim // self.cluster_shape[1], grid_size=self.mlp1_3_grid(seq_len), # N150/Llama 8B: default K block 8 exceeded available L1 by # 60,256 bytes with trace-owned tensors live. Four 32-wide tiles # still divide K=4096 and halve the input CB staging vs. 8. # Validated at 512/1024/2048 tokens (MLP PCC > 0.9996); this is # a measured fit, not an assertion that 4 is throughput-optimal. in0_block_w=( 4 if self.device_name == "N150" and self.base_model_name == "Llama-3.1-8B" and seq_len >= self.prefill_len_cutoff else None ), per_core_N=( math.ceil( (self.hidden_dim // self.cluster_shape[1]) / (ttnn.TILE_SIZE * self.dram_shard_grid_width) ) if not self.is_galaxy else None ), ) @lru_cache(maxsize=None) def get_mlp_ff2_prg_config(self, mode: Mode, seq_len: int = 1, prefetcher: Prefetcher = None): if mode == Mode.DECODE: if self.dim >= 4096 and self.is_galaxy: return self.matmul_1d_config_from_tensor_shapes( ( 1, 1, 32, self.hidden_dim // 8, ), ( 1, 1, self.hidden_dim // 8, self.dim // 4, ), grid=ttnn.CoreGrid(x=8, y=2), overwrite_subblock_h=1, overwrite_subblock_w=1, ) else: if prefetcher is not None: return self.matmul_1d_ring_config( 1, 32, self.hidden_dim // self.cluster_shape[1], self.dim, # Use padded N prefetcher.ring_size, num_global_cb_receivers=prefetcher.num_receiver_cores, ) else: return self.dram_matmul_config( m=self.tile_padded_batch_rows, k=self.hidden_dim // self.cluster_shape[1], n=self.dim, num_cores=self.mlp2_core_grid.num_cores, num_workers_per_dram_bank=self.get_dram_sharded_matmul_num_workers(TensorGroup.FF2, self.dim), ) elif mode == Mode.PREFILL: if self.use_minimal_prefill_matmul(seq_len): grid = self.mlp2_grid(seq_len) return ttnn.MinimalMatmulConfig( M_block_size=8, K_block_size=8, N_block_size=8, compute_with_storage_grid_size=ttnn.CoreCoord(grid[0], grid[1]), ) else: return self.matmul_config( m=min(seq_len, self.prefill_len_cutoff), # 512 if BH, 1024 if WH k=self.hidden_dim // (self.cluster_shape[1] if self.is_galaxy else 1), n=self.dim, grid_size=self.mlp2_grid(seq_len), per_core_N=( math.ceil(self.dim / (ttnn.TILE_SIZE * self.dram_shard_grid_width)) if not self.is_galaxy else None ), ) else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_mlp_ff1_3_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): if mode == Mode.DECODE: if prefetcher is not None: return ttnn.create_sharded_memory_config( shape=(32, self.hidden_dim // self.cluster_shape[1] // prefetcher.ring_size), # Use padded N core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return ttnn.L1_WIDTH_SHARDED_MEMORY_CONFIG elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_mlp_ff2_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): if mode == Mode.DECODE: if prefetcher is not None: return ttnn.create_sharded_memory_config( shape=(32, self.dim // prefetcher.ring_size), # Use padded N core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return ttnn.L1_WIDTH_SHARDED_MEMORY_CONFIG elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") # NOTE: Cannot use @lru_cache here because tensor parameter would cause memory leak # by keeping references to all tensors passed to this function def get_mlp_ff2_all_reduce_mem_config(self, mode: Mode, tensor: ttnn.Tensor): if mode == Mode.DECODE: if self.is_galaxy: return ( ttnn.create_sharded_memory_config( shape=(32, self.dim // 8 // 4), # shard_grid_cores = 8, num_devices=4 core_grid=ttnn.CoreRangeSet({ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(7, 0))}), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) if self.dim == 8192 else ttnn.create_sharded_memory_config( shape=(32 * 8, self.dim // 4 // 8), core_grid=ttnn.CoreGrid(y=1, x=8), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) ) else: return tensor.memory_config() elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_mlp_binary_mult_mem_config(self, mode: Mode): """Get the memory config for MLP binary mult (w2 input) - replaces SHARDED_MLP2_INPUT_MEMCFG.""" if mode == Mode.DECODE: return ttnn.create_sharded_memory_config( ( 32 if self.is_galaxy else self.tile_padded_batch_rows, self.hidden_dim // self.cluster_shape[1] // self.mlp2_core_grid.num_cores, ), self.mlp2_core_grid, ttnn.ShardStrategy.WIDTH, ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) elif mode == Mode.PREFILL: return None else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_mlp_output_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): if mode == Mode.DECODE: if prefetcher is not None: num_mlp_output_cores = 32 if self.num_devices == 4 else 16 return ttnn.create_sharded_memory_config( shape=(1, 1, 32, self.dim // self.cluster_shape[1] // num_mlp_output_cores), core_grid=prefetcher.dynamic_worker_core_grid(num_mlp_output_cores), strategy=ttnn.ShardStrategy.WIDTH, use_height_and_width_as_shard_shape=True, ) elif self.is_galaxy: return ttnn.create_sharded_memory_config( shape=(32, nearest_32(self.dim // (8 * self.lm_head_core_grid.y) // 4)), core_grid=ttnn.CoreGrid(y=self.lm_head_core_grid.y, x=8), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return self.get_residual_mem_config(mode, None) elif mode == Mode.PREFILL: return None else: raise ValueError(f"Invalid mode: {mode}") # NOTE: get_mlp_act_mem_config is a TG-specific MLP to memory config @lru_cache(maxsize=None) def get_mlp_act_mem_config(self, mode: Mode): """Get the memory config for MLP activation (TG specific).""" if mode == Mode.DECODE: if self.dim >= 4096: return ttnn.create_sharded_memory_config( shape=(32, self.dim // 4 // 16), core_grid=ttnn.CoreGrid(x=8, y=2), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: full_grid = ttnn.CoreRangeSet({ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(7, 7))}) return ttnn.MemoryConfig( ttnn.TensorMemoryLayout.WIDTH_SHARDED, ttnn.BufferType.L1, ttnn.ShardSpec(full_grid, [32, nearest_32(56)], ttnn.ShardOrientation.ROW_MAJOR), ) elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") # ========================================================================= # ATTENTION PROGRAM AND MEMORY CONFIGS # ========================================================================= @lru_cache(maxsize=None) def get_attn_sdpa_prefill_program_config(self, seq_len: int = 1, chunk_start_idx: int = None): """Get the SDPA program config for prefill mode.""" # Sequence length and chunk start index are both required for prefill # Chunk values based on what works best empirically # We want 256 if seqlen >= 2048 else 64. BUT: # SPDA limitation: chunk_start_idx must be a multiple of q_chunk_size # Here (x & -x) is the highest power of 2 that divides x. # When chunk_start_idx=0, we use default values since 0 is a multiple of any number. q_chunk = ( 256 if seq_len >= 2048 and (chunk_start_idx is None or chunk_start_idx == 0) else ( 64 if seq_len < 2048 and (chunk_start_idx is None or chunk_start_idx == 0) else ( min(256, chunk_start_idx & -chunk_start_idx) if seq_len >= 2048 else min(64, chunk_start_idx & -chunk_start_idx) ) ) ) # Workaround for https://github.com/tenstorrent/tt-metal/issues/35225: k_chunk = ( 256 if seq_len >= 2048 and (chunk_start_idx is None or chunk_start_idx == 0) else ( 64 if seq_len < 2048 and (chunk_start_idx is None or chunk_start_idx == 0) else ( min(256, chunk_start_idx & -chunk_start_idx) if seq_len >= 2048 else min(64, chunk_start_idx & -chunk_start_idx) ) ) ) return ttnn.SDPAProgramConfig( compute_with_storage_grid_size=(8, 8), exp_approx_mode=False, q_chunk_size=q_chunk, k_chunk_size=k_chunk, ) @lru_cache(maxsize=None) def get_attn_sdpa_decode_program_config(self, prefetcher: Prefetcher = None): """Get the SDPA program config for decode mode. The KV chunk sizes come from self.sdpa_decode_{q,k}_chunk_size (default 0, i.e. the op auto-selects). Architectures that need an explicit small chunk set these in _set_model_specific_params(); e.g. Gemma-2 (head_dim=256), where the auto chunk drives the flash-decode op into a multi-chunk path whose cross-chunk reduction is numerically wrong for head_dim=256, producing a sharp accuracy cliff once the context exceeds one chunk. """ q_chunk, k_chunk = self.sdpa_decode_q_chunk_size, self.sdpa_decode_k_chunk_size if prefetcher is not None: sdpa_grid_size = (8, 8) start_core = ttnn.CoreCoord(1, 0) num_sdpa_cores = sdpa_grid_size[0] * sdpa_grid_size[1] return ttnn.SDPAProgramConfig( compute_with_storage_grid_size=sdpa_grid_size, sub_core_grids=ttnn.num_cores_to_corerangeset_in_subcoregrids( start_core, num_sdpa_cores, prefetcher.all_worker_cores_range_set, row_wise=True ), exp_approx_mode=False, q_chunk_size=q_chunk, k_chunk_size=k_chunk, ) else: return ttnn.SDPAProgramConfig( compute_with_storage_grid_size=(8, 8), exp_approx_mode=False, q_chunk_size=q_chunk, k_chunk_size=k_chunk, ) @lru_cache(maxsize=None) def get_attn_sdpa_program_config( self, mode: Mode, seq_len: int = 1, chunk_start_idx: int = None, prefetcher: Prefetcher = None ): """Get the SDPA program config for attention.""" if mode == Mode.DECODE: return self.get_attn_sdpa_decode_program_config(prefetcher) elif mode == Mode.PREFILL: return self.get_attn_sdpa_prefill_program_config(seq_len, chunk_start_idx) else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_input_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): if mode == Mode.DECODE: if prefetcher is not None: return ttnn.create_sharded_memory_config( shape=(32, self.dim // self.cluster_shape[0] // prefetcher.ring_size), core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) elif self.is_galaxy: return ttnn.create_sharded_memory_config( shape=(32, nearest_32(self.dim // (8 * self.lm_head_core_grid.y) // 4)), core_grid=ttnn.CoreGrid(y=self.lm_head_core_grid.y, x=8), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return ttnn.create_sharded_memory_config( ( self.tile_padded_batch_rows, self.dim // self.attn_input_grid.num_cores, ), # Shard shape: [32, 128] -> 1 shard per core self.attn_input_grid, ttnn.ShardStrategy.WIDTH, ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") # ========================================================================= # ATTENTION PROGRAM AND MEMORY CONFIGS (continued) # QKV, WO, All-Reduce, All-Gather configs # ========================================================================= @lru_cache(maxsize=None) def get_attn_qkv_program_config(self, mode: Mode, seq_len: int = 1, prefetcher: Prefetcher = None): """Get the program config for the QKV matmul in attention.""" if mode == Mode.DECODE: if prefetcher is not None: return self.matmul_1d_ring_config( 1, 32, self.dim // self.cluster_shape[0], self.qkv_size // self.cluster_shape[1], # Use padded N prefetcher.ring_size, num_global_cb_receivers=prefetcher.num_receiver_cores, untilize_out=True, ) else: return self.dram_matmul_config( m=self.tile_padded_batch_rows, k=self.dim, n=self.qkv_size // self.num_devices, num_cores=self.attn_input_grid.num_cores, num_workers_per_dram_bank=self.get_dram_sharded_matmul_num_workers( TensorGroup.WQKV, self.qkv_size // self.num_devices ), ) elif mode == Mode.PREFILL: self.MAX_QKV_MM_SEQ_LEN = 2048 if self.use_minimal_qkv_prefill_matmul(seq_len): return ttnn.MinimalMatmulConfig( M_block_size=8, K_block_size=8, N_block_size=8, compute_with_storage_grid_size=ttnn.CoreCoord(8, 10) if is_blackhole() else ttnn.CoreCoord(8, 8), ) else: return ttnn.MatmulMultiCoreReuseMultiCastProgramConfig( compute_with_storage_grid_size=(8, 10) if is_blackhole() else (8, 8), in0_block_w=1, # FIXME: optimize this config for prefill, careful use DI_DT_WORKAROUND if necessary out_subblock_h=1, # Must be divisible by per_core_M out_subblock_w=1, # Must be divisible by per_core_N, out_subblock_w * out_subblock_h <= 4 # This branch is only reached when use_minimal_qkv_prefill_matmul() is False, # i.e. seq_len <= 128, so the former `8 if seq_len >= MAX_QKV_MM_SEQ_LEN` arm # (MAX_QKV_MM_SEQ_LEN == 2048) was unreachable and has been removed. At every # reachable seq_len the max(1, ...) floor makes this 1. # NOTE: P100 runs OOM in L1 with a larger per_core_M (workaround for issue #50656). per_core_M=1 if self.device_name == "P100" else max(1, math.ceil(seq_len / ttnn.TILE_SIZE / 8)), per_core_N=math.ceil( self.qkv_size / self.cluster_shape[1] / 32 / self.dram_shard_grid_width ), # N / TILE_WIDTH / grid width transpose_mcast=False, fused_activation=None, fuse_batch=seq_len <= self.MAX_QKV_MM_SEQ_LEN, ) else: raise ValueError(f"Invalid mode: {mode}") def use_minimal_prefill_matmul(self, seq_len: int) -> bool: # Qwen's 128-token prompts can run singly or be flattened into a larger # batched prefill. return seq_len > 128 or (seq_len == 128 and self.base_model_name == "Qwen3-32B" and self.device_name == "T3K") def use_minimal_qkv_prefill_matmul(self, seq_len: int) -> bool: if self.use_minimal_prefill_matmul(seq_len): return True # The regular 128-token QKV prefill matmul over-allocates L1 on Llama 8B # Galaxy DP4; use minimal matmul only for that row-submesh case. return self.base_model_name == "Llama-3.1-8B" and seq_len == 128 and self.is_galaxy_8_device_row_submesh @lru_cache(maxsize=None) def get_attn_qkv_mm_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): """Get the memory config for QKV matmul output in attention.""" if mode == Mode.DECODE: if prefetcher is not None: qkv_out_shard_shape_ring = ( 32, self.qkv_size // self.cluster_shape[1] // prefetcher.ring_size, ) # Use padded N return ttnn.create_sharded_memory_config( shape=qkv_out_shard_shape_ring, core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return ttnn.L1_WIDTH_SHARDED_MEMORY_CONFIG elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_qkv_all_reduce_output_mem_config(self, mode: Mode, mesh_cols: int = 1, prefetcher: Prefetcher = None): """Get the memory config for QKV all-reduce output in attention.""" if mode == Mode.DECODE: if prefetcher is not None: # When using prefetcher, return the current memory config (pass-through) return None # Caller should use tensor's memory_config() else: num_cores = 40 if self.dim == 8192 else (24 if self.dim == 4096 else (20 if self.dim == 3072 else 12)) qkv_core_grid = ( self.dram_shard_core_grid_for_k(self.dim) if not self.is_galaxy else num_to_coregrid(num_cores) ) shard_height = ttnn.TILE_SIZE * mesh_cols shard_width = self.dim // qkv_core_grid.num_cores if not self.is_galaxy else 32 return ttnn.create_sharded_memory_config( ( shard_height, shard_width, ), # Shard shape: [32, 128] -> 1 shard per core core_grid=qkv_core_grid, strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_create_head_input_mem_config(self, mode: Mode): """Get the memory config for create_head input (TG specific).""" if mode == Mode.DECODE: return self.model_config["CREATE_HEAD_INPUT_MEMCFG"] elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_create_head_output_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): """Get the memory config for create_qkv_heads output in attention.""" if mode == Mode.DECODE: if prefetcher is not None: return ttnn.MemoryConfig( ttnn.TensorMemoryLayout.HEIGHT_SHARDED, ttnn.BufferType.L1, ttnn.ShardSpec( prefetcher.all_worker_cores_range_set, [32, self.head_dim], ttnn.ShardOrientation.ROW_MAJOR, ), ) else: # CREATE_QKV_DECODE_SHARD equivalent if is_blackhole(): return ttnn.create_sharded_memory_config( shape=(ttnn.TILE_SIZE, self.head_dim), core_grid=ttnn.CoreGrid(y=4, x=8), strategy=ttnn.ShardStrategy.HEIGHT, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return ttnn.L1_HEIGHT_SHARDED_MEMORY_CONFIG elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_sdpa_output_mem_config( self, mode: Mode, batch_size_per_device_group: int = 1, prefetcher: Prefetcher = None ): """Get the memory config for SDPA output in attention.""" if mode == Mode.DECODE: if prefetcher is not None: start_core = ttnn.CoreCoord(1, 0) return ttnn.create_sharded_memory_config( shape=(math.ceil(self.n_local_heads / ttnn.TILE_SIZE) * ttnn.TILE_SIZE, self.head_dim), core_grid=ttnn.num_cores_to_corerangeset_in_subcoregrids( start_core, batch_size_per_device_group, prefetcher.all_worker_cores_range_set, row_wise=True, ), strategy=ttnn.ShardStrategy.HEIGHT, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return ttnn.create_sharded_memory_config( shape=(math.ceil(self.n_local_heads / ttnn.TILE_SIZE) * ttnn.TILE_SIZE, self.head_dim), core_grid=ttnn.CoreRangeSet({num_to_corerange(batch_size_per_device_group)}), strategy=ttnn.ShardStrategy.HEIGHT, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_concat_heads_output_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): """Get the memory config for attention concat_heads output before WO matmul.""" if mode == Mode.DECODE: if prefetcher is not None: wo_out_shard_shape_ring = ( 32, self.dim // self.cluster_shape[1] // prefetcher.ring_size, ) # Use padded N return ttnn.create_sharded_memory_config( shape=wo_out_shard_shape_ring, core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return ttnn.MemoryConfig( ttnn.TensorMemoryLayout.WIDTH_SHARDED, ttnn.BufferType.L1, ttnn.ShardSpec( num_to_core_range_set(self.num_devices), [ self.tile_padded_batch_rows, self.dim // self.num_devices, ], ttnn.ShardOrientation.ROW_MAJOR, ), ) elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_all_gather_output_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): """Get the memory config for attention all-gather output.""" if mode == Mode.DECODE: if prefetcher is not None: return ttnn.create_sharded_memory_config( shape=(32, self.dim // prefetcher.ring_size), # Use padded N core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: # All gather matmuls currently only supported on T3K # We need it sharded on num_cores = num_devices return ttnn.MemoryConfig( ttnn.TensorMemoryLayout.WIDTH_SHARDED, ttnn.BufferType.L1, ttnn.ShardSpec( num_to_core_range_set(self.num_devices), [ self.tile_padded_batch_rows, self.dim // self.num_devices, ], ttnn.ShardOrientation.ROW_MAJOR, ), ) elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_all_gather_matmul_program_config(self, mode: Mode, prefetcher: Prefetcher = None): """Get the program config for fused all-gather matmul in attention.""" if mode == Mode.DECODE: if prefetcher is not None: k_wo = self.dim n_wo = self.dim // self.cluster_shape[1] return self.matmul_1d_ring_config( 1, 32, k_wo, n_wo, prefetcher.ring_size, num_global_cb_receivers=prefetcher.num_receiver_cores, ) else: if self.use_fused_all_gather_matmul: do_core_grid_size = (8, 1) do_per_core_N = ( self.dim // self.num_devices // ttnn.TILE_SIZE // (do_core_grid_size[0] * do_core_grid_size[1]) ) return ttnn.MatmulMultiCoreReuseMultiCast1DProgramConfig( compute_with_storage_grid_size=do_core_grid_size, in0_block_w=self.dim // ttnn.TILE_SIZE // (do_core_grid_size[0] * do_core_grid_size[1]), # [32 x 8k] x [8k x 1k] = [32 x 1k] out_subblock_h=1, out_subblock_w=get_out_subblock_w( do_per_core_N, out_subblock_h=1 ), # Max out_subblock_w = 4, needs to be divisible by per_core_N per_core_M=self.tile_padded_batch_rows // ttnn.TILE_SIZE, per_core_N=do_per_core_N, fuse_batch=True, fused_activation=None, mcast_in0=True, ) else: return None elif mode == Mode.PREFILL: return None else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_output_program_config(self, mode: Mode): """Get the program config for attention output matmul (replaces ATTN_OUTPUT_PROGCFG).""" if mode == Mode.DECODE: if self.is_galaxy: return None else: return self.dram_matmul_config( m=self.tile_padded_batch_rows, k=(self.n_heads * self.head_dim) // self.num_devices, n=self.dim, num_cores=self.n_heads // self.num_devices, num_workers_per_dram_bank=self.get_dram_sharded_matmul_num_workers(TensorGroup.WO, self.dim), ) elif mode == Mode.PREFILL: return None else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_wo_program_config(self, mode: Mode, seq_len: int = 1, prefetcher: Prefetcher = None): """Get the program config for WO (dense output) matmul in attention.""" if mode == Mode.DECODE: if prefetcher is not None: k_wo = self.dim n_wo = self.dim // self.cluster_shape[1] return self.matmul_1d_ring_config( 1, 32, k_wo, n_wo, prefetcher.ring_size, num_global_cb_receivers=prefetcher.num_receiver_cores, ) elif self.is_galaxy: return None # TG uses core_grid parameter instead else: return self.get_attn_output_program_config(Mode.DECODE) elif mode == Mode.PREFILL: dram_sharded_wo = not (self._use_fused_all_gather_matmul or self.is_galaxy) n_dim = ( self.dim // self.cluster_shape[1] if self.is_galaxy else ( 1024 if self.num_devices == 8 and getattr(self, "_use_t3k_fused_agmm_config", not self.is_galaxy_cluster) and not is_blackhole() and 1024 % (self.dim // self.num_devices) == 0 else self.dim ) ) # Attention output is not necessarily the same dimension as the self.dim, e.g. in Mistral k_dim = ( (self.n_heads * self.head_dim) // self.cluster_shape[0] if self.is_galaxy else (self.n_heads * self.head_dim) // self.num_devices ) return self.matmul_config( m=min(seq_len, 1024), k=k_dim, n=n_dim, grid_size=self.find_prefill_grid(self.prefill_rows, k_dim // ttnn.TILE_SIZE), in0_block_w=1 if self.is_galaxy else None, fuse_batch=seq_len <= 1024, per_core_N=( math.ceil(n_dim / (ttnn.TILE_SIZE * self.dram_shard_grid_width)) if dram_sharded_wo else None ), ) else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_wo_output_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): """Get the memory config for WO matmul output in attention.""" if mode == Mode.DECODE: if prefetcher is not None: wo_out_shard_shape_ring = (32, self.dim // prefetcher.ring_size) return ttnn.create_sharded_memory_config( shape=wo_out_shard_shape_ring, core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return ttnn.L1_WIDTH_SHARDED_MEMORY_CONFIG elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_dense_output_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): """Get the memory config for dense output (after all-reduce) in attention.""" if mode == Mode.DECODE: if self.ccl_topology() == ttnn.Topology.Ring and prefetcher is None: return self.get_residual_mem_config(Mode.DECODE, None) else: if prefetcher is not None: wo_out_shard_shape_ring = ( 32, self.dim // self.cluster_shape[1] // prefetcher.ring_size, ) return ttnn.create_sharded_memory_config( shape=wo_out_shard_shape_ring, core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return self.get_residual_mem_config(Mode.DECODE, None) elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_all_reduce_output_mem_config( self, mode: Mode, hidden_size: int = 0, mesh_rows: int = 1, prefetcher: Prefetcher = None ): """Get the memory config for attention all-reduce output (TG specific configs).""" if mode == Mode.DECODE: if self.is_galaxy: if hidden_size == 8192: return self.model_config["SELF_OUT_REDUCE_SCATTER_MEMCFG"] else: return self.model_config["SELF_OUT_GATHERED_MEMCFG"](mesh_rows) else: if prefetcher is not None: return ttnn.create_sharded_memory_config( shape=(1, 1, 32, self.dim // self.cluster_shape[1] // prefetcher.ring_size), core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, use_height_and_width_as_shard_shape=True, ) else: return self.get_residual_mem_config(Mode.DECODE, None) elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_gather_users_mem_config(self, mode: Mode, mesh_cols: int = 1, prefetcher: Prefetcher = None): """Get the memory config for gather users in attention (TG path).""" if mode == Mode.DECODE: if prefetcher is not None: wo_out_shard_shape_ring = ( 32, self.dim // self.cluster_shape[1] // prefetcher.ring_size, ) return ttnn.create_sharded_memory_config( shape=wo_out_shard_shape_ring, core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return self.model_config["GATHER_USERS_MEMCFG"](mesh_cols) elif mode == Mode.PREFILL: return ttnn.DRAM_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_attn_kv_prefill_mem_config(self, seq_len: int = 1): """Get the memory config for KV cache fill during prefill.""" # max(1, ...): KV replication (n_kv_heads < num_devices) makes the integer divide 0, # which would build a 0-row shard config; clamp to one KV head per device. return ttnn.create_sharded_memory_config( ((max(1, self.n_kv_heads // self.cluster_shape[1]) * seq_len // (8 * 8)), self.head_dim), ttnn.CoreGrid(y=8, x=8), ttnn.ShardStrategy.HEIGHT, ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) # ========================================================================= # NORM CONFIGS # ========================================================================= @lru_cache(maxsize=None) def get_norm_config(self, norm_type: str, mode: Mode, prefetcher: Prefetcher = None): """Get the norm config dict for attention, ff, or lm_head norms.""" prefetcher_norm_grid = ttnn.CoreGrid(y=8, x=4) match norm_type: case "attn": if mode == Mode.DECODE and prefetcher is not None: return { "sharded_program_config": self.create_sharded_norm_config(prefetcher_norm_grid), "sharded_output_config": ttnn.create_sharded_memory_config( shape=( 32, self.dim // self.cluster_shape[0] // prefetcher.dynamic_worker_core_grid(32).num_cores(), ), core_grid=prefetcher.dynamic_worker_core_grid(32), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ), "output_mem_config": ttnn.create_sharded_memory_config( shape=(32, self.dim // self.cluster_shape[0] // prefetcher.ring_size), core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ), } else: return { "sharded_program_config": self.create_sharded_norm_config(self.attn_input_grid), "sharded_output_config": self.get_attn_input_mem_config(Mode.DECODE), "output_mem_config": None, } case "ff": if mode == Mode.DECODE and prefetcher is not None: return { "sharded_program_config": self.create_sharded_norm_config(prefetcher_norm_grid), "sharded_output_config": ttnn.create_sharded_memory_config( shape=(32, self.dim // prefetcher.dynamic_worker_core_grid(32).num_cores()), core_grid=prefetcher.dynamic_worker_core_grid(32), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ), "output_mem_config": ttnn.create_sharded_memory_config( shape=(32, self.dim // self.cluster_shape[0] // prefetcher.ring_size), core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ), } else: return { "sharded_program_config": self.create_sharded_norm_config(self.mlp_core_grid), "sharded_output_config": ttnn.create_sharded_memory_config( (self.tile_padded_batch_rows, self.dim // self.mlp_core_grid.num_cores), self.mlp_core_grid, ttnn.ShardStrategy.WIDTH, ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ), "output_mem_config": None, } case "lm_head": if mode == Mode.DECODE and prefetcher is not None: return { "sharded_program_config": self.create_sharded_norm_config(prefetcher_norm_grid), "sharded_output_config": ttnn.create_sharded_memory_config( shape=(32, self.dim // prefetcher.dynamic_worker_core_grid(32).num_cores()), core_grid=prefetcher.dynamic_worker_core_grid(32), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ), "output_mem_config": None, } else: return { "sharded_program_config": self.create_sharded_norm_config(self.lm_head_core_grid), "sharded_output_config": self.get_lm_head_input_mem_config(mode, None), "output_mem_config": None, } case _: raise ValueError(f"Invalid norm_type: {norm_type}") # ========================================================================= # LM HEAD PROGRAM AND MEMORY CONFIGS AND HELPER METHODS # ========================================================================= def get_lm_head_max_columns_per_device(self, core_grid: ttnn.CoreGrid, prefetcher: Prefetcher = None): # 128256 comes from original llama 3 vocab size. 128256 / 4 was experimentally the maximum columns that worked per device. # The LM head for that was on 48 cores, so we know 128256 / 4 / 48 = 668 columns per core is close to the L1 limit. # FIXME: Update blackhole figure to be per-core as well. LLAMA_VOCAB_SIZE = 128256 NUM_LM_HEAD_CORES = 48 NUM_LM_HEAD_COLUMNS = 8 max_columns_per_device = ( 668 * core_grid.num_cores ) # 668 columns per core is close to the L1 limit in LM head matmul. if is_blackhole(): if self.num_devices == 4: max_columns_per_device = LLAMA_VOCAB_SIZE // self.num_devices // NUM_LM_HEAD_COLUMNS elif self.num_devices == 8: max_columns_per_device = LLAMA_VOCAB_SIZE // self.num_devices // (NUM_LM_HEAD_COLUMNS * 2) else: max_columns_per_device = LLAMA_VOCAB_SIZE // NUM_LM_HEAD_COLUMNS if prefetcher is not None: return math.ceil(max_columns_per_device / (ttnn.TILE_SIZE * prefetcher.ring_size)) * ( ttnn.TILE_SIZE * prefetcher.ring_size ) else: return max_columns_per_device @lru_cache(maxsize=None) def get_lm_head_input_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): """Get the memory config for LM head input.""" if mode == Mode.DECODE: if prefetcher is not None: return ttnn.create_sharded_memory_config( shape=(32, self.dim // prefetcher.ring_size), core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return ttnn.create_sharded_memory_config( ( self.tile_padded_batch_rows, nearest_32((self.dim // (4 if self.is_galaxy else 1)) // self.lm_head_core_grid.num_cores), ), # Shard shape: [32, 128] -> 1 shard per core self.lm_head_core_grid, ttnn.ShardStrategy.WIDTH, ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) elif mode == Mode.PREFILL: return ttnn.create_sharded_memory_config( ( self.tile_padded_batch_rows, nearest_32((self.dim // (4 if self.is_galaxy else 1)) // self.lm_head_core_grid.num_cores), ), # Shard shape: [32, 128] -> 1 shard per core self.lm_head_core_grid, ttnn.ShardStrategy.WIDTH, ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_lm_head_program_config(self, split_size: int = 1, prefetcher: Prefetcher = None): """Get the program config for LM head matmul.""" if prefetcher is not None: return self.matmul_1d_ring_config( 1, 32, self.dim, self.max_columns_per_device_lm_head, prefetcher.ring_size, prefetch=False, num_global_cb_receivers=1, untilize_out=True, ) else: return self.dram_matmul_config( self.tile_padded_batch_rows, self.dim, split_size, self.lm_head_core_grid.num_cores, ) @lru_cache(maxsize=None) def get_lm_head_output_mem_config(self, mode: Mode, prefetcher: Prefetcher = None): """Get the memory config for LM head output.""" if mode == Mode.DECODE: if prefetcher is not None: return ttnn.create_sharded_memory_config( shape=(32, self.max_columns_per_device_lm_head // prefetcher.ring_size), core_grid=prefetcher.to_core_range_set( prefetcher.receiver_cores(sender_active=True, receiver_active=True) ), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: return ttnn.L1_WIDTH_SHARDED_MEMORY_CONFIG elif mode == Mode.PREFILL: return ttnn.L1_WIDTH_SHARDED_MEMORY_CONFIG else: raise ValueError(f"Invalid mode: {mode}") @lru_cache(maxsize=None) def get_lm_head_sharded_output_mem_config(self, prefetcher: Prefetcher = None): """Get the output memory config for LM head (after sharded_to_interleaved).""" if prefetcher is not None: return ttnn.DRAM_MEMORY_CONFIG else: return ttnn.L1_MEMORY_CONFIG @lru_cache(maxsize=None) def get_lm_head_reshard_mem_config(self, prefetcher: Prefetcher = None): """Get the memory config for LM head output resharding.""" if prefetcher is None: return ttnn.L1_MEMORY_CONFIG lm_head_output_core_range_set = ttnn.CoreRangeSet( [ ttnn.CoreRange(ttnn.CoreCoord(1, 0), ttnn.CoreCoord(prefetcher.num_receiver_cores, 7)), ttnn.CoreRange(ttnn.CoreCoord(8, 0), ttnn.CoreCoord(8 + prefetcher.num_receiver_cores - 1, 7)), ] ) def next_power_of_2(n): if n <= 0: return 1 return 1 << ((n - 1).bit_length()) lm_head_size = ( (self.padded_vocab_size // self.num_devices) if self.padded_vocab_size else math.ceil(self.vocab_size / self.num_devices) ) lm_head_size_per_device = next_power_of_2(lm_head_size) return ttnn.create_sharded_memory_config( shape=(32, lm_head_size_per_device // prefetcher.ring_size), core_grid=prefetcher.to_core_range_set(prefetcher.receiver_cores(sender_active=True, receiver_active=True)), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) # NOTE: These attention helpers are placed here for historical reasons def get_sharded_wo_ring_mem_config(self): """Get the memory config for WO weights in ring mode.""" wo_shape_ring = ( self.dim // self.cluster_shape[0], self.dim // self.cluster_shape[1], ) return self.create_dram_sharded_mem_config( k=wo_shape_ring[0], n=wo_shape_ring[1], ) def get_attn_weights_layout(self): """Get the layout for attention weights.""" return ttnn.TILE_LAYOUT # ========================================================================= # UTILITY METHODS # ========================================================================= def get_max_prefill_chunk_size(self): # Set the max number of tokens for each prefill chunk based on the model and device max_prefill_chunk_size_div1024 = os.getenv("MAX_PREFILL_CHUNK_SIZE") if max_prefill_chunk_size_div1024 is None: # TODO Improve this to be more general to more devices and models MAX_PREFILL_CHUNK_SIZES_DIV1024 = { "Llama-3.2-1B": {"N150": 128, "N300": 128, "T3K": 128, "TG": 128, "P150x4": 128}, "Llama-3.2-3B": {"N150": 8, "N300": 128, "T3K": 128, "TG": 128, "P150x4": 128}, "Llama-3.1-8B": {"N150": 4, "N300": 64, "T3K": 128, "TG": 128, "P150x4": 128}, "Llama-3.2-11B": {"N150": 4, "N300": 64, "T3K": 128, "TG": 128, "P150x4": 128}, "Llama-3.1-70B": {"N150": None, "N300": None, "T3K": 32, "TG": 128, "P150x4": 128}, "Llama-3.2-90B": {"N150": None, "N300": None, "T3K": 32, "TG": 128, "P150x4": 128}, "DeepSeek-R1-Distill-Llama-70B": {"N150": None, "N300": None, "T3K": 32, "TG": 128, "P150x4": 128}, "Qwen2.5-7B": {"N150": 4, "N300": 32, "T3K": 128, "TG": 128, "P150x4": 128}, "Qwen2.5-32B": {"N150": None, "N300": None, "T3K": 64, "TG": 128, "P150x4": 128, "P150x8": 128}, # Large single-chunk prefill matters for EXAONE: chunked prefill is # unsupported on sliding-window layers (48 of its 64 layers). "EXAONE-4.5-33B": {"P150x8": 128}, "Qwen2.5-Coder-32B": {"N150": None, "N300": None, "P150x4": 128}, "Qwen2.5-72B": {"N150": None, "N300": None, "T3K": 16, "TG": 128, "P150x4": 128, "P150x8": 128}, "Qwen2.5-VL-3B": {"N150": 128, "N300": 128, "T3K": None, "TG": None, "P150x4": None}, "Qwen2.5-VL-7B": {"N150": 64, "N300": 128, "T3K": None, "TG": None, "P150x4": None}, "Qwen2.5-VL-32B": {"N150": None, "N300": None, "T3K": 64, "TG": None, "P150x4": None}, "Qwen2.5-VL-72B": {"N150": None, "N300": None, "T3K": 32, "TG": None, "P150x4": None}, "Qwen3-VL-32B": {"N150": None, "N300": None, "T3K": 64, "TG": None, "P150x4": None, "P150x8": 64}, "DeepSeek-R1-Distill-Qwen-14B": {"N150": 4, "N300": 64, "T3K": 128, "TG": None, "P150x4": None}, "Phi-3.5-mini-instruct": {"N150": 128, "N300": 128, "T3K": 128, "TG": 128, "P150x4": 128}, "Phi-3-mini-128k-instruct": {"N150": 32, "N300": 64, "T3K": 128, "TG": 128, "P150x4": 128}, "QwQ-32B": {"N150": None, "N300": None, "T3K": 64, "TG": 128, "P150x4": 128}, "Qwen3-32B": {"N150": None, "N300": None, "T3K": 64, "TG": 128, "P150x4": 128}, "Qwen3-Embedding-8B": {"N150": 4, "N300": 64, "T3K": 128, "TG": 128, "P150x4": 128}, "Phi-4": {"N150": 4, "N300": 64, "T3K": 128, "TG": 128, "P150x4": 128}, "Mistral-Small-3.1-24B": {"N150": 8, "N300": 128, "T3K": 128, "TG": 128, "P150x4": 128}, "gemma-2-2b": { "N150": 32, "N300": 32, "T3K": 32, "TG": 32, "P100": 32, "P150": 32, "P150x4": 32, }, "gemma-2-9b": { # No N150 entry: the full 42-layer model exhausts single-chip DRAM # during weight caching, so gemma-2-9b is not supported on N150. "N300": 32, "T3K": 32, "TG": 32, "P100": 32, "P150": 32, "P150x4": 32, }, "gemma-3-1b": {"N150": 32, "N300": 32, "T3K": 32, "TG": 32, "P150x4": 32}, "gemma-3-4b": {"N150": 128, "N300": 128, "T3K": 128, "TG": 128, "P150x4": 128}, "medgemma-4b": {"N150": 128, "N300": 128, "T3K": 128, "TG": 128, "P150x4": 128}, "gemma-3-27b": {"N150": 128, "N300": 128, "T3K": 128, "TG": 128, "P150x4": 128}, "medgemma-27b": {"N150": 128, "N300": 128, "T3K": 128, "TG": 128, "P150x4": 128}, "LFM2.5-VL-1.6B": {"N150": 128, "N300": 128, "T3K": 128, "TG": None, "P150x4": 128}, } try: max_prefill_chunk_size_div1024 = MAX_PREFILL_CHUNK_SIZES_DIV1024[self.base_model_name][self.device_name] except KeyError: logger.warning( f"Unknown model {self.model_name} on device {self.device_name}, setting MAX_PREFILL_CHUNK_SIZE to 4 for compatibility" ) logger.warning( f"Try setting MAX_PREFILL_CHUNK_SIZE to larger powers of 2 up to e.g. 128 for faster performance (if you run out of L1 memory it was too high)" ) max_prefill_chunk_size_div1024 = 4 assert ( max_prefill_chunk_size_div1024 is not None ), f"Unsupported model {self.model_name} on device {self.device_name}" else: max_prefill_chunk_size_div1024 = int(max_prefill_chunk_size_div1024) return max_prefill_chunk_size_div1024 * 1024 def get_trace_prefill_supported_seq_lens(self): default_supported_seq_lens = { "N150": [128], "N300": [128, 1024], "T3K": [128, 1024], "TG": [128, 1024], "P150": [128, 1024], "P300": [128, 1024], "P150x4": [128, 1024], "P150x8": [128, 1024], } # TODO: If no specific sequence lengths are listed for a model and device, the default one will be used (from the default_supported_seq_lens dictionary) model_specific_supported_seq_lens = { "Llama-3.1-8B": { "P100": [128, 1024], "N150": [128, 1024], "N300": [128, 1024, 2048, 4096, 8192], "T3K": [128, 1024, 2048, 4096, 8192], "TG": [128, 1024, 2048, 4096, 8192], }, "Llama-3.1-70B": { "T3K": [128, 1024, 2048, 4096, 8192], "TG": [128, 1024, 2048, 4096, 8192], }, "Llama-3.3-70B": { "T3K": [128], "TG": [128, 1024, 2048, 4096, 8192], }, "Qwen3-Embedding-8B": { "N150": [128, 1024], "N300": [128, 1024, 2048, 4096, 8192], "T3K": [128, 1024, 2048, 4096, 8192], "TG": [128, 1024, 2048, 4096, 8192], "P150x4": [128, 1024, 2048, 4096, 8192], }, "Llama-3.2-3B": { "N150": [], }, } model_name = self.base_model_name device_name = self.device_name # If there is no entry for a model in model_specific_supported_seq_lens, use the entry in default_supported_seq_lens result = model_specific_supported_seq_lens.get(model_name, {}).get( device_name, default_supported_seq_lens.get(device_name) ) if result is not None: return cap_seq_lens_to_max_prefill_chunk_size(result, self.capped_warmup_seq_len) else: return [] @staticmethod def __get_llama_local_params_name(model_name): if "3.2-1B" in model_name: local_params = "LLAMA3_2_1B_PARAMS" elif "3.2-3B" in model_name: local_params = "LLAMA3_2_3B_PARAMS" elif "3.1-8B" in model_name or "Meta-Llama-3-8B" in model_name: local_params = "LLAMA3_1_8B_PARAMS" elif "3.2-11B" in model_name: local_params = "LLAMA3_2_11B_PARAMS" elif "3.1-70B" in model_name: local_params = "LLAMA3_1_70B_PARAMS" elif "3.2-90B" in model_name: local_params = "LLAMA3_2_90B_PARAMS" else: local_params = None return local_params def is_distributed_norm(self, mode: Mode): if not self.is_multichip: return False if all([dim > 1 for dim in list(self.mesh_device.shape)]): # 2D grid return True elif ( self.dim > 4096 and mode == Mode.PREFILL ): # Somewhere between 4k and 8k WH runs out of L1 if not distributed return True return False def ccl_topology(self): # Use ring on a T3K or 6U galaxy or P300x2 or P150x4/8 submesh if ttnn.cluster.get_cluster_type() in [ ttnn.cluster.ClusterType.P300_X2, ttnn.cluster.ClusterType.P150_X4, ttnn.cluster.ClusterType.P150_X8, ]: return ttnn.Topology.Ring elif ttnn.cluster.get_cluster_type() in [ ttnn.cluster.ClusterType.T3K, ttnn.cluster.ClusterType.GALAXY, ttnn.cluster.ClusterType.TG, ttnn.cluster.ClusterType.BLACKHOLE_GALAXY, ]: if self.num_devices >= 8: return ttnn.Topology.Ring else: # e.g., 1x4 submesh does not support ring topology; fallback to linear return ttnn.Topology.Linear if self.num_devices > 1: # All other multi chip devices return ttnn.Topology.Linear return None def prepare_residual_tensor_decode(self, x, input_mem_cfg, force_replicated=False, on_host=False): """ Prepare inputs for decode mode. x: (batch, seq, dim) """ dims = (None, None) if force_replicated else (None, -1) mesh_mapper = ttnn.ShardTensor2dMesh(self.mesh_device, dims=dims, mesh_shape=self.cluster_shape) if len(x.shape) == 3: batch = x.shape[0] seq_len = x.shape[1] assert x.shape[2] == self.dim elif len(x.shape) == 4: seq_len = x.shape[0] assert x.shape[1] == 1 batch = x.shape[2] assert x.shape[3] == self.dim assert seq_len == 1, "Only supporting decode mode" # Support input on device if torch.is_tensor(x): # Input on host -> Use torch x = x.transpose(0, 1).unsqueeze(1) # [seq_len, 1, batch, dim] # Pad small batches to 32 if batch < 32: zeros = torch.zeros(1, seq_len, 32, self.dim) zeros[:, :, :batch, :] = x x = zeros elif len(x.shape) == 3: # Input on device -> Use ttnn x = ttnn.reshape(x, (batch, seq_len, 1, self.dim)) # [batch, seqlen, dim] -> [batch, seqlen, 1, dim] x = ttnn.permute(x, (1, 2, 0, 3)) # [seq_len, 1, batch, dim] elif len(x.shape) == 4: pass # already in [seq_len, 1, batch, dim] if torch.is_tensor(x): x = ttnn.from_torch( x, device=self.mesh_device if not on_host else None, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, mesh_mapper=mesh_mapper, memory_config=input_mem_cfg if not on_host else None, ) else: # Convert the row major layout from embedding back to tile layout x = ttnn.to_layout(x, layout=ttnn.TILE_LAYOUT) return x def prepare_residual_tensor_prefill(self, x_bsh, force_replicated=False): """ Prepare inputs for prefill mode. x: (batch, seq, hidden_dim) B: batch (1) S: sequence len H: dim """ x_1BSH = x_bsh.unsqueeze(0) dims = (None, None) if force_replicated else (None, -1) mesh_mapper = ttnn.ShardTensor2dMesh(self.mesh_device, dims=dims, mesh_shape=self.cluster_shape) # input goes to DRAM xs_1BSH = ttnn.from_torch( x_1BSH, device=self.mesh_device, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, memory_config=ttnn.DRAM_MEMORY_CONFIG, mesh_mapper=mesh_mapper, ) return xs_1BSH def _get_text_prefix(self): if self.is_llama_vision(): return "text_model." else: return "" def _get_vision_prefix(self): return "visual." def _get_hidden_activation_type(self, config): activation_map = { "gelu": ttnn.UnaryWithParam(ttnn.UnaryOpType.GELU, 0.0), "gelu_pytorch_tanh": ttnn.UnaryWithParam(ttnn.UnaryOpType.GELU, 1.0), "relu": ttnn.UnaryOpType.RELU, "silu": ttnn.UnaryOpType.SILU, "swish": ttnn.UnaryOpType.SILU, } hidden_activation = config.get("hidden_act") or config.get("hidden_activation") if not hidden_activation: # Default to SILU if no activation is specified return ttnn.UnaryOpType.SILU hidden_activation = hidden_activation.lower() if hidden_activation not in activation_map: raise NotImplementedError(f"Unsupported activation '{hidden_activation}'") return activation_map.get(hidden_activation, ttnn.UnaryOpType.SILU) def _set_model_specific_params(self): # Gemma-family text models (gemma, gemma2) store RMSNorm weights as (1 + w) # and scale input embeddings by sqrt(hidden_size). Gemma-3 keeps this in its # own ModelArgs subclass; enabling it here (config-gated) brings up Gemma-2 # text models via HF_MODEL without touching any non-Gemma model. if self.model_type is not None and str(self.model_type).lower() in ("gemma", "gemma2"): self.rms_norm_add_unit_offset = True self.embed_scale = self.dim**0.5 # head_dim=256 decode-attention fix: pin a small KV chunk and use the # flash-decode op's default compute config. Matches the known-good Gemma-3 # decode SDPA (models/demos/gemma4/tt/attention/decode.py); the framework's # auto chunk + forced fp32-accum compute config corrupt the cross-chunk # online-softmax reduction for head_dim=256. self.sdpa_decode_q_chunk_size = 32 self.sdpa_decode_k_chunk_size = 64 self.sdpa_decode_use_default_compute_config = True # EXAONE-4.x (incl. the EXAONE-4.5 text decoder, whose text model_type is # exaone4): post-norm residual order with no input_layernorm — # h = x + post_attention_layernorm(attn(x)); out = h + post_feedforward_layernorm(mlp(h)) # — and hybrid rope inverted vs Gemma-3: RoPE (llama3-scaled, theta 1e6) is # applied ONLY on sliding_attention layers while full_attention layers are # NoPE (no positional encoding). See HF modeling_exaone4.py. if self.model_type is not None and str(self.model_type).lower() in ("exaone4", "exaone4_5", "exaone4_5_text"): self.use_post_norm = True # Sliding (local) layers take the checkpoint's rope scaling; global # (full-attention) layers take identity cos/sin via use_global_nope. self.rope_scaling_local = self.rope_scaling self.rope_scaling = None self.use_global_nope = True # EXAONE-4.5 checkpoints ship a vision tower (model.visual.*) and an MTP # module (mtp.*); this port runs the text decoder only. transformers # itself drops mtp.* on load (_keys_to_ignore_on_load_unexpected). self.force_text_only = True self.is_multimodal = False def _set_params_from_dict(self, config): eos_token_id = config.get("eos_token_id", None) self.image_token_index = config.get("image_token_index", None) # Try to get text_config, if it doesn't exist everything is text config text_config = config.get("text_config", config) self.eos_token_id = None if isinstance(eos_token_id, int) else eos_token_id layer_types = text_config["layer_types"] if "layer_types" in text_config else None # Architecture family (e.g. "gemma2", "gemma3", "llama"). Used to enable # architecture-specific behaviour without affecting other models. self.model_type = text_config.get("model_type", config.get("model_type", None)) # Common params with different names between Meta and HF self.dim = text_config.get("dim", text_config.get("hidden_size")) self.n_heads = text_config.get("n_heads", text_config.get("num_attention_heads")) self.n_kv_heads = text_config.get("n_kv_heads", text_config.get("num_key_value_heads")) self.n_layers = text_config.get("n_layers", text_config.get("num_hidden_layers")) # multimodal llama additionally adds cross attention layers # they are calculated in HF but not calculated in Meta self.n_layers -= len(text_config.get("cross_attention_layers", ())) self.vision_num_cross_attention_layers = len(text_config.get("cross_attention_layers", ())) self.sliding_window_pattern = ( [lt == "sliding_attention" for lt in layer_types] if layer_types is not None else [False] * self.n_layers ) self.full_model_n_layers = self.n_layers self.norm_eps = text_config.get( "norm_eps", text_config.get("rms_norm_eps", text_config.get("layer_norm_eps")) ) # layer_norm_eps: Command-R (cohere) HF key self.vocab_size = text_config["vocab_size"] # Pad vocab_size to be divisible by (32 * num_devices) for proper shard alignment tile_size = 32 if self.is_galaxy: self.padded_vocab_size = compute_galaxy_padded_vocab_size(self.vocab_size, self.num_devices) elif self.num_devices == 0: # No mesh (e.g. reference-output generation): pad to tile_size only self.padded_vocab_size = math.ceil(self.vocab_size / tile_size) * tile_size else: self.padded_vocab_size = compute_padded_vocab_size(self.vocab_size, self.num_devices) self.head_dim = text_config.get("head_dim", self.dim // self.n_heads) or self.dim // self.n_heads self.num_experts_per_tok = text_config.get("num_experts_per_tok", 0) self.num_local_experts = text_config.get("num_local_experts", 0) self.max_context_len = text_config.get("max_position_embeddings") # Handle different MLP dimension specifications if "intermediate_size" in text_config: self.hidden_dim = text_config["intermediate_size"] self.ffn_dim_multiplier = None self.multiple_of = None # temporary solution for using HF_MODEL for LLaMA until llama_model references are removed local_params = self.__get_llama_local_params_name(self.model_name) if local_params in self.LOCAL_LLAMA_PARAMS: params_file = os.path.join(self.LOCAL_LLAMA_PARAMS[local_params], "params.json") if os.path.exists(params_file): with open(params_file, "r") as f: params = json.load(f) self.ffn_dim_multiplier = params["ffn_dim_multiplier"] self.multiple_of = params["multiple_of"] else: self.ffn_dim_multiplier = text_config["ffn_dim_multiplier"] self.multiple_of = text_config["multiple_of"] self.hidden_dim = calculate_hidden_dim(self.dim, self.ffn_dim_multiplier, self.multiple_of) if "_name_or_path" in config and config["_name_or_path"]: normalized_path = os.path.normpath(config["_name_or_path"]) # For HF paths, they might end with `/snapshots//` if "snapshots" in normalized_path: full_model_name = normalized_path.split(os.path.sep)[-3] self.model_name = full_model_name.split("--")[-1] else: self.model_name = os.path.basename(normalized_path) logger.info(f"Model name from config: {self.model_name}") if self.base_model_name in ["Qwen2.5-7B", "Qwen2.5-VL-7B"] and self.num_devices not in [0, 2, 4]: raise AssertionError( "Qwen2.5-7B and Qwen2.5-VL-7B is only supported on 2 or 4 devices, run on an N300 or use MESH_DEVICE=N150x4" ) if self.num_devices > 0 and self.cluster_shape != [1, 1]: self.pad_logits_to_power_of_2 = should_pad_sampling_logits_to_power_of_2( self.base_model_name, self.padded_vocab_size, self.num_devices ) else: # Off on [1, 1]: an A/B on the multi-step split path (PR #53167) # measured no end-to-end decode benefit from padding the topk chunks # to a power of two, so the flag stays multi-device only. self.pad_logits_to_power_of_2 = False self.unpadded_hidden_dim = self.hidden_dim # Don't need to pad for CPU runs if self.num_devices: # Default padding cores for each model, 0 if not set here default_padded_cores = { "Qwen2.5-VL-72B": 32, "Qwen2.5-VL-32B": 16, "Qwen2.5-72B": 32, "Qwen2.5-32B": 16, "Qwen2.5-7B": 16, "QwQ-32B": 16, # 27392/8 = 3424 = 107 tiles (prime) -> find_grid_k_n degenerates to a # single core (1x1 norm grid overflows L1; 1-core MLP matmuls). Pad to # 28672 (112 tiles/device, gcd(160,112)=16 -> 16-core grids). "EXAONE-4.5-33B": 16, }.get(self.base_model_name, 0) # Override MLP padding cores from env var mlp_padded_cores = int(os.environ.get("PAD_MLP_CORES", default_padded_cores)) # Only pad if MLP_PADDED_CORES is non-zero if mlp_padded_cores > 0: padded_hidden_dim = nearest_multiple( self.hidden_dim, mlp_padded_cores * ttnn.TILE_SIZE * self.num_devices ) if padded_hidden_dim != self.hidden_dim: logger.info( f"PAD_MLP_CORES={mlp_padded_cores}, padding hidden dim from {self.hidden_dim} to {padded_hidden_dim}" ) self.hidden_dim = padded_hidden_dim self.layer_types = text_config.get("layer_types", None) # Sliding window attention self.sliding_window = text_config.get("sliding_window", None) # Gemma-2 alternates local (sliding-window) and global attention. Upstream # Gemma2Config.__post_init__ already fills `layer_types` before to_dict(), so # for a real HF checkpoint this block is usually a no-op; keep it as a # defensive fallback for dict-style configs that omit the list. # Pattern: even layers = sliding (matches HF is_sliding = not bool(layer_idx % 2)). if ( self.model_type is not None and str(self.model_type).lower().startswith("gemma2") and self.sliding_window is not None and self.layer_types is None ): self.layer_types = ["sliding_attention" if (i % 2 == 0) else "full_attention" for i in range(self.n_layers)] self.sliding_window_pattern = [lt == "sliding_attention" for lt in self.layer_types] # RoPE params (transformers 5.x nests these under `rope_parameters`) self.rope_theta = get_rope_theta(text_config) self.rope_theta_local = get_rope_local_base_freq(text_config) self.use_sliding_window = text_config.get("use_sliding_window", None) if ( self.sliding_window is not None and self.rope_theta_local is None and (self.use_sliding_window == True or self.use_sliding_window is None) ): # For interleaved attention self.rope_theta_local = self.rope_theta rope_scaling_params = get_rope_scaling(text_config) self.original_max_context_len = text_config.get("original_max_position_embeddings", None) self.rope_scaling = ( rope_scaling_model_factory(rope_scaling_params, original_max_context_len=self.original_max_context_len) if rope_scaling_params else None ) self.query_pre_attn_scalar = text_config.get("query_pre_attn_scalar", None) # Command-R (cohere): final-logit scalar applied post-linear on the LM head. self.logit_scale = text_config.get("logit_scale", None) # Final logit soft-capping (Gemma-2): logits -> tanh(logits / cap) * cap. # Attn-score softcapping is not applied (see __init__ comment); only the # final-logit cap is consumed, via Transformer._apply_final_logit_softcapping. self.final_logit_softcapping = text_config.get("final_logit_softcapping", None) # Configurable MLP activation type self.mlp_activation_type = self._get_hidden_activation_type(text_config) self._set_vision_params(config) self.is_multimodal = not self.force_text_only and ("vision_config" in config or self.is_vision()) self.vision_chunk_size = config.get("vision_chunk_size", 896) self.vision_max_num_chunks = config.get("vision_max_num_chunks", 4) if "vision_num_cross_attention_layers" in config: self.vision_num_cross_attention_layers = config["vision_num_cross_attention_layers"] self.vision_dim = 1280 self.vision_mlp_ratio = 4 self.vision_hidden_dim = int(self.vision_dim * self.vision_mlp_ratio) self.vision_act_layer = ttnn.UnaryOpType.GELU self.vision_dropout = 0.0 self.vision_attn_n_heads = 16 self.vision_head_dim = self.vision_dim // self.vision_attn_n_heads self.vision_n_layers = 32 self.vision_n_global_layers = 8 self.vision_max_num_tiles = 4 self.vision_patch_size = 14 self.vision_in_channels = 3 self.state_dict_text_prefix = self._get_text_prefix() self.state_dict_vision_prefix = self._get_vision_prefix() self._set_model_specific_params() @property def use_scaled_rope(self): return self.rope_scaling is not None @property def base_model_name(self): return get_base_model_name(self.model_name) @property def vision_chunk_ntok(self): """ Returns the number of tokens per chunk, accounting for the extra class token """ return (self.vision_chunk_size // self.vision_patch_size) ** 2 + 1 def _set_vision_params(self, config): vision_config = config.get("vision_config", config) # Get vision_dim from config (same key for all models) self.vision_dim = vision_config.get("hidden_size", 1280) # Get vision_head_dim - Mistral has it in config, others calculate it if "head_dim" in vision_config: self.vision_head_dim = vision_config["head_dim"] else: num_heads = vision_config.get("num_attention_heads") or vision_config.get("num_heads") or 16 self.vision_head_dim = self.vision_dim // num_heads # Get image_size from config (same key for all models) self.image_size = vision_config.get("image_size", -1) # Optional values with reasonable fallbacks chunk_size_fallback = self.image_size if self.image_size != -1 else vision_config.get("image_size", -1) self.vision_chunk_size = vision_config.get("vision_chunk_size", chunk_size_fallback) self.vision_max_num_chunks = vision_config.get("vision_max_num_chunks", vision_config.get("max_num_tiles", 4)) # Common vision parameters for all models intermediate_size = vision_config.get("intermediate_size", self.vision_dim * 4) self.vision_image_size = vision_config.get("image_size", 1540) self.vision_rope_theta = get_rope_theta(vision_config, default=10000.0) self.image_token_index = vision_config.get("image_token_index", 10) self.vision_mlp_ratio = intermediate_size // self.vision_dim self.vision_hidden_dim = int(self.vision_dim * self.vision_mlp_ratio) self.vision_attn_n_heads = vision_config.get("num_attention_heads") or vision_config.get("num_heads") or 16 # Note: vision_head_dim is already set above (from config for Mistral, calculated for others) # Default to 32 layers - this is the standard for Llama vision models (e.g., Llama-3.2-11B-Vision uses 32) # This default is only used when the config doesn't specify num_hidden_layers or depth # Models that specify these values in their config (e.g., Mistral-Small-3.1-24B-Instruct-2503 uses 24) # will use their specified values, not this default # The default of 32 comes from the main branch and matches Llama vision model architecture self.vision_n_layers = vision_config.get("num_hidden_layers") or vision_config.get("depth") or 32 self.vision_patch_size = vision_config.get("patch_size", 14) self.vision_in_channels = vision_config.get("num_channels", 3) self.vision_dropout = vision_config.get("attention_dropout", 0.0) self.mm_tokens_per_image = vision_config.get("mm_tokens_per_image", config.get("mm_tokens_per_image", 256)) # Optional vision activation layer, defaults to GELU act_layer = vision_config.get("act_layer", "gelu").lower() self.vision_act_layer = { "gelu": ttnn.UnaryOpType.GELU, "relu": ttnn.UnaryOpType.RELU, "silu": ttnn.UnaryOpType.SILU, }.get(act_layer, ttnn.UnaryOpType.GELU) # Optional tuning knobs self.vision_max_num_tiles = vision_config.get("max_num_tiles", 4) self.vision_n_global_layers = vision_config.get("n_global_layers", vision_config.get("num_global_layers", 8)) def _set_hf_params(self, checkpoint_dir): def merge_text_config(base_config): text_config = base_config.get("text_config", {}) # Merge non-nested keys into text_config text_config.update({k: v for k, v in base_config.items() if k not in ["text_config", "vision_config"]}) return text_config def merge_vision_config(base_config): vision_config = base_config.get("vision_config", {}) # Merge non-nested keys into vision_config vision_config.update({k: v for k, v in base_config.items() if k not in ["text_config", "vision_config"]}) return vision_config from transformers import AutoConfig if self.dummy_weights: logger.info(f"Loading state param for dummy {self.model_name} from {self.LOCAL_HF_PARAMS[self.model_name]}") self.hf_config = AutoConfig.from_pretrained( self.LOCAL_HF_PARAMS[self.model_name], trust_remote_code=self.trust_remote_code_hf ) else: self.hf_config = AutoConfig.from_pretrained( self.CKPT_DIR, trust_remote_code=self.trust_remote_code_hf, local_files_only=os.getenv("CI") == "true", ) config = self.hf_config.to_dict() if "text_config" in config or "vision_config" in config: merged_text_config = merge_text_config(config) self._set_params_from_dict(merged_text_config) if "Mistral-Small-3.1-24B-Instruct-2503" in self.model_name: self._set_vision_params(config["vision_config"]) else: if "vision_config" in config: merged_vision_config = merge_vision_config(config) self._set_vision_params({"vision_config": merged_vision_config}) self.is_multimodal = not self.force_text_only and ("vision_config" in config or self.is_vision()) else: self._set_params_from_dict(config) # compatibility with _set_params if "llama" in self.model_name.lower(): if "3.2-11B" in checkpoint_dir: logger.warning(f"-Vision is removed from model_name {self.model_name}") # TODO: do not remove "-Vision" part self.model_name = "Llama-3.2-11B" + ("-Instruct" if self.instruct else "") elif "3.1-70B" in checkpoint_dir: self.is_70b = True # self.dim == 8192 and self.n_layers == 80 elif "3.2-90B" in checkpoint_dir: logger.warning(f"-Vision is removed from model_name {self.model_name}") # TODO: do not remove "-Vision" part self.model_name = "Llama-3.2-90B" + ("-Instruct" if self.instruct else "") self.is_90b = True def __repr__(self): return f"""ModelArgs( dim={self.dim}, n_layers={self.n_layers}, n_heads={self.n_heads}, n_kv_heads={self.n_kv_heads}, vocab_size={self.vocab_size}, multiple_of={self.multiple_of}, ffn_dim_multiplier={self.ffn_dim_multiplier}, norm_eps={self.norm_eps}, rope_theta={self.rope_theta}, rope_scaling_factor={self.rope_scaling.factor if self.rope_scaling is not None else None}, max_batch_size={self.max_batch_size}, max_seq_len={self.max_seq_len}, vision_chunk_size={self.vision_chunk_size}, vision_max_num_chunks={self.vision_max_num_chunks}, vision_num_cross_attention_layers={self.vision_num_cross_attention_layers} )""" def can_enable_trace(self, prefill_seq_len, num_cached_tokens=0): """ This function is used to determine if trace should be enabled for the prefill. Tracing is used only for certain sequence lengths, because for bigger sequence lengths, op2op gaps are already small, so we don't need tracing. # TODO: Support chunked prefill with tracing - https://github.com/tenstorrent/tt-metal/issues/32056 """ allowed_seq_lens = self.trace_prefill_supported_seq_lens return ( prefill_seq_len in allowed_seq_lens and prefill_seq_len <= self.max_prefill_chunk_size and prefill_seq_len <= self.max_seq_len ) def is_llama_vision(self): return self.CKPT_DIR is not None and ("llama" in self.CKPT_DIR.lower()) and ("vision" in self.CKPT_DIR.lower()) def is_vision(self): """Check if this is a vision-capable model (Llama vision or Mistral multimodal)""" return self.is_llama_vision() or ( "mistral" in self.model_name.lower() and ( (self.CKPT_DIR is not None and "vision" in self.CKPT_DIR.lower()) or "Mistral-Small-3.1-24B-Instruct-2503" in self.model_name ) ) def get_state_dict_prefix(self, module_name, layer_num, is_vision=False): # Llama vision models use "text_model." prefix for text keys # Other vision models (Mistral, etc.) don't use text_model prefix for text if self.is_llama_vision(): text_prefix = self.state_dict_text_prefix else: # Standard models and non-Llama vision: no prefix for text, prefix for vision text_prefix = "" if not is_vision else self.state_dict_text_prefix vision_prefix = self.state_dict_vision_prefix if is_vision else "" layer_prefix = f"layers.{layer_num}." if layer_num is not None else "" text_module_map = { "MLP": "feed_forward", "Attention": "attention", "TransformerBlock": "", "": "", # If no module is given, just get layer prefix } vision_module_map = { "MLP": "mlp.", "Attention": "self_attn.", "TransformerBlock": "", "": "", } module_map = vision_module_map if is_vision else text_module_map prefix = vision_prefix if is_vision else text_prefix return prefix + layer_prefix + module_map[module_name] def weight_cache_path(self, dtype): # Keep the weight cache separate for generative and instruct weights if self.instruct: return ( self.model_cache_path / {ttnn.bfloat16: "tensor_cache_instruct_bf16", ttnn.bfloat8_b: "tensor_cache_instruct_bfp8"}[dtype] ) else: return ( self.model_cache_path / {ttnn.bfloat16: "tensor_cache_bf16", ttnn.bfloat8_b: "tensor_cache_bfp8"}[dtype] ) # Name of the marker file dropped into a weight-cache directory once every weight for that # (model, dtype, mesh shape) has been materialized to disk. Generalizes the GPT-OSS # warm-cache detector (#48531) to every tt_transformers e2e model. # Marker filename and schema version both come from models/common/weight_cache.py -- this # class must not define its own, or the two writers diverge (they previously shared a filename # and version number while encoding mesh_shape incompatibly). See that module for the version # history and what each field guarantees. WEIGHT_CACHE_MARKER = _WC_MARKER WEIGHT_CACHE_FORMAT_VERSION = _WC_FORMAT_VERSION def _weight_cache_build_variant(self): """A signature of the build options that change an ``as_tensor`` cache *filename*. The cache name is ``{name}_dtype_{dtype}_layout_{layout}.tensorbin``, so anything that moves a weight to a different dtype -- or adds weights outright -- produces a different file set. Two knobs do that here: the DRAM prefetcher (pins every layer to decoder 0's dtype and adds the ring-matmul splits in lm_head) and the precision/optimizations config (per-decoder, per-tensor-group dtypes). Recording them means a cache seeded under one variant is not accepted for a build that needs a different one -- which matters because a filename this build needs but the seed never wrote would otherwise be regenerated by as_tensor FROM the placeholder. (#45400 review)""" try: # get_tensor_dtype lives on the DecodersPrecision held in self.optimizations, NOT on # ModelArgs. The original self.get_tensor_dtype(...) raised AttributeError on every # model -- unnoticed because the old except collapsed it to the match-anything # "unknown", and no hardware run executed this method until the 0ec5959bade sweep, # where the fail-closed sentinel surfaced it on the first cold build. (#45400 review) dtypes = [ str(self.optimizations.get_tensor_dtype(decoder_id, group, prefetcher=bool(self.prefetcher))) for decoder_id in range(self.n_layers) for group in TensorGroup ] precision = hashlib.sha1("|".join(dtypes).encode()).hexdigest()[:12] except Exception as e: # A variant we cannot compute is a variant we cannot verify. Do NOT collapse to a # match-anything constant (that silently reopened the placeholder-persistence hole for # every precision variant); return an unverifiable sentinel that the completeness gate # rejects and mark_weight_cache_complete refuses to write, so the build cold-loads -- # slow but correct -- and the log says why. (#45400 review, finding R3) logger.warning(f"Could not compute the weight-cache build variant ({e!r}); warm-cache skip disabled.") return {"unverifiable": True, "error": f"{type(e).__name__}: {e}"} return { "prefetcher": bool(self.prefetcher), "precision": precision, # attention.py caches wqkv_bias_decode_sharded_{batch_size}, so batch is in a filename. "batch": int(getattr(self, "max_batch_size", 0) or 0), # attention.py picks cache_name("wo_width_sharded_2d") vs cache_name("wo") off this. "fused_ag": bool(getattr(self, "use_fused_all_gather_matmul", False)), # load_state_dict permutes QKV differently per rope mode, so the SAME cache filename # carries different content across modes. The Llama CI job runs both modes against one # cache dir; keeping the mode in the variant stops a marker seeded under one mode from # certifying the other. (#45400 review, finding R2) "hf_rope": bool(getattr(self, "use_hf_rope", False)), } def _weight_cache_identity(self, components=None): """Marker identity for this build. `components` names the parts being constructed, so a text-only seed cannot certify a cache for a build that also needs the vision tower.""" return dict( model_name=self.model_name, n_layers=self.n_layers, mesh_shape=tuple(self.mesh_device.shape), components=components, build_variant=self._weight_cache_build_variant(), ) def weight_cache_is_complete(self, dtype, components=None): """True when the on-disk ttnn weight cache for this (model, dtype, mesh shape, components) was fully built by a previous run and every tensorbin that build produced is still present. When True, ttnn.as_tensor loads every weight from its cached .tensorbin and the HF state_dict is never read, so the caller can skip the expensive from_pretrained host load entirely (the load that OOMs/hangs during prefill, #48509). Set TT_TRANSFORMERS_FORCE_MODEL_LOAD=1 to force a fresh load (e.g. to regenerate the cache). Delegates to models/common/weight_cache.py so there is a SINGLE marker reader/writer: the two used to encode mesh_shape differently while sharing one filename and format_version, so a model reachable from both (gemma3 inherits this class but its demos call the shared helper) had each side reject the other's marker and cold-load forever. (#45400 review)""" return _weight_cache_is_complete(self.weight_cache_path(dtype), **self._weight_cache_identity(components)) def mark_weight_cache_complete(self, dtype, state_dict=None, components=None): """Record that the ttnn weight cache for this (model, dtype, mesh shape, components) was fully built, so subsequent runs can skip the HF state_dict load. state_dict (the real, just-loaded weights) is captured as a {key: [shape, dtype]} manifest so a later warm run can reconstruct a dataless placeholder without touching HF. Must be called AFTER the model is constructed, so the tensorbins exist to be recorded.""" if state_dict is None: return _mark_weight_cache_complete( self.weight_cache_path(dtype), state_dict, is_moe=bool(getattr(self, "is_mixture_of_experts", False)), **self._weight_cache_identity(components), ) def placeholder_state_dict(self, dtype): """Warm-cache build: return a lazy, dataless stand-in for the HF state_dict. Every weight is served as an uninitialized CPU torch.empty of the shape/dtype recorded in the marker manifest -- no from_pretrained, no weight bytes read from disk, so the prefill host-OOM (#48509) is avoided. ttnn.as_tensor(cache_file_name=...) ignores this placeholder on a cache hit (ttnn/operations/core.py) and loads the real weight from its .tensorbin; the placeholder exists only to satisfy the host-side reshape ops (permute/chunk/cat/transpose) the modules run before calling as_tensor. Guarded by weight_cache_is_complete, so every as_tensor is guaranteed to hit and discard it. The mapping is falsy (__bool__ -> False) so reference-building callers that test `if not state_dict` (e.g. test_model_prefill) load real weights explicitly, while modules that index it during construction still work.""" import collections.abc marker = _wc_marker_path(self.weight_cache_path(dtype), self._weight_cache_build_variant()) meta = json.loads(marker.read_text()) manifest = meta["weights"] self.is_mixture_of_experts = bool(meta.get("is_moe", False)) # Same reasoning as build_cached_state_dict: these are set by load_state_dict from the # checkpoint keys, which the warm path never reads. The manifest has the same keys. self.fuse_qkv = any("qkv" in k for k in manifest) self.fuse_mlp = any("gate_up" in k for k in manifest) if self.is_mixture_of_experts: self.moe = True expert_indices = [int(k[-11]) + 1 for k in manifest if "block_sparse_moe.experts" in k] self.num_experts = max(expert_indices) if expert_indices else self.num_local_experts dtype_cache = {} def _to_dtype(s): if s not in dtype_cache: dtype_cache[s] = getattr(torch, s.rsplit(".", 1)[-1]) return dtype_cache[s] class _PlaceholderStateDict(collections.abc.Mapping): is_placeholder = True def __init__(self, spec): self._spec = spec # key -> (shape, dtype_str) def __bool__(self): return False def __getitem__(self, key): shape, dt = self._spec[key] return torch.empty(tuple(shape), dtype=_to_dtype(dt)) def __iter__(self): return iter(self._spec) def __len__(self): return len(self._spec) logger.info(f"Warm ttnn weight cache: built placeholder state_dict for {len(manifest)} weights (no HF load).") return _PlaceholderStateDict(manifest) def get_model_config(self): return self.model_config def get_hf_model_cls(self): from transformers import AutoModelForCausalLM, AutoModelForImageTextToText if not self.is_multimodal: # Text-only ports of multimodal checkpoints (e.g. EXAONE-4.5): the # composite config class only maps to an image-text-to-text model — # load that; text-only inputs bypass its vision tower entirely. if ( type(self.hf_config) not in AutoModelForCausalLM._model_mapping and type(self.hf_config) in AutoModelForImageTextToText._model_mapping ): return AutoModelForImageTextToText return AutoModelForCausalLM # AutoModelForVision2Seq was removed in transformers 5.x; its model mapping # was folded into AutoModelForImageTextToText (available since 4.46). for model_cls in (AutoModelForImageTextToText,): if type(self.hf_config) in model_cls._model_mapping: return model_cls raise ValueError(f"Unknown model for config {type(self.hf_config)}") def _load_hf_state_dict_for_configured_layers(self): """Read only the decoder layers this ModelArgs was trimmed to (plus embeddings, norm and lm_head) straight from the safetensors shards, instead of materialising the whole checkpoint through from_pretrained. Returns None when the full HF model is needed or the checkpoint cannot be read that way, and the caller falls back to from_pretrained. """ full_n_layers = getattr(self, "full_model_n_layers", self.n_layers) if self.n_layers >= full_n_layers: return None # The HF model object itself is needed for reference modules, custom (remote-code) layouts and # the multimodal key remapping; those keep the from_pretrained path. if self.cache_hf_flag or self.trust_remote_code_hf or self.is_multimodal or self.force_text_only: return None try: state_dict = load_hf_state_dict_for_layers( self.CKPT_DIR, self.n_layers, local_files_only=os.getenv("CI") == "true" ) except (OSError, KeyError, ValueError, ImportError) as exc: logger.warning(f"Per-layer safetensors load of {self.CKPT_DIR} failed ({exc}); loading the full model") return None if not any(_HF_LAYER_KEY_RE.match(k) for k in state_dict): logger.warning( f"Per-layer safetensors load of {self.CKPT_DIR} found no decoder layers; loading the full model" ) return None # from_pretrained materialises the tied lm_head; raw shards only carry the embedding. if "lm_head.weight" not in state_dict and getattr(self.hf_config, "tie_word_embeddings", False): if "model.embed_tokens.weight" in state_dict: state_dict["lm_head.weight"] = state_dict["model.embed_tokens.weight"] logger.info(f"Loaded {len(state_dict)} weights for {self.n_layers} of {full_n_layers} layers from safetensors") return state_dict def load_state_dict(self): # by default, the model is not a mixture-of-expert. This will be set to True if we find any `.experts.` in the keys if self.dummy_weights: from transformers import AutoConfig config = AutoConfig.from_pretrained( self.LOCAL_HF_PARAMS[self.model_name], trust_remote_code=self.trust_remote_code_hf ) if hasattr(config, "text_config"): config.text_config.num_layers = self.n_layers config.text_config.num_hidden_layers = self.n_layers else: config.num_layers = self.n_layers config.num_hidden_layers = self.n_layers model_cls = self.get_hf_model_cls() from_config_exc = None try: # Avoid loading checkpoint weights when dummy_weights is set. try: model = model_cls.from_config(config, trust_remote_code=self.trust_remote_code_hf) except TypeError: model = model_cls.from_config(config) except Exception as exc: from_config_exc = exc logger.info("Error loading dummy weights using .from_config. Error: {}", exc) if hasattr(model_cls, "_from_config"): try: try: model = model_cls._from_config(config, trust_remote_code=self.trust_remote_code_hf) except TypeError: model = model_cls._from_config(config) except Exception as fallback_exc: logger.info("Error loading dummy weights using ._from_config. Error: {}", fallback_exc) if from_config_exc is not None: raise fallback_exc from from_config_exc raise else: raise # model.load_state_dict({k: torch.randn_like(v) for k, v in model.state_dict().items()}) state_dict = model.state_dict() else: # Always HuggingFace since we only support HF_MODEL now state_dict = self._load_hf_state_dict_for_configured_layers() if state_dict is None: model_cls = self.get_hf_model_cls() model = model_cls.from_pretrained( self.CKPT_DIR, torch_dtype="auto", trust_remote_code=self.trust_remote_code_hf, local_files_only=os.getenv("CI") == "true", # Note that the default setting is torch.dtype.float32, but model weights are # may come in any dtype. If the model's weights are in torch.dtype.bfloat16, this would result in 2x memory usage from an # unnecessary cast. ) if self.cache_hf_flag: self.cached_hf_model = model state_dict = model.state_dict() self.is_mixture_of_experts = any([".experts." in k for k in state_dict.keys()]) if self.is_multimodal: state_dict = standardize_hf_keys_multimodal(state_dict) if self.is_llama_vision(): if self.use_hf_rope: # For HF-style RoPE: skip QKV format conversion state_dict = convert_hf_to_meta_mllama_no_qkv_permute(state_dict, self.head_dim, self.hf_config) else: # Standard: convert to Meta format state_dict = convert_hf_to_meta_mllama(state_dict, self.head_dim, self.hf_config) else: if self.use_hf_rope: # For HF-style RoPE: skip QKV format conversion state_dict = convert_vision_hf_to_meta_no_qkv_permute(state_dict, self.head_dim) else: # Standard: convert to Meta format state_dict = convert_vision_hf_to_meta(state_dict, self.head_dim) else: if self.force_text_only: # Text-only port of a multimodal checkpoint (e.g. EXAONE-4.5): drop # the vision tower and auxiliary heads (MTP) before key sniffing — # vision `attn.qkv` keys would otherwise trip fuse_qkv, and mtp.layers.* # keys would collide with the layer-trim loop below. state_dict = { k: v for k, v in state_dict.items() if not (k.startswith(("model.visual.", "visual.", "mtp."))) } self.fuse_qkv = any(["qkv" in layer_name for layer_name in state_dict.keys()]) self.fuse_mlp = any(["gate_up" in layer_name for layer_name in state_dict.keys()]) state_dict = standardize_hf_keys(state_dict) if self.use_hf_rope: # For Attention: skip QKV format conversion state_dict = convert_hf_to_meta_no_qkv_permute(state_dict, self.head_dim, self.n_heads, self.n_kv_heads) elif self.model_type == "cohere": # Command-R rotates Q/K INTERLEAVED-native (HF modeling_cohere overrides # rotate_half: adjacent pairs (2i,2i+1) + repeat_interleave cache) — unlike # llama's NeoX half-split. The stock NeoX->Meta reverse_permute therefore # SCRAMBLES already-interleaved cohere Q/K pairs; the ttnn interleaved # rotary op is correct only with the unpermuted layout. Root-caused # 2026-08-28 (quality defect): layer-0 PCC 0.9324 -> 0.9998 at seq 36 # (served math probe 422 restored) by skipping the permute. state_dict = convert_hf_to_meta_no_qkv_permute(state_dict, self.head_dim, self.n_heads, self.n_kv_heads) else: # Standard: convert to Meta format state_dict = convert_hf_to_meta(state_dict, self.head_dim, self.n_heads, self.n_kv_heads) keys_dict = list(state_dict.keys())[:] remv = [f"layers.{i}." for i in list(range(self.n_layers, self.full_model_n_layers))] for k in keys_dict: if any([r in k for r in remv]): state_dict.pop(k) if getattr(self, "is_mixture_of_experts", False): self.moe = True # transformers 5.x fused Mixtral experts into batched params (mlp.experts.*), # dropping the per-expert block_sparse_moe.experts.{i} keys. Derive the count # from those keys when present (<5.x), else fall back to the config value. expert_indices = [int(item[-11]) + 1 for item in keys_dict if "block_sparse_moe.experts" in item] self.num_experts = max(expert_indices) if expert_indices else self.num_local_experts return state_dict # ========================================================================= # MATMUL / CONFIG HELPERS # ========================================================================= def create_dram_sharded_mem_config(self, k, n, dram_grid=None): """Create DRAM-sharded memory config for width-sharded tensors""" dram_cores = self.dram_grid_size.x # WH has 12 dram cores, P150 has 8, P100 has 7 assert self.dram_grid_size.y == 1, "Current dram sharding assumes y dim is 1" padded_size = math.ceil(n / (ttnn.TILE_SIZE * dram_cores)) * (ttnn.TILE_SIZE * dram_cores) if dram_grid is None: dram_grid = self.dram_weight_grid shard_spec = ttnn.ShardSpec(dram_grid, (k, padded_size // dram_cores), ttnn.ShardOrientation.ROW_MAJOR) return ttnn.MemoryConfig(ttnn.TensorMemoryLayout.WIDTH_SHARDED, ttnn.BufferType.DRAM, shard_spec) def matmul_config( self, m: int, k: int, n: int, grid_size: Tuple[int, int], in0_block_w: int = None, fuse_batch: bool = False, fused_activation=None, per_core_M=None, per_core_N=None, ): if per_core_M is None: per_core_M = math.ceil(m / (ttnn.TILE_SIZE * grid_size[1])) if per_core_N is None: per_core_N = math.ceil(n / (ttnn.TILE_SIZE * grid_size[0])) out_subblock_h = 1 out_subblock_w = ( get_out_subblock_w(per_core_N, out_subblock_h) if not self.is_galaxy else 1 ) # TODO: Needed for TG hang workaround if in0_block_w is None: assert ( k % (ttnn.TILE_SIZE * grid_size[1]) == 0 ), f"Input width must be divisible by tile size times grid size" in0_block_w = self.find_largest_divisor(k // (ttnn.TILE_SIZE * grid_size[1])) return ttnn.MatmulMultiCoreReuseMultiCastProgramConfig( compute_with_storage_grid_size=grid_size, in0_block_w=in0_block_w, out_subblock_h=out_subblock_h, out_subblock_w=out_subblock_w, per_core_M=per_core_M, per_core_N=per_core_N, transpose_mcast=False, fused_activation=fused_activation, fuse_batch=fuse_batch, ) def dram_decode_in0_block_w(self, k: int, n: int, num_cores: int) -> int: """in0_block_w for a DRAM-sharded decode matmul with one reader per bank. The activation multicast is one semaphore-gated block per sender, chained across senders, so the chain costs the block count K / in0_block_w. The factory lets a block span several consecutive activation shards, which takes the block count off the shard grid. Widening is not free. A block wider than a shard is gathered over the NoC into the sender before the multicast, so those shards cross the NoC twice. The multicast overlaps the block's weight read, so the chain is exposed only while a hop costs more than the weight a worker reads for that block: once bw x per_worker_N covers it, a wider block buys nothing and still pays the gather. Sweeping every legal in0_block_w over 21 decode shapes on P300 (Llama-3.2-1B, Llama-3.1-8B, Qwen3-8B and Qwen3-32B at one, four and eight devices) puts a hop at about HOP_IN_WEIGHT_TILES weight tiles: the rule then picks the measured optimum on 20 of the 21, one step short on the last, and no width slower than stock on any. 100 and 112 give the same result, so it sits in a plateau rather than on a cliff. The constant is a hardware ratio, one hop's fixed cost over the time a worker takes to stream one weight tile, so it moves with the bank count, the DRAM bandwidth and the NoC, not with the model. Anything that changes the fixed cost of a hop moves it down and this has to be refitted. """ HOP_IN_WEIGHT_TILES = 100 k_tiles = k // ttnn.TILE_SIZE per_core_k = k_tiles // num_cores per_worker_n = math.ceil(n / (ttnn.TILE_SIZE * self.dram_grid_size.x)) l1_budget_tiles = 800 * 1024 // (3 * 1088) base = self.find_largest_divisor(per_core_k) # Only ever widen, and only while a block's own weight read is shorter than one hop. # Widening halves the hops and doubles the weight a hop has to hide behind, so past the # crossover it buys nothing and still pays the gather. The L1 budget is the looser of the # two bounds and stays as a cap. A wide-N projection such as the LM head is already past # the crossover at the stock width, which is why the stock width is also the floor. for bw in range(min(16, k_tiles), base, -1): if k_tiles % bw: continue if per_core_k % bw and bw % per_core_k: continue if bw * per_worker_n > HOP_IN_WEIGHT_TILES: continue if bw * per_worker_n <= l1_budget_tiles: return bw return base def dram_shard_core_grid_for_k(self, k: int) -> Tuple[int, int]: rows, cols = self.find_grid(k // ttnn.TILE_SIZE) return ttnn.CoreGrid(x=cols, y=rows) def find_grid(self, N): """ Find the number of rows and columns for a grid of cores such that the total number of tiles N can be evenly divided among the cores. Each core will have the same integer number of tiles. The grid size is limited to a maximum of 2 rows and 8 columns. Parameters: N (int): Total number of tiles to be distributed. Returns: tuple: A tuple (rows, cols) representing the grid dimensions. Raises: AssertionError: If it's not possible to find such a grid configuration. """ max_rows = 8 if is_wormhole_b0() else 10 max_cols = 8 if is_wormhole_b0() else 12 max_cores = max_rows * max_cols # Find all possible numbers of cores that divide N and are less than or equal to max_cores target = 32 possible_cores = [k for k in range(1, max_cores + 1) if N % k == 0] possible_cores.sort(key=lambda x: abs(x - target)) # Sort by closest to target for cores in possible_cores: # Try to find a grid configuration with the current number of cores for rows in range(1, max_rows + 1): if cores % rows == 0: cols = cores // rows if cols <= max_cols: return rows, cols # If no configuration is found, assert an error raise AssertionError( f"Cannot find a grid configuration for {N} tiles that evenly divides into {max_cores} cores of max size {max_rows}x{max_cols}." ) def find_prefill_grid(self, row_tiles, col_tiles): """Find a grid such that the number of row tiles evenly divides into the number of rows and the number of column tiles evenly divides into the number of columns """ max_rows = 8 max_cols = 8 # TODO Improve configuration for BH (higher core grid than WH) # Find number of cols that evenly divides into the number of columns cols = None rows = None for i in range(max_cols, 0, -1): if col_tiles % i == 0: cols = i break for i in range(max_rows, 0, -1): if row_tiles % i == 0: rows = i break assert cols is not None, f"Cannot find a number of columns that evenly divides into {col_tiles}, not even 1(!)." assert rows is not None, f"Cannot find a number of rows that evenly divides into {row_tiles}, not even 1(!)." return rows, cols def dram_shard_core_grid_for_k_and_n(self, k: int, n: int) -> Tuple[int, int]: rows, cols = self.find_grid_k_n(k // ttnn.TILE_SIZE, n // ttnn.TILE_SIZE) return ttnn.CoreGrid(x=cols, y=rows) def find_grid_k_n(self, K, N): """ Find the number of rows and columns for a grid of cores such that the total number of tiles N can be evenly divided among the cores. Each core will have the same integer number of tiles. Parameters: N (int): Total number of tiles to be distributed. Returns: tuple: A tuple (rows, cols) representing the grid dimensions. Raises: AssertionError: If it's not possible to find such a grid configuration. """ max_rows = 8 max_cols = 8 # Maximum number of rows or columns max_cores = max_rows * max_cols # Maximum number of cores # Find all possible numbers of cores that divide N and are less than or equal to max_cores possible_cores = [c for c in range(1, max_cores + 1) if K % c == 0 and N % c == 0] possible_cores.sort(reverse=True) # Start checking from the largest number of cores for cores in possible_cores: # Try to find a grid configuration with the current number of cores for rows in range(1, max_rows + 1): if cores % rows == 0: cols = cores // rows if cols <= max_cols: return rows, cols # If no configuration is found, assert an error raise AssertionError( f"Cannot find a grid configuration such that both {K} and {N} tiles evenly divide into cores of max size {max_rows}x{max_cols}." ) def find_largest_divisor(self, n, max_divisor=8): for i in range(max_divisor, 0, -1): if n % i == 0: return i return 1 # Fallback to 1 if no divisor found def get_dram_sharded_matmul_num_workers(self, tensor_group: TensorGroup, n: int) -> int: """Return the validated P150 reader count for a Llama 3.1 8B decode projection.""" if self.base_model_name != "Llama-3.1-8B" or self.device_name != "P150": return 1 if tensor_group not in (TensorGroup.FF1_FF3, TensorGroup.FF2, TensorGroup.WQKV, TensorGroup.WO): return 1 num_workers = 2 shard_width_tiles = math.ceil(n / (ttnn.TILE_SIZE * self.dram_grid_size.x)) return num_workers if shard_width_tiles % num_workers == 0 else 1 def dram_matmul_config( self, m: int, k: int, n: int, num_cores=None, fused_activation=None, num_workers_per_dram_bank: int = 1, ): # in0_block_w must evenly divide k and be no larger than tile_size * num_cores if num_cores is None: # num_cores = self.dram_shard_core_grid_for_k(k).num_cores num_cores = self.dram_shard_core_grid_for_k_and_n(k, n).num_cores assert ( k % (ttnn.TILE_SIZE * num_cores) == 0 ), f"k must be divisible by tile_size * num_cores: {k} % {ttnn.TILE_SIZE * num_cores} != 0" # assert n % (ttnn.TILE_SIZE * num_cores) == 0, f"n must be divisible by tile_size * num_cores: {n} % {ttnn.TILE_SIZE * num_cores} != 0" if not self.is_galaxy and self.prefetcher is None: in0_block_w = self.dram_decode_in0_block_w(k, n, num_cores) else: in0_block_w = self.find_largest_divisor(k // (ttnn.TILE_SIZE * num_cores)) return ttnn.MatmulMultiCoreReuseMultiCastDRAMShardedProgramConfig( in0_block_w=in0_block_w, per_core_M=math.ceil(m / ttnn.TILE_SIZE), per_core_N=math.ceil(n / (ttnn.TILE_SIZE * num_cores)), fused_activation=fused_activation, num_workers_per_dram_bank=num_workers_per_dram_bank, ) def matmul_1d_ring_config( self, B, M, K, N, num_cores, num_global_cb_receivers, prefetch=True, untilize_out=False, ): M *= B # Fuse batch always enabled in0_block_w = K // num_cores // ttnn.TILE_SIZE out_block_h = M // ttnn.TILE_SIZE out_block_w = N // num_cores // ttnn.TILE_SIZE num_blocks_y = (M // ttnn.TILE_SIZE - 1) // out_block_h + 1 num_blocks_x = (N // ttnn.TILE_SIZE - 1) // out_block_w + 1 num_blocks_total = num_blocks_y * num_blocks_x if num_blocks_total != num_cores: assert False, f"num_blocks_total {num_blocks_total} != num_cores {num_cores}" out_subblock_h = 1 # For Blackhole with fp32_dest_acc_en=True, out_subblock_h * out_subblock_w must be <= 4 # For Wormhole/Grayskull with fp32_dest_acc_en=False, the limit is 8 max_subblock_w = 8 out_subblock_w = max_subblock_w while out_block_w % out_subblock_w != 0: out_subblock_w -= 1 hop_grid = [] # FIXME: Make not hard coded hop_core_range_set = ttnn.CoreRangeSet( { ttnn.CoreRange( ttnn.CoreCoord(x, y), ttnn.CoreCoord(x, y), ) for x, y in hop_grid } ) grid = num_to_coregrid(num_cores) # tt-metal/ttnn/cpp/ttnn/operations/matmul/device/matmul_op_multi_core_reuse_mcast_1d_program_factory.cpp program_config = ttnn.MatmulMultiCoreReuseMultiCast1DProgramConfig( compute_with_storage_grid_size=(grid.x, grid.y), in0_block_w=in0_block_w, out_subblock_h=out_subblock_h, out_subblock_w=out_subblock_w, per_core_M=out_block_h, per_core_N=out_block_w, fuse_batch=True, fused_activation=None, mcast_in0=False, gather_in0=True, hop_cores=hop_core_range_set, num_global_cb_receivers=num_global_cb_receivers if prefetch else 1, untilize_out=untilize_out, ) return program_config def matmul_1d_config( self, m, k, n, grid=ttnn.CoreGrid(x=8, y=8), act=None, is_fp32_accumulate=False, overwrite_per_core_k=None, overwrite_subblock_w=None, overwrite_subblock_h=None, ): tile_width = ttnn.TILE_SIZE tile_height = ttnn.TILE_SIZE if ( n // tile_width // grid.num_cores < 1 ): # use less number of cores in case we have more N num tiles than cores # assert (n // tile_width) % grid.x == 0 grid_y = n // tile_width // grid.x grid = ttnn.CoreGrid(x=grid.x, y=grid_y) per_core_m = m // tile_height per_core_k = self.find_largest_divisor(k // (ttnn.TILE_SIZE * grid.num_cores)) per_core_n = math.ceil(n / tile_width / grid.num_cores) if is_fp32_accumulate: max_subblock_w_h = 4 else: max_subblock_w_h = 8 # find the largest value between 1 and 8 that is a factor of per_core_n # e.g. if per_core_n is 14, then out_subblock_w = 7 out_subblock_w = max([i for i in range(1, max_subblock_w_h + 1) if per_core_n % i == 0]) # find the largest value that is a factor of per_core_m such that # out_subblock_w * out_subblock_h <= 8 out_subblock_h = max( [ i for i in range(1, max_subblock_w_h + 1) if per_core_m % i == 0 and i * out_subblock_w <= max_subblock_w_h ] ) if overwrite_per_core_k is not None: per_core_k = overwrite_per_core_k if overwrite_subblock_w is not None: out_subblock_w = overwrite_subblock_w if overwrite_subblock_h is not None: out_subblock_h = overwrite_subblock_h return ttnn.MatmulMultiCoreReuseMultiCast1DProgramConfig( compute_with_storage_grid_size=(grid.x, grid.y), in0_block_w=per_core_k, out_subblock_h=out_subblock_h, out_subblock_w=out_subblock_w, per_core_M=per_core_m, per_core_N=per_core_n, fuse_batch=True, fused_activation=act, mcast_in0=True, ) def matmul_1d_config_from_tensor_shapes( self, in0_shape, in1_shape, grid=ttnn.CoreGrid(x=8, y=8), act=None, is_fp32_accumulate=False, overwrite_subblock_w=None, overwrite_subblock_h=None, ): m, k, n = in0_shape[0] * in0_shape[1] * in0_shape[2], in0_shape[3], in1_shape[3] return self.matmul_1d_config( m, k, n, grid, act, is_fp32_accumulate, overwrite_subblock_w=overwrite_subblock_w, overwrite_subblock_h=overwrite_subblock_h, ) def create_sharded_norm_config(self, grid): """Helper function to create LayerNormShardedMultiCoreProgramConfig for RMS NORM. Args: grid (ttnn.CoreGrid): Grid specification for the norm operation """ block_w = self.dim // grid.num_cores // ttnn.TILE_SIZE # Find largest value <= 4 that evenly divides block_w subblock_w = 4 while subblock_w > 0: if block_w % subblock_w == 0: break subblock_w -= 1 return ttnn.LayerNormShardedMultiCoreProgramConfig( compute_with_storage_grid_size=[grid.x, grid.y], subblock_w=subblock_w, block_h=self.tile_padded_batch_rows // ttnn.TILE_SIZE, block_w=block_w, inplace=False, ) def create_tokenizer(self): from transformers import AutoTokenizer # Mapping of base model names to their known tokenizer paths # These are the original models that have proper tokenizers base_model_tokenizer_mapping = { "Qwen2.5-0.5B": "Qwen/Qwen2.5-Coder-0.5B-Instruct", "Qwen2.5-1.5B": "Qwen/Qwen2.5-1.5B-Instruct", "Qwen2.5-3B": "Qwen/Qwen2.5-3B-Instruct", "Qwen2.5-7B": "Qwen/Qwen2.5-7B-Instruct", "Qwen2.5-14B": "Qwen/Qwen2.5-14B-Instruct", "Qwen2.5-32B": "Qwen/Qwen2.5-32B-Instruct", "Qwen2.5-72B": "Qwen/Qwen2.5-72B-Instruct", "Qwen2.5-VL-32B": "Qwen/Qwen2.5-VL-32B-Instruct", "Qwen2.5-VL-72B": "Qwen/Qwen2.5-VL-72B-Instruct", "Qwen3-VL-32B": "Qwen/Qwen3-VL-32B-Instruct", "Qwen2.5-32B-Instruct": "Qwen/Qwen2.5-32B-Instruct", "Llama-3-8B": "meta-llama/Llama-3-8B", "Meta-Llama-3-8B": "meta-llama/Meta-Llama-3-8B", "Llama-3.1-8B": "meta-llama/Llama-3.1-8B-Instruct", "Llama-3.1-70B": "meta-llama/Llama-3.1-70B-Instruct", "Llama-3.2-1B": "meta-llama/Llama-3.2-1B-Instruct", "Llama-3.2-3B": "meta-llama/Llama-3.2-3B-Instruct", "Llama-3.2-11B": "meta-llama/Llama-3.2-11B-Vision-Instruct", "Llama-3.2-90B": "meta-llama/Llama-3.2-90B-Vision-Instruct", "Mistral-7B": "mistralai/Mistral-7B-Instruct-v0.3", "Mistral-Small-3.1-24B": "mistralai/Mistral-Small-3.1-24B-Instruct-2503", "Phi-3-mini-128k-instruct": "microsoft/Phi-3-mini-128k-instruct", "gemma-3-4b": "google/gemma-3-4b-it", "gemma-3-27b": "google/gemma-3-27b-it", "Qwen3.6-27B": "Qwen/Qwen3.6-27B", "LFM2.5-VL-1.6B": "LiquidAI/LFM2.5-VL-1.6B", "EXAONE-4.5-33B": "LGAI-EXAONE/EXAONE-4.5-33B", } logger.info(f"Tokenizer path: {self.TOKENIZER_PATH}") logger.info(f"Model name: {self.model_name}") logger.info(f"Base model name: {self.base_model_name}") tokenizer = None try: # Try to load tokenizer from the original model path # If there is no Processor, it will return Tokenizer (useful for multimodal models) tokenizer = AutoTokenizer.from_pretrained( self.TOKENIZER_PATH, local_files_only=os.getenv("CI") == "true", trust_remote_code=self.trust_remote_code_hf, ) logger.info(f"Successfully loaded tokenizer from {self.TOKENIZER_PATH}") except Exception as e: logger.warning(f"Failed to load tokenizer from {self.TOKENIZER_PATH}: {e}") # Only try fallback if initial load failed if tokenizer is None: # Try to use base model tokenizer as fallback fallback_tokenizer_path = base_model_tokenizer_mapping.get(self.base_model_name) # If no direct match, try to infer from model name patterns if not fallback_tokenizer_path: model_name_lower = self.model_name.lower() if "qwen2.5" in model_name_lower and "0.5b" in model_name_lower: fallback_tokenizer_path = "Qwen/Qwen2.5-Coder-0.5B-Instruct" elif "qwen2.5" in model_name_lower and "1.5b" in model_name_lower: fallback_tokenizer_path = "Qwen/Qwen2.5-1.5B-Instruct" elif "qwen2.5" in model_name_lower and "3b" in model_name_lower: fallback_tokenizer_path = "Qwen/Qwen2.5-3B-Instruct" elif "qwen2.5" in model_name_lower and "7b" in model_name_lower: fallback_tokenizer_path = "Qwen/Qwen2.5-7B-Instruct" elif "qwen2.5" in model_name_lower and "14b" in model_name_lower: fallback_tokenizer_path = "Qwen/Qwen2.5-14B-Instruct" elif "qwen2.5" in model_name_lower and "32b" in model_name_lower: fallback_tokenizer_path = "Qwen/Qwen2.5-32B-Instruct" elif "qwen2.5" in model_name_lower and "72b" in model_name_lower: fallback_tokenizer_path = "Qwen/Qwen2.5-72B-Instruct" elif "qwen2.5" in model_name_lower and "32b" in model_name_lower: fallback_tokenizer_path = "Qwen/Qwen2.5-32B-Instruct" elif "llama" in model_name_lower and "3.1" in model_name_lower and "8b" in model_name_lower: fallback_tokenizer_path = "meta-llama/Llama-3.1-8B-Instruct" elif "llama" in model_name_lower and "3.1" in model_name_lower and "70b" in model_name_lower: fallback_tokenizer_path = "meta-llama/Llama-3.1-70B-Instruct" elif "llama" in model_name_lower and "3.2" in model_name_lower and "1b" in model_name_lower: fallback_tokenizer_path = "meta-llama/Llama-3.2-1B-Instruct" elif "llama" in model_name_lower and "3.2" in model_name_lower and "3b" in model_name_lower: fallback_tokenizer_path = "meta-llama/Llama-3.2-3B-Instruct" elif "mistral" in model_name_lower and "7b" in model_name_lower: fallback_tokenizer_path = "mistralai/Mistral-7B-Instruct-v0.3" elif "mistral" in model_name_lower and "small" in model_name_lower and "24b" in model_name_lower: fallback_tokenizer_path = "mistralai/Mistral-Small-3.1-24B-Instruct-2503" elif "phi-3-mini" in model_name_lower and "128k" in model_name_lower and "instruct" in model_name_lower: fallback_tokenizer_path = "microsoft/Phi-3-mini-128k-instruct" if fallback_tokenizer_path: logger.info(f"Attempting to use fallback tokenizer: {fallback_tokenizer_path}") try: tokenizer = AutoTokenizer.from_pretrained( fallback_tokenizer_path, local_files_only=os.getenv("CI") == "true" ) logger.info(f"Successfully loaded fallback tokenizer from {fallback_tokenizer_path}") except Exception as fallback_e: logger.error(f"Failed to load fallback tokenizer from {fallback_tokenizer_path}: {fallback_e}") raise fallback_e else: logger.error(f"No fallback tokenizer found for base model: {self.base_model_name}") raise Exception(f"No fallback tokenizer found for base model: {self.base_model_name}") # Add meta-compatible stop token list to the HF tokenizer if not hasattr(tokenizer, "stop_tokens") or tokenizer.stop_tokens is None: tokenizer.stop_tokens = [tokenizer.eos_token_id] # Phi-3-mini uses "<|end|>" as EOS token if "phi-3-mini" in self.base_model_name.lower(): tokenizer.stop_tokens.append(tokenizer.encode("<|end|>")[0]) return tokenizer def create_processor(self): from transformers import AutoProcessor processor = None try: processor = AutoProcessor.from_pretrained(self.TOKENIZER_PATH, local_files_only=os.getenv("CI") == "true") logger.info(f"Successfully loaded processor from {self.TOKENIZER_PATH}") except Exception as e: logger.warning(f"Failed to load processor from {self.TOKENIZER_PATH}: {e}") return processor def encode_prompt(self, prompt_text, system_prompt_text=None, instruct=True): if instruct: try: return encode_prompt_hf(self.tokenizer, prompt_text, system_prompt_text) except ValueError as e: logger.warning(f"Failed to encode chat prompt, are you sure this is an instruct model? Error: {e}") logger.warning(f"Falling back to base model encoding with no chat template") return self.tokenizer.encode(prompt_text, add_special_tokens=False) def reference_lm_head(self): model = self.reference_transformer(wrap=False) layer = model.lm_head layer._load_state_dict = layer.load_state_dict if self.use_hf_rope: layer.load_state_dict = lambda x: layer._load_state_dict(convert_meta_to_hf_no_qkv_permute(x)) else: layer.load_state_dict = lambda x: layer._load_state_dict(convert_meta_to_hf(x, self.head_dim)) return layer def reference_transformer(self, wrap=True, load_checkpoint=False): from transformers import AutoConfig model_cls = self.get_hf_model_cls() # HF is much faster at loading from a checkpoint than generating from config # so use that by preference unless we don't have a checkpoint if self.dummy_weights and not load_checkpoint: config = AutoConfig.from_pretrained( self.LOCAL_HF_PARAMS[self.model_name], trust_remote_code=self.trust_remote_code_hf, local_files_only=os.getenv("CI") == "true", ) if hasattr(config, "text_config"): config.text_config.num_layers = self.n_layers config.text_config.num_hidden_layers = self.n_layers else: config.num_layers = self.n_layers config.num_hidden_layers = self.n_layers if os.getenv("CI") == "true": # In CI, from_pretrained on NFS can spend ~54s scanning before failing for models # that don't have clean checkpoint artifacts. Skip straight to from_config. model = model_cls.from_config(config, trust_remote_code=self.trust_remote_code_hf) else: try: # .from_pretrained + _init_weights works faster than .from_config model = model_cls.from_pretrained( self.CKPT_DIR, config=config, torch_dtype="auto", trust_remote_code=self.trust_remote_code_hf, local_files_only=True, ) model.apply(model._init_weights) except Exception as e: logger.info(f"Error loading dummy weights using .from_pretrained. Using .from_config. Error: {e}") model = model_cls.from_config(config, trust_remote_code=self.trust_remote_code_hf) # model.load_state_dict({k: torch.randn_like(v) for k, v in model.state_dict().items()}) else: model_cls = self.get_hf_model_cls() # HF is much faster at loading from a checkpoint than generating from config # so use that by preference unless we don't have a checkpoint if self.dummy_weights and not load_checkpoint: config = AutoConfig.from_pretrained( self.LOCAL_HF_PARAMS[self.model_name], trust_remote_code=self.trust_remote_code_hf, local_files_only=os.getenv("CI") == "true", ) if hasattr(config, "text_config"): config.text_config.num_layers = self.n_layers config.text_config.num_hidden_layers = self.n_layers else: config.num_layers = self.n_layers config.num_hidden_layers = self.n_layers if os.getenv("CI") == "true": # In CI, from_pretrained on NFS can spend ~54s scanning before failing for models # that don't have clean checkpoint artifacts. Skip straight to from_config. model = model_cls.from_config(config, trust_remote_code=self.trust_remote_code_hf) else: try: # .from_pretrained + _init_weights works faster than .from_config model = model_cls.from_pretrained( self.CKPT_DIR, config=config, torch_dtype="auto", trust_remote_code=self.trust_remote_code_hf, local_files_only=True, ) model.apply(model._init_weights) except Exception as e: logger.info( f"Error loading dummy weights using .from_pretrained. Using .from_config. Error: {e}" ) model = model_cls.from_config(config, trust_remote_code=self.trust_remote_code_hf) # model.load_state_dict({k: torch.randn_like(v) for k, v in model.state_dict().items()}) else: if self.cache_hf_flag and self.cached_hf_model is None: model = model_cls.from_pretrained( self.CKPT_DIR, torch_dtype="auto", local_files_only=os.getenv("CI") == "true", trust_remote_code=self.trust_remote_code_hf, ) self.cached_hf_model = model elif self.cache_hf_flag and self.cached_hf_model is not None: model = self.cached_hf_model else: # No caching - load fresh each time model = model_cls.from_pretrained( self.CKPT_DIR, torch_dtype="auto", trust_remote_code=self.trust_remote_code_hf, local_files_only=os.getenv("CI") == "true", ) # HACK: Assume that we want the language model layers only. # transformers 5.x nests the text model under model.model.language_model for # multimodal models (e.g. Mllama); <5 exposed model.language_model directly. if hasattr(model, "language_model"): # transformers <5 multimodal model.model = model.language_model # We keep language_model because transformers don't let us change or delete it elif hasattr(model.model, "language_model"): # transformers >=5 multimodal model.model = model.model.language_model model.model.layers = model.model.layers[: self.n_layers] if wrap: wrapper = HfModelWrapper(model, self.head_dim, config=self.hf_config, use_hf_rope=self.use_hf_rope) return wrapper else: return model def reference_vision_multi_modal(self): model = self.reference_vision_transformer(wrap=False) # transformers 5.x nests multi_modal_projector under model.model layer = getattr(model, "multi_modal_projector", None) or model.model.multi_modal_projector layer._load_state_dict = layer.load_state_dict layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_vision_rms_norm(self): model = self.reference_vision_transformer(wrap=False) # transformers 5.x nests multi_modal_projector under model.model mmp = getattr(model, "multi_modal_projector", None) or model.model.multi_modal_projector layer = mmp.mm_soft_emb_norm layer._load_state_dict = layer.load_state_dict layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_rms_norm(self): # Always HuggingFace since we only support HF_MODEL now model = self.reference_transformer(wrap=False) layers = getattr(model, "layers", getattr(model, "model", {}).layers) layer = layers[0].input_layernorm layer._load_state_dict = layer.load_state_dict if self.use_hf_rope: layer.load_state_dict = lambda x: layer._load_state_dict(convert_meta_to_hf_no_qkv_permute(x)) else: layer.load_state_dict = lambda x: layer._load_state_dict(convert_meta_to_hf(x, self.head_dim)) return layer def reference_vision_transformer(self, wrap=True, load_checkpoint=False): # Always HuggingFace since we only support HF_MODEL now from transformers import AutoConfig model_cls = self.get_hf_model_cls() if self.dummy_weights and not load_checkpoint: config = AutoConfig.from_pretrained(self.LOCAL_HF_PARAMS[self.model_name]) if hasattr(config, "text_config"): config.text_config.num_layers = self.n_layers config.text_config.num_hidden_layers = self.n_layers else: config.num_layers = self.n_layers config.num_hidden_layers = self.n_layers if os.getenv("CI") == "true": # In CI, from_pretrained on NFS can spend ~54s scanning before failing for models # that don't have clean checkpoint artifacts. Skip straight to from_config. model = model_cls.from_config(config, trust_remote_code=self.trust_remote_code_hf) else: try: # .from_pretrained + _init_weights works faster than .from_config model = model_cls.from_pretrained( self.CKPT_DIR, config=config, torch_dtype="auto", trust_remote_code=self.trust_remote_code_hf, local_files_only=True, ) model.apply(model._init_weights) except Exception as e: logger.info(f"Error loading dummy weights using .from_pretrained. Using .from_config. Error: {e}") model = model_cls.from_config(config, trust_remote_code=self.trust_remote_code_hf) # model.load_state_dict({k: torch.randn_like(v) for k, v in model.state_dict().items()}) else: if self.cached_hf_model is None: model = model_cls.from_pretrained( self.CKPT_DIR, torch_dtype="auto", local_files_only=os.getenv("CI") == "true" ) self.cached_hf_model = model else: model = self.cached_hf_model inner = model.model if hasattr(inner, "layers"): inner.layers = inner.layers[: self.n_layers] elif hasattr(inner, "language_model") and hasattr(inner.language_model, "layers"): inner.language_model.layers = inner.language_model.layers[: self.n_layers] if wrap: wrapper = HfModelWrapper(model, self.head_dim, use_hf_rope=self.use_hf_rope) return wrapper else: return model def reference_gemma_model(self): model = self.reference_vision_transformer(wrap=False) layer = model layer._load_state_dict = layer.load_state_dict layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_vision_model(self): model = self.reference_vision_transformer(wrap=False) if "Mistral-Small-3.1-24B-Instruct-2503" in self.model_name: # Mistral-Small-3.1-24B-Instruct-2503 has a different structure layer = self._get_vision_tower(model) else: layer = self._get_vision_tower(model).vision_model layer._load_state_dict = layer.load_state_dict layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_vision_mlp(self, layer_idx=0): model = self.reference_vision_transformer(wrap=False) vision_tower = self._get_vision_tower(model) if "Mistral-Small-3.1-24B" in self.model_name: layer = vision_tower.transformer.layers[layer_idx].feed_forward else: layer = vision_tower.vision_model.encoder.layers[0].mlp layer._load_state_dict = layer.load_state_dict layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_siglip_patch_embed(self): model = self.reference_vision_transformer(wrap=False) layer = self._get_vision_tower(model).vision_model.embeddings.patch_embedding # layer._load_state_dict = layer.load_state_dict # layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_vision_pos_embedding(self): model = self.reference_vision_transformer(wrap=False) layer = self._get_vision_tower(model).vision_model.embeddings.position_embedding # layer._load_state_dict = layer.load_state_dict # layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_vision_embedding(self): model = self.reference_vision_transformer(wrap=False) layer = self._get_vision_tower(model).vision_model.embeddings # layer._load_state_dict = layer.load_state_dict # layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_vision_layernorm(self, layer_name="layer_norm1"): model = self.reference_vision_transformer(wrap=False) if layer_name == "layer_norm1": layer = self._get_vision_tower(model).vision_model.encoder.layers[0].layer_norm1 elif layer_name == "layer_norm2": layer = self._get_vision_tower(model).vision_model.encoder.layers[0].layer_norm2 else: layer = self._get_vision_tower(model).vision_model.post_layernorm # layer._load_state_dict = layer.load_state_dict # layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_vision_attention(self, layer_idx=0): model = self.reference_vision_transformer(wrap=False) vision_tower = self._get_vision_tower(model) if "Mistral-Small-3.1-24B" in self.model_name: layer = vision_tower.transformer.layers[layer_idx].attention else: layer = vision_tower.vision_model.encoder.layers[0].self_attn layer._load_state_dict = layer.load_state_dict layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_vision_encoder_block(self): model = self.reference_vision_transformer(wrap=False) layer = self._get_vision_tower(model).vision_model.encoder.layers[0] # layer._load_state_dict = layer.load_state_dict # layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def _get_vision_tower(self, model): """Handle different HF model architectures: Mistral3 nests vision_tower under model.model, while Mllama has it directly on model.""" if hasattr(model, "vision_tower"): return model.vision_tower return model.model.vision_tower def reference_vision_encoder(self): model = self.reference_vision_transformer(wrap=False) vision_tower = self._get_vision_tower(model) if "Mistral-Small-3.1-24B" in self.model_name: layer = vision_tower.transformer else: layer = vision_tower.vision_model.encoder layer._load_state_dict = layer.load_state_dict layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_pixtral_image_block(self, layer_num=0): model = self.reference_vision_transformer(wrap=False) vision_tower = self._get_vision_tower(model) layer = vision_tower.transformer.layers[layer_num] layer._load_state_dict = layer.load_state_dict layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_vision_rms(self): model = self.reference_vision_transformer(wrap=False) vision_tower = self._get_vision_tower(model) layer = vision_tower.transformer.layers[0].ffn_norm layer._load_state_dict = layer.load_state_dict layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_conv2d_patch(self): model = self.reference_vision_transformer(wrap=False) vision_tower = self._get_vision_tower(model) layer = vision_tower.patch_conv layer._load_state_dict = layer.load_state_dict layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_vision_rot_emb(self): model = self.reference_vision_transformer(wrap=False) vision_tower = self._get_vision_tower(model) if "Mistral-Small-3.1-24B" in self.model_name: layer = vision_tower.patch_positional_embedding else: raise NotImplementedError(f"reference_vision_rot_emb not implemented for {self.model_name}") layer._load_state_dict = layer.load_state_dict layer.load_state_dict = lambda x: layer._load_state_dict(convert_vision_meta_to_hf(x, self.head_dim)) return layer def reference_mlp(self): model = self.reference_transformer(wrap=False) layer = model.model.layers[0].mlp layer._load_state_dict = layer.load_state_dict if self.use_hf_rope: layer.load_state_dict = lambda x: layer._load_state_dict( convert_meta_to_hf_no_qkv_permute(x, fuse_mlp=self.fuse_mlp) ) else: layer.load_state_dict = lambda x: layer._load_state_dict( convert_meta_to_hf(x, self.head_dim, fuse_mlp=self.fuse_mlp) ) return layer def reference_embedding(self, reference_model=None): if reference_model is None: model = self.reference_transformer(wrap=False) layer = model.model.embed_tokens else: layer = reference_model.model.model.embed_tokens layer._load_state_dict = layer.load_state_dict if self.use_hf_rope: layer.load_state_dict = lambda x: layer._load_state_dict(convert_meta_to_hf_no_qkv_permute(x)) else: layer.load_state_dict = lambda x: layer._load_state_dict(convert_meta_to_hf(x, self.head_dim)) return layer def reference_decoder(self, load_checkpoint=False): model = self.reference_transformer(wrap=False, load_checkpoint=load_checkpoint) layer = model.model.layers[0] use_position_embeddings = layer.__class__.__name__ != "Phi3DecoderLayer" or self.base_model_name in ("phi-4",) if hasattr(model.model, "rotary_emb_local"): rotary_emb_local = model.model.rotary_emb_local else: rotary_emb_local = None wrapper = HfDecoderWrapper( layer, self.head_dim, model.model.rotary_emb if use_position_embeddings else None, rotary_emb_local, self.use_hf_rope, ) return wrapper def reference_attention(self, load_checkpoint=False): model = self.reference_transformer(wrap=False, load_checkpoint=load_checkpoint) layer = model.model.layers[0].self_attn use_position_embeddings = "position_embeddings" in inspect.signature(layer.forward).parameters wrapper = HfAttentionWrapper( layer, self.head_dim, model.model.rotary_emb if use_position_embeddings else None, use_hf_rope=self.use_hf_rope, ) return wrapper def set_tg_attention_config(self): shard_spec_n_cores_grid = ttnn.CoreRangeSet({num_to_corerange(40)}) self.model_config["CREATE_HEAD_INPUT_MEMCFG"] = ( None if self.dim < 4096 else ttnn.MemoryConfig( ttnn.TensorMemoryLayout.WIDTH_SHARDED, ttnn.BufferType.L1, ttnn.ShardSpec( shard_spec_n_cores_grid, [ 32, 32, ], ttnn.ShardOrientation.ROW_MAJOR, ), ) ) if self.is_galaxy: num_cores = 40 if self.dim == 8192 else (24 if self.dim == 4096 else (20 if self.dim == 3072 else 12)) self.model_config["QKV_OUT_GATHERED_MEMCFG"] = lambda mesh_cols: ttnn.create_sharded_memory_config( shape=(32 * mesh_cols, 32), # mesh_cols = 4 core_grid=num_to_coregrid(num_cores), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) # Prefer a tile-aligned width shard, but never hand num_to_coregrid a core # count it cannot map -- it returns None, which create_sharded_memory_config # rejects with "Invalid core_grid type". Fall back to the legacy expression # in that case so dims that built a config before still build one. self_out_width = self.dim // 4 self_out_cores = compute_galaxy_width_shard_cores(self_out_width) if num_to_coregrid(self_out_cores) is None: self_out_cores = min(32, self_out_width // ttnn.TILE_SIZE) self_out_grid = num_to_coregrid(self_out_cores) self.model_config["SELF_OUT_GATHERED_MEMCFG"] = lambda mesh_rows: ttnn.create_sharded_memory_config( shape=(32 * mesh_rows, self_out_width // self_out_cores), core_grid=self_out_grid, strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) # The gathered attention output is n_local_heads * head_dim wide, which is # only equal to dim // cluster_shape[0] when n_heads * head_dim == dim. # Qwen3-32B (dim 5120, 64 heads, 8 KV heads, head_dim 128) gathers 1024 # columns, while Qwen2.5-32B / QwQ-32B (40 heads) gather 640 -- so keying # this off dim would break the latter. Derive it from the head geometry. gather_users_cores = min(32, (self.n_heads // self.n_kv_heads) * self.head_dim // ttnn.TILE_SIZE) self.model_config["GATHER_USERS_MEMCFG"] = lambda mesh_cols: ttnn.create_sharded_memory_config( shape=(32 * mesh_cols, 32), # mesh_cols = 4 core_grid=num_to_coregrid(gather_users_cores), strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) else: qkv_core_grid = self.dram_shard_core_grid_for_k(self.dim) self.model_config["QKV_OUT_GATHERED_MEMCFG"] = lambda mesh_rows: ttnn.create_sharded_memory_config( ( ttnn.TILE_SIZE * mesh_rows, self.dim // qkv_core_grid.num_cores, ), # Shard shape: [32, 128] -> 1 shard per core core_grid=qkv_core_grid, strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) gather_core_grid = self.dram_shard_core_grid_for_k(self.dim // 4) self.model_config["SELF_OUT_GATHERED_MEMCFG"] = lambda mesh_rows: ttnn.create_sharded_memory_config( ( ttnn.TILE_SIZE * mesh_rows, self.dim // 4 // gather_core_grid.num_cores, ), core_grid=gather_core_grid, strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) users_core_grid = self.dram_shard_core_grid_for_k(self.dim // 8) self.model_config["GATHER_USERS_MEMCFG"] = lambda mesh_cols: ttnn.create_sharded_memory_config( ( ttnn.TILE_SIZE * mesh_cols, self.dim // 8 // users_core_grid.num_cores, ), core_grid=users_core_grid, strategy=ttnn.ShardStrategy.WIDTH, orientation=ttnn.ShardOrientation.ROW_MAJOR, use_height_and_width_as_shard_shape=True, ) class HfAttentionWrapper: def __init__(self, attention, head_dim, rotary_emb, use_hf_rope=False, rope_layer_type=None): from transformers import DynamicCache super().__init__() self.attention = attention self.past_key_value = DynamicCache() self.head_dim = head_dim self.rotary_emb = rotary_emb self.use_hf_rope = use_hf_rope # transformers 5.x Gemma3 rotary picks `{layer_type}_inv_freq`. When the caller chose a # specific rope module (e.g. global vs local), pin the layer_type to match it instead of # the attention layer's own type; otherwise fall back to the attention's layer_type. self.rope_layer_type = rope_layer_type def forward(self, x, start_pos, freqs_cis_i, mask=None): position_ids = torch.tensor([list(range(start_pos, start_pos + x.shape[1]))] * x.shape[0]) if mask is not None: while len(mask.shape) < 4: mask = mask.unsqueeze(0) # transformers 5.x renamed the attention cache kwarg past_key_value -> past_key_values. # Passing the wrong name lets it fall into **kwargs (ignored), so the cache is never # populated. Pick the name present in the attention's signature. cache_kw = ( "past_key_values" if "past_key_values" in inspect.signature(self.attention.forward).parameters else "past_key_value" ) if self.rotary_emb is not None: # transformers 5.x Gemma3 uses per-layer-type RoPE: the rotary forward takes a # `layer_type` and selects `{layer_type}_inv_freq` (layer_type=None -> AttributeError # 'None_inv_freq'). Pass the wrapped attention's own layer_type when the rotary accepts # it (the authoritative type for this layer); <5 / non-Gemma3 rotaries don't take it. _layer_type = ( self.rope_layer_type if self.rope_layer_type is not None else getattr(self.attention, "layer_type", None) ) if _layer_type is not None and "layer_type" in inspect.signature(self.rotary_emb.forward).parameters: position_embeddings = self.rotary_emb(x, position_ids, layer_type=_layer_type) else: position_embeddings = self.rotary_emb(x, position_ids) output, *_ = self.attention( x, position_embeddings=position_embeddings, use_cache=True, attention_mask=mask, **{cache_kw: self.past_key_value}, ) else: output, _, self.past_key_value = self.attention( x, use_cache=True, position_ids=position_ids, attention_mask=mask, **{cache_kw: self.past_key_value}, ) return output def __call__(self, *args, **kwargs): return self.forward(*args, **kwargs) def load_state_dict(self, state_dict): if self.use_hf_rope: raise NotImplementedError("Not supported if `use_hf_rope` is True") try: # Checking for fused qkv layer fuse_qkv = hasattr(self.attention, "qkv_proj") except: fuse_qkv = False if self.use_hf_rope: return self.attention.load_state_dict(convert_meta_to_hf_no_qkv_permute(state_dict, fuse_qkv)) else: return self.attention.load_state_dict(convert_meta_to_hf(state_dict, self.head_dim, fuse_qkv)) @property def cache_k(self): [(k, v)] = [(kk, vv) for (kk, vv) in hf_cache_to_legacy(self.past_key_value) if kk is not None] hf_k = k.permute(0, 2, 1, 3) # match meta-style reference which uses (batch_size, seq, n_kv_heads, head_dim) if self.use_hf_rope: # No transformation needed for HF-style RoPE return hf_k # Llama-style: apply reverse_permute transformation batch_size, seq_len, n_heads, head_dim = hf_k.shape meta_k = torch.zeros_like(hf_k) for b in range(batch_size): for s in range(seq_len): # Flatten just heads and head_dim flat = hf_k[b, s].flatten() # Apply reverse_permute transformed = reverse_permute(flat.unsqueeze(-1), n_heads, flat.shape[0], 1).squeeze(-1) # Restore heads and head_dim shape meta_k[b, s] = transformed.reshape(n_heads, head_dim) return meta_k @property def cache_v(self): [(k, v)] = [(kk, vv) for (kk, vv) in hf_cache_to_legacy(self.past_key_value) if kk is not None] return v.permute(0, 2, 1, 3) # match meta-style reference which uses (batch_size, seq, n_kv_heads, head_dim) class HfDecoderWrapper: def __init__(self, decoder, head_dim, rotary_emb, rotary_emb_local=None, use_hf_rope=False): from transformers import DynamicCache self.decoder = decoder self.head_dim = head_dim self.rotary_emb = rotary_emb self.rotary_emb_local = rotary_emb_local self.past_key_values = DynamicCache() self.use_hf_rope = use_hf_rope def forward(self, x, start_pos, freqs_cis_i, mask=None): position_ids = torch.tensor([list(range(start_pos, start_pos + x.shape[1]))] * x.shape[0]) # transformers 5.x consolidated Gemma3 RoPE into a module that selects `{layer_type}_inv_freq` # (layer_type=None -> AttributeError 'None_inv_freq'). Pass the matching layer_type when the # rotary forward accepts it (global rotary -> full_attention, local -> sliding_attention); # pre-5.x / non-Gemma3 rotaries don't take the kwarg. position_embeddings = None if self.rotary_emb is not None: if "layer_type" in inspect.signature(self.rotary_emb.forward).parameters: position_embeddings = self.rotary_emb(x, position_ids, layer_type="full_attention") else: position_embeddings = self.rotary_emb(x, position_ids) if mask is not None: while len(mask.shape) < 4: mask = mask.unsqueeze(0) # transformers 5.x renamed the decoder-layer cache kwarg past_key_value -> past_key_values. cache_kw = ( "past_key_values" if "past_key_values" in inspect.signature(self.decoder.forward).parameters else "past_key_value" ) if self.rotary_emb_local is not None: # transformers <5 Gemma3 decoder layer takes split global/local rope and selects via # is_sliding internally; pass both (global=full_attention, local=sliding_attention). if "layer_type" in inspect.signature(self.rotary_emb_local.forward).parameters: position_embeddings_local = self.rotary_emb_local(x, position_ids, layer_type="sliding_attention") else: position_embeddings_local = self.rotary_emb_local(x, position_ids) result = self.decoder.forward( x, position_embeddings_global=position_embeddings, position_embeddings_local=position_embeddings_local, use_cache=True, position_ids=position_ids, attention_mask=mask, **{cache_kw: self.past_key_values}, ) else: # transformers 5.x Gemma3 decoder layer takes a single `position_embeddings` and expects # the CALLER to supply the rope for this layer's type (the model picks full_attention vs # sliding_attention per layer). Layer 0 is sliding, so computing global rope here (as the # default above does) feeds the wrong rope and the reference diverges from the TT decoder # (which applies the correct per-layer rope). Recompute with this layer's own layer_type. _layer_type = getattr(getattr(self.decoder, "self_attn", None), "layer_type", None) if ( self.rotary_emb is not None and _layer_type is not None and "layer_type" in inspect.signature(self.rotary_emb.forward).parameters ): position_embeddings = self.rotary_emb(x, position_ids, layer_type=_layer_type) result = self.decoder.forward( x, position_embeddings=position_embeddings, use_cache=True, position_ids=position_ids, attention_mask=mask, **{cache_kw: self.past_key_values}, ) # transformers 5.x decoder layers return the hidden-states tensor directly instead of a # (hidden_states, ...) tuple; only unwrap [0] when it's actually a tuple, otherwise result[0] # would index the batch dim and drop a leading dimension (e.g. [1,1,dim] -> [1,dim]). output = result[0] if isinstance(result, tuple) else result return output def __call__(self, *args, **kwargs): return self.forward(*args, **kwargs) def load_state_dict(self, state_dict): try: # Checking for fused qkv and mlp layers fuse_qkv = hasattr(self.decoder.self_attn, "qkv_proj") fuse_mlp = hasattr(self.decoder.mlp, "gate_up_proj") except: fuse_qkv, fuse_mlp = False, False if self.use_hf_rope: return self.decoder.load_state_dict(convert_meta_to_hf_no_qkv_permute(state_dict, fuse_qkv, fuse_mlp)) else: return self.decoder.load_state_dict(convert_meta_to_hf(state_dict, self.head_dim, fuse_qkv, fuse_mlp)) @property def cache_k(self): [(k, v)] = [(kk, vv) for (kk, vv) in hf_cache_to_legacy(self.past_key_values) if kk is not None] hf_k = k.permute(0, 2, 1, 3) # match meta-style reference which uses (batch_size, seq, n_kv_heads, head_dim) if self.use_hf_rope: # No transformation needed for HF-style RoPE return hf_k # Llama-style: apply reverse_permute transformation batch_size, seq_len, n_heads, head_dim = hf_k.shape meta_k = torch.zeros_like(hf_k) for b in range(batch_size): for s in range(seq_len): # Flatten just heads and head_dim flat = hf_k[b, s].flatten() # Apply reverse_permute transformed = reverse_permute(flat.unsqueeze(-1), n_heads, flat.shape[0], 1).squeeze(-1) # Restore heads and head_dim shape meta_k[b, s] = transformed.reshape(n_heads, head_dim) return meta_k @property def cache_v(self): [(k, v)] = [(kk, vv) for (kk, vv) in hf_cache_to_legacy(self.past_key_values) if kk is not None] return v.permute(0, 2, 1, 3) # match meta-style reference which uses (batch_size, seq, n_kv_heads, head_dim) class HfModelWrapper: def __init__(self, model, head_dim, config=None, use_hf_rope=False): from transformers import DynamicCache self.model = model self.head_dim = head_dim self.config = config self.past_key_values = DynamicCache() self.use_hf_rope = use_hf_rope def forward(self, inputs_embeds, start_pos, mode="decode"): position_ids = torch.tensor( [list(range(start_pos, start_pos + inputs_embeds.shape[1]))] * inputs_embeds.shape[0] ) # In prefill mode the reference must match the TT model, which returns the last decoder # layer's output *before* the final norm. HF's output_hidden_states does not expose that # tensor: the tuple is (embeddings, out_0, ..., out_{N-2}, norm(out_{N-1})), so hidden_states[-2] # is the embeddings for a 1-layer model and the second-to-last layer otherwise. Capture the # input to the final norm via a forward pre-hook to get the true pre-norm last-layer output. captured = {} handle = None if mode != "decode": handle = self.model.model.norm.register_forward_pre_hook( lambda module, args: captured.__setitem__("pre_norm", args[0]) ) try: logits, new_cache, hidden_states = self.model.forward( inputs_embeds=inputs_embeds, position_ids=position_ids, use_cache=True, past_key_values=self.past_key_values, return_dict=False, output_hidden_states=True, ) finally: if handle is not None: handle.remove() self.past_key_values = new_cache return logits if mode == "decode" else captured["pre_norm"] def __call__(self, *args, **kwargs): return self.forward(*args, **kwargs) def load_state_dict(self, state_dict): try: # Checking for fused qkv and mlp layers fuse_qkv = hasattr(self.model.model.layers[0].self_attn, "qkv_proj") fuse_mlp = hasattr(self.model.model.layers[0].mlp, "gate_up_proj") except: fuse_qkv, fuse_mlp = False, False if self.use_hf_rope: return self.model.load_state_dict( convert_meta_to_hf_no_qkv_permute(state_dict, fuse_qkv, fuse_mlp, self.config) ) else: return self.model.load_state_dict( convert_meta_to_hf(state_dict, self.head_dim, fuse_qkv, fuse_mlp, self.config) ) def eval(self): self.model.eval() @property def cache_k(self): kvs = hf_cache_to_legacy(self.past_key_values) meta_ks = [] for k, v in kvs: hf_k = k.permute( 0, 2, 1, 3 ) # match meta-style reference which uses (batch_size, seq, n_kv_heads, head_dim) if self.use_hf_rope: # No transformation needed for HF-style RoPE meta_ks.append(hf_k) continue # Llama-style: apply reverse_permute transformation batch_size, seq_len, n_heads, head_dim = hf_k.shape meta_k = torch.zeros_like(hf_k) for b in range(batch_size): for s in range(seq_len): # Flatten just heads and head_dim flat = hf_k[b, s].flatten() # Apply reverse_permute transformed = reverse_permute(flat.unsqueeze(-1), n_heads, flat.shape[0], 1).squeeze(-1) # Restore heads and head_dim shape meta_k[b, s] = transformed.reshape(n_heads, head_dim) meta_ks.append(meta_k) return meta_ks @property def cache_v(self): kvs = hf_cache_to_legacy(self.past_key_values) return [ v.permute(0, 2, 1, 3) for k, v in kvs ] # match meta-style reference which uses (batch_size, seq, n_kv_heads, head_dim) class DecodersPrecision: @classmethod def from_string(cls, optimizations: str): if optimizations == "performance": return cls.performance elif optimizations == "accuracy": return cls.accuracy else: raise ValueError( f"Invalid optimization configuration: {optimizations}. Allowed values are 'performance' or 'accuracy'" ) @classmethod def accuracy(cls, num_decoders, model_name): inst = cls._precision_factory(num_decoders, model_name, ModelOptimizations.accuracy) inst.__name__ = "accuracy" return inst @classmethod def performance(cls, num_decoders, model_name): inst = cls._precision_factory(num_decoders, model_name, ModelOptimizations.performance) inst.__name__ = "performance" return inst def __init__(self, num_decoders, model_name, decoder_conf: dict = None): if decoder_conf is None: decoder_conf = ModelOptimizations.accuracy(model_name) self.decoder_optimizations = {decoder_id: decoder_conf for decoder_id in range(num_decoders)} self._update_full_name() def set_decoder_conf(self, decoder_id, conf: ModelOptimizations): self.decoder_optimizations[decoder_id] = conf self._update_full_name() def get_tensor_dtype(self, decoder_id, tensor: TensorGroup, prefetcher: bool = False): """ Get the dtype for a specific tensor in a decoder layer. Args: decoder_id: The decoder layer index tensor: The tensor group (FF1_FF3, FF2, WQKV, WO, etc.) prefetcher: If True, returns the same dtype for all layers (uses decoder 0's dtype). This is required when using the prefetcher to avoid race conditions caused by different block sizes across layers. Returns: The ttnn dtype for the tensor, or None if original dtype should be used. """ precision_setting_lookup = { PrecisionSetting.BFP4: ttnn.bfloat4_b, PrecisionSetting.BFP8: ttnn.bfloat8_b, PrecisionSetting.BF16: ttnn.bfloat16, None: None, # this signals that original dtype should be used } # When prefetcher is enabled, use decoder 0's dtype for all layers # to ensure consistent block sizes across all layers effective_decoder_id = 0 if prefetcher else decoder_id if ( effective_decoder_id not in self.decoder_optimizations or tensor not in self.decoder_optimizations[effective_decoder_id].tensor_dtype_settings ): # When prefetcher is enabled and no config exists, default to BFP8 for weight tensors # but keep None for ACTIVATION (which uses special handling in the callers) if prefetcher and tensor != TensorGroup.ACTIVATION: return ttnn.bfloat8_b return None key = self.decoder_optimizations[effective_decoder_id].tensor_dtype_settings[tensor] if key is None or key not in precision_setting_lookup: # When prefetcher is enabled and key is invalid, default to BFP8 for weight tensors # but keep None for ACTIVATION (which uses special handling in the callers) if prefetcher and tensor != TensorGroup.ACTIVATION: return ttnn.bfloat8_b return None return precision_setting_lookup[key] def get_math_fidelity(self, decoder_id, op: OpGroup, configuration: ModelArgs): math_fidelity_setting_lookup = { MathFidelitySetting.LOFI: configuration.compute_kernel_config_lofi, MathFidelitySetting.HIFI2: configuration.compute_kernel_config_hifi2, MathFidelitySetting.HIFI2_NA: configuration.compute_kernel_config_hifi2_na, MathFidelitySetting.HIFI2_FP16: configuration.compute_kernel_config_hifi2_fp16, MathFidelitySetting.HIFI2_NOL1ACC: configuration.compute_kernel_config_hifi2_nol1acc, MathFidelitySetting.HIFI4: configuration.compute_kernel_config_hifi4, MathFidelitySetting.HIFI4_FP16: configuration.compute_kernel_config_hifi4_fp16, MathFidelitySetting.HIFI4_FP32: configuration.compute_kernel_config_hifi4_fp32, } return math_fidelity_setting_lookup[self.decoder_optimizations[decoder_id].op_fidelity_settings[op]] def _update_full_name(self): self._full_name = " | ".join( f"Decoder {decoder_id}: {conf._full_name}" for decoder_id, conf in self.decoder_optimizations.items() ) @classmethod def _precision_factory(cls, num_decoders, model_name, optimization_level): # use respective configuration for each optimization level decoder_config_filename = None match optimization_level: case ModelOptimizations.accuracy: decoder_config_filename = ACCURACY_DECODER_CONFIG_FILENAME case ModelOptimizations.performance: decoder_config_filename = PERFORMANCE_DECODER_CONFIG_FILENAME case _: raise ValueError(f"optimization_level ({optimization_level}) not implemented") # check if decoder config exists, if it exists load it else use optimization_level model_params_dir = Path(__file__).parent.parent decoder_config_path = model_params_dir / "model_params" / model_name / decoder_config_filename inst = None if decoder_config_path.exists(): inst = parse_decoder_json(decoder_config_path, default_optimization=optimization_level) logger.info( f"Model {model_name} requires specific TensorPrecision and OpFidelity configuration, using {decoder_config_path}" ) else: inst = cls(num_decoders, model_name, optimization_level(model_name)) return inst def num_to_corerange( x: int, start_core: ttnn.CoreCoord = ttnn.CoreCoord(0, 0), grid_x: int = 8, grid_y: int = 8, ) -> ttnn.CoreRange: """ Construct a rectangular CoreRange of exactly ``x`` cores starting at ``start_core`` on a ``grid_x × grid_y`` core grid. The CoreRange is allocated in row-major order semantics but must form a single contiguous rectangle representable by ``ttnn.CoreRange``. Defaults to an 8×8 grid for backward compatibility. """ # --- basic sanity --- assert x > 0, "x must be positive" assert grid_x > 0 and grid_y > 0 assert 0 <= start_core.x < grid_x assert 0 <= start_core.y < grid_y sx, sy = start_core.x, start_core.y # --- linear availability (row-major correctness) --- remaining_linear_cores = (grid_x - sx) + (grid_y - sy - 1) * grid_x # remainder of start row # full rows below assert remaining_linear_cores >= x, ( f"Not enough cores from start_core {start_core} " f"to allocate {x} cores (only {remaining_linear_cores} available)" ) # --- rectangular availability --- remaining_x = grid_x - sx remaining_y = grid_y - sy # --- shape rule --- assert x < grid_x or x % grid_x == 0, f"x must be < grid_x ({grid_x}) or a multiple of grid_x" # --- choose rectangle dimensions --- num_x = min(x, remaining_x) num_y = x // num_x assert num_x * num_y == x, f"x={x} cannot form a rectangular CoreRange starting at {start_core}" # --- bounds check --- assert num_y <= remaining_y, f"CoreRange height {num_y} exceeds available rows {remaining_y}" end_x = sx + num_x - 1 end_y = sy + num_y - 1 return ttnn.CoreRange( start_core, ttnn.CoreCoord(end_x, end_y), ) def num_to_coregrid(x): if x % 8 == 0: return ttnn.CoreGrid(y=x // 8, x=8) if x == 12: return ttnn.CoreGrid(y=2, x=6) if x == 20: return ttnn.CoreGrid(y=4, x=5) def determine_device_name(mesh_device: ttnn.MeshDevice) -> str: """ Determine device name based on number of devices and architecture. Args: mesh_device (MeshDevice): MeshDevice object Returns: str: Device name (e.g., "CPU", "N150", "P100", etc.) Raises: ValueError: If architecture or device count is unsupported """ num_devices = mesh_device.get_num_devices() dram_grid_size = mesh_device.dram_grid_size() # CoreCoord with (x, y) if ttnn.device.is_blackhole(mesh_device): dict_device_names = { 1: "P100" if dram_grid_size and dram_grid_size.x == 7 else "P150", # P100 DRAM grid is 7x1, P150 is 8x1 2: "P300", 4: "P150x4", 8: "P150x8", 32: "BHGLX", } elif ttnn.device.is_wormhole_b0(mesh_device): dict_device_names = { 1: "N150", 2: "N300", 4: "N150x4", 8: "T3K", 32: "TG", } else: raise ValueError(f"Unsupported architecture: {ttnn.get_arch_name()}") if num_devices in dict_device_names: return dict_device_names[num_devices] else: raise ValueError(f"Unsupported number of devices: {num_devices} for {ttnn.get_arch_name()}")