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