Download code/models/tt_transformers/tt/model_config.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 256 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/tt_transformers/tt/model_config.py
- Command line
-
hf download hf://tt-hous/clef/code/models/tt_transformers/tt/model_config.py
-
curl -L -o model_config.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/tt_transformers/tt/model_config.py
256 kB
| # 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: | |
| 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 | |
| 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, | |
| }, | |
| } | |
| def tensor_dtype_settings(self): | |
| return self._opt_settings["TensorPrecision"] | |
| 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() | |
| def decoders_optimizations(self): | |
| """Get the decoders optimizations configuration.""" | |
| return self.model_config["DECODERS_OPTIMIZATIONS"] | |
| 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) | |
| 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 | |
| # ========================================================================= | |
| 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 | |
| # ========================================================================= | |
| 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}") | |
| 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 | |
| ), | |
| ) | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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 | |
| 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 | |
| # ========================================================================= | |
| 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, | |
| ) | |
| 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, | |
| ) | |
| 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}") | |
| 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 | |
| # ========================================================================= | |
| 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 | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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}") | |
| 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 | |
| # ========================================================================= | |
| 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 | |
| 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}") | |
| 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, | |
| ) | |
| 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}") | |
| 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 | |
| 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 [] | |
| 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 `<model_name>/snapshots/<snapshot_id>/` | |
| 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() | |
| def use_scaled_rope(self): | |
| return self.rope_scaling is not None | |
| def base_model_name(self): | |
| return get_base_model_name(self.model_name) | |
| 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)) | |
| 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 | |
| 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)) | |
| 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 | |
| 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() | |
| 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 | |
| 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: | |
| 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'" | |
| ) | |
| def accuracy(cls, num_decoders, model_name): | |
| inst = cls._precision_factory(num_decoders, model_name, ModelOptimizations.accuracy) | |
| inst.__name__ = "accuracy" | |
| return inst | |
| 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() | |
| ) | |
| 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()}") | |