# SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc. # SPDX-License-Identifier: Apache-2.0 """ Tests for the MLP1D module (1D mesh topology: N150, N300, T3K). This test suite verifies: 1. Unit tests for config dataclasses (no device needed) 2. MLP1D class matches HuggingFace/Meta reference model 3. MLP1D correctly rejects TG/Galaxy devices """ import math import os import time from dataclasses import replace from functools import lru_cache from pathlib import Path import pytest import torch from loguru import logger from transformers import AutoConfig, AutoModelForCausalLM # transformers 5.x moved no_init_weights to transformers.initialization; fall back # to the old location for transformers < 5.x. try: from transformers.initialization import no_init_weights except ImportError: from transformers.modeling_utils import no_init_weights import ttnn from models.common.auto_compose import to_torch_auto_compose from models.common.modules.lazy_weight import LazyWeight from models.common.modules.mlp.mlp_1d import MLP1D, MLP1DConfig, _matmul_config from models.common.tensor_utils import TILE_SIZE from models.common.utility_functions import comp_allclose, comp_pcc from models.tt_transformers.tt.common import Mode # 1D module suites target the T3K; skip when the host system is a Galaxy. pytestmark = pytest.mark.usefixtures("skip_on_galaxy_system") def get_mlp_weights_from_ref_model(reference_mlp) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """ Extract w1, w2, w3 weights from a reference MLP module in TTNN layout (transposed). Handles both standard LLaMA-style MLPs (gate_proj, up_proj, down_proj) and fused gate_up_proj models (Phi-3/Phi-4). Returns: (w1, w2, w3) tensors in TTNN layout: (in_features, out_features) """ if hasattr(reference_mlp, "gate_proj"): w1_torch = reference_mlp.gate_proj.weight.T # (dim, hidden_dim) w3_torch = reference_mlp.up_proj.weight.T # (dim, hidden_dim) elif hasattr(reference_mlp, "gate_up_proj"): # Handle models like Phi-3/Phi-4 that use fused gate_up_proj gate_up_weight = reference_mlp.gate_up_proj.weight hidden_dim = gate_up_weight.shape[0] // 2 w1_torch = gate_up_weight[:hidden_dim, :].T # (dim, hidden_dim) w3_torch = gate_up_weight[hidden_dim:, :].T # (dim, hidden_dim) else: raise AttributeError(f"Reference MLP {type(reference_mlp)} has no gate_proj or gate_up_proj") w2_torch = reference_mlp.down_proj.weight.T # (hidden_dim, dim) return w1_torch, w2_torch, w3_torch def _get_prefill_len_cutoff(hf_model_name: str, mesh_shape: tuple[int, int]) -> int | None: """ Get model/device-specific prefill_len_cutoff override. Root cause: The matmul program config computes per_core_M = ceil(m / (tile_size * grid_height)), where m = min(seq_len, prefill_len_cutoff). Larger per_core_M requires more L1 memory for circular buffers. Combined with in0_block_w=8 and BFP8 weights, certain model/device combinations overflow L1. Symptom: RuntimeError: TT_FATAL ... "Statically allocated circular buffers ... grow to ... beyond max L1 size" Fix: Reduce prefill_len_cutoff from 1024 to 512 for affected models. This halves m, reducing per_core_M (e.g., from 4 to 2), which fits in L1. Matches tt_transformers/tt/model_config.py:577-584 logic: - Llama-3.1-8B, Llama-3.2-11B, Mistral-7B, gemma-3-4b on N150 (1x1) → 512 - Qwen2.5-7B on N300 (1x2) → 512 - Mixtral-8x7B on T3K (1x8) → 512 - Others → None (use default) """ # Extract base model name from HF model name base_name = hf_model_name.split("/")[-1].rsplit("-Instruct", 1)[0] # Map mesh_shape to device type if mesh_shape == (1, 1): device = "N150" elif mesh_shape == (1, 2): device = "N300" elif mesh_shape == (1, 8): device = "T3K" else: device = None # Apply model_config.py logic if base_name in ["Llama-3.1-8B", "Llama-3.2-11B", "Mistral-7B", "gemma-3-4b"] and device == "N150": return 512 elif base_name in ["Qwen2.5-7B"] and device == "N300": return 512 elif base_name in ["Mixtral-8x7B"] and device == "T3K": return 512 return None # Use default # ============================================================================ # Weight Caching - Avoid expensive torch.randn_like() per test # ============================================================================ _CACHED_MLP_WEIGHTS: dict[str, dict[str, torch.Tensor]] = {} def _get_or_init_mlp_weights(model_name: str, reference_mlp) -> None: """Initialize MLP weights once per model, cache and reuse across tests. torch.randn_like() is very slow for large models (18s for 70B). This caches the random weights and reuses them across tests. """ if model_name not in _CACHED_MLP_WEIGHTS: logger.info(f"\033[33m[cache miss]\033[0m Initializing weights for {model_name}") _CACHED_MLP_WEIGHTS[model_name] = {} with torch.no_grad(): for name, param in reference_mlp.named_parameters(): _CACHED_MLP_WEIGHTS[model_name][name] = torch.randn_like(param) else: logger.info(f"\033[32m[cache hit]\033[0m Reusing cached weights for {model_name}") # Load cached weights into model with torch.no_grad(): for name, param in reference_mlp.named_parameters(): param.copy_(_CACHED_MLP_WEIGHTS[model_name][name]) # ============================================================================ # Unit Tests - No device required # ============================================================================ # Note: These tests only test the MLP1DConfig dataclass creation. # The _resolve_mlp1d_config function is tested via integration tests # (test_mlp_1d_vs_reference) since it requires real LazyWeight instances. def test_mlp_1d_config_creation(): """Test that MLP1DConfig dataclass can be created with explicit values.""" from unittest.mock import MagicMock from models.common.modules.mlp.mlp_1d import MLP1DConfig mock_mesh_device = MagicMock() mock_tt_ccl = MagicMock() mock_w1 = MagicMock() mock_w2 = MagicMock() mock_w3 = MagicMock() # Create config with explicit values config = MLP1DConfig( w1=mock_w1, w2=mock_w2, w3=mock_w3, mesh_device=mock_mesh_device, tt_ccl=mock_tt_ccl, dim=4096, hidden_dim=14336, max_batch_size=64, topology=ttnn.Topology.Ring, ) # Verify explicit values are preserved assert config.w1 is mock_w1 assert config.w2 is mock_w2 assert config.w3 is mock_w3 assert config.mesh_device is mock_mesh_device assert config.tt_ccl is mock_tt_ccl assert config.dim == 4096 assert config.hidden_dim == 14336 assert config.max_batch_size == 64 assert config.topology == ttnn.Topology.Ring def test_mlp_1d_config_defaults(): """Test that MLP1DConfig has sensible defaults.""" from unittest.mock import MagicMock from models.common.modules.mlp.mlp_1d import MLP1DConfig # Minimal creation - only required fields config = MLP1DConfig(w1=MagicMock(), w2=MagicMock(), w3=MagicMock()) # Check defaults assert config.max_batch_size == 32 assert config.mlp_activation_type == ttnn.UnaryOpType.SILU # Optional fields default to None assert config.mesh_device is None assert config.tt_ccl is None assert config.dim is None assert config.hidden_dim is None def test_mlp_1d_config_power_user_overrides(): """Test that MLP1DConfig accepts power-user overrides for program configs.""" from unittest.mock import MagicMock from models.common.modules.mlp.mlp_1d import MLP1DConfig mock_prg_config = MagicMock() mock_mem_config = MagicMock() config = MLP1DConfig( w1=MagicMock(), w2=MagicMock(), w3=MagicMock(), decode_w1_w3_prg_config=mock_prg_config, decode_w2_prg_config=mock_prg_config, decode_mlp2_input_memcfg=mock_mem_config, decode_residual_memcfg=mock_mem_config, activation_dtype=ttnn.bfloat16, ) # User-provided overrides should be preserved assert config.decode_w1_w3_prg_config is mock_prg_config assert config.decode_w2_prg_config is mock_prg_config assert config.decode_mlp2_input_memcfg is mock_mem_config assert config.decode_residual_memcfg is mock_mem_config assert config.activation_dtype == ttnn.bfloat16 # Pulled from deduped perf sweep of existing model tests in CI LLAMA_8B = "meta-llama/Llama-3.1-8B-Instruct" LLAMA_70B = "meta-llama/Llama-3.3-70B-Instruct" LLAMA_1B = "meta-llama/Llama-3.2-1B-Instruct" LLAMA_3B = "meta-llama/Llama-3.2-3B-Instruct" LLAMA_11B = "meta-llama/Llama-3.2-11B-Vision-Instruct" LLAMA_90B = "meta-llama/Llama-3.2-90B-Vision-Instruct" MISTRAL_7B = "mistralai/Mistral-7B-Instruct-v0.3" QWEN2_7B = "Qwen/Qwen2-7B-Instruct" QWEN25_7B = "Qwen/Qwen2.5-7B-Instruct" QWEN25_72B = "Qwen/Qwen2.5-72B-Instruct" QWEN25_CODER_32B = "Qwen/Qwen2.5-Coder-32B-Instruct" DEEPSEEK_R1_14B = "deepseek-ai/DeepSeek-R1-Distill-Qwen-14B" PHI_4 = "microsoft/phi-4" QWEN3_32B = "Qwen/Qwen3-32B" _slow = pytest.mark.slow # [INFO] Galaxy DP run multiple copies of the following on 1x1, 1x2, and 1x8 meshes. def _list_glx_test_cases() -> list[pytest.param]: # fmt: off return [ # === Fast tests (minimal coverage set) === # Single device pytest.param((1, 1), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x1-prefill-128-mixed-8B"), pytest.param((1, 1), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x1-decode-32-uniform-8B"), # Multi-device (1x8) pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x8-prefill-128-mixed-8B"), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x8-prefill-1024-mixed-8B"), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x8-decode-32-mixed-8B"), # 70B (larger dims) pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_70B, 0.98, id="1x8-prefill-1024-mixed-70B"), # === Slow tests (full coverage from models sweep) === # (1,1) mesh - from DP-32 (8B) pytest.param((1, 1), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x1-prefill-128-uniform-8B", marks=_slow), pytest.param((1, 1), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x1-prefill-256-mixed-8B", marks=_slow), pytest.param((1, 1), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x1-prefill-512-mixed-8B", marks=_slow), pytest.param((1, 1), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x1-prefill-1024-mixed-8B", marks=_slow), pytest.param((1, 1), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x1-decode-32-mixed-8B", marks=_slow), # (1,2) mesh - from DP-16 (8B) pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x2-prefill-128-mixed-8B", marks=_slow), pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x2-prefill-128-uniform-8B", marks=_slow), pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x2-prefill-256-mixed-8B", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x2-prefill-512-mixed-8B", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x2-prefill-1024-mixed-8B", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x2-decode-32-mixed-8B", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x2-decode-32-uniform-8B", marks=_slow), # (1,8) mesh - from DP-4 (8B) pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x8-prefill-128-uniform-8B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x8-prefill-256-mixed-8B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x8-prefill-512-mixed-8B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x8-decode-32-uniform-8B", marks=_slow), # (1,8) mesh - from DP-4_70B (70B, mixed dtype only) pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_70B, 0.98, id="1x8-prefill-128-mixed-70B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_70B, 0.98, id="1x8-prefill-256-mixed-70B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_70B, 0.98, id="1x8-prefill-512-mixed-70B", marks=_slow), pytest.param((1, 8), 1, 2048, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_70B, 0.98, id="1x8-prefill-2048-mixed-70B", marks=_slow), pytest.param((1, 8), 1, 4096, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_70B, 0.98, id="1x8-prefill-4096-mixed-70B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_70B, 0.98, id="1x8-decode-32-mixed-70B", marks=_slow), ] # fmt: on # [INFO] Non-Galaxy test cases from N150/N300/T3K/BH runs. def _list_non_glx_test_cases() -> list[pytest.param]: # fmt: off return [ # === Fast tests (minimal coverage set) === # Single device (1x1) - small model pytest.param((1, 1), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x1-prefill-128-mixed-1B"), pytest.param((1, 1), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x1-decode-32-uniform-1B"), # Multi-device (1x2) - vision model 11B pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_11B, 0.99, id="1x2-prefill-128-uniform-11B"), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_11B, 0.99, id="1x2-decode-32-uniform-11B"), # Multi-device (1x8) - vision model 90B pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_90B, 0.98, id="1x8-prefill-128-mixed-90B"), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_90B, 0.98, id="1x8-decode-32-mixed-90B"), # Non-Llama model families pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, PHI_4, 0.99, id="1x2-prefill-128-uniform-phi-4"), pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN3_32B, 0.98, id="1x8-prefill-128-mixed-Qwen3-32B"), pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN25_7B, 0.99, id="1x2-prefill-128-uniform-Qwen2.5-7B"), pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, DEEPSEEK_R1_14B, 0.98, id="1x2-prefill-128-mixed-DeepSeek-R1-14B"), # === Slow tests (full coverage) === # (1,1) LLAMA_1B - remaining cases not in fast set pytest.param((1, 1), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x1-prefill-256-mixed-1B", marks=_slow), pytest.param((1, 1), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x1-prefill-512-mixed-1B", marks=_slow), pytest.param((1, 1), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x1-decode-32-mixed-1B", marks=_slow), pytest.param((1, 1), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x1-prefill-1024-mixed-1B", marks=_slow), pytest.param((1, 1), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x1-prefill-128-uniform-1B", marks=_slow), pytest.param((1, 1), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x1-prefill-256-uniform-1B", marks=_slow), pytest.param((1, 1), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x1-prefill-512-uniform-1B", marks=_slow), pytest.param((1, 1), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x1-prefill-128-mixed-3B", marks=_slow), pytest.param((1, 1), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x1-prefill-256-mixed-3B", marks=_slow), pytest.param((1, 1), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x1-prefill-512-mixed-3B", marks=_slow), pytest.param((1, 1), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x1-decode-32-mixed-3B", marks=_slow), pytest.param((1, 1), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x1-prefill-1024-mixed-3B", marks=_slow), pytest.param((1, 1), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x1-prefill-128-uniform-3B", marks=_slow), pytest.param((1, 1), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x1-prefill-256-uniform-3B", marks=_slow), pytest.param((1, 1), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x1-prefill-512-uniform-3B", marks=_slow), pytest.param((1, 1), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x1-decode-32-uniform-3B", marks=_slow), pytest.param((1, 1), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x1-prefill-128-mixed-8B", marks=_slow), pytest.param((1, 1), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x1-prefill-128-uniform-8B", marks=_slow), pytest.param((1, 1), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x1-prefill-256-mixed-8B", marks=_slow), pytest.param((1, 1), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x1-prefill-256-uniform-8B", marks=_slow), pytest.param((1, 1), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x1-prefill-512-mixed-8B", marks=_slow), pytest.param((1, 1), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x1-prefill-512-uniform-8B", marks=_slow), pytest.param((1, 1), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x1-prefill-1024-mixed-8B", marks=_slow), pytest.param((1, 1), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x1-prefill-1024-uniform-8B", marks=_slow), pytest.param((1, 1), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x1-decode-32-mixed-8B", marks=_slow), pytest.param((1, 1), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x1-decode-32-uniform-8B", marks=_slow), pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x2-prefill-128-mixed-1B", marks=_slow), pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x2-prefill-256-mixed-1B", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x2-prefill-512-mixed-1B", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x2-prefill-1024-mixed-1B", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x2-decode-32-mixed-1B", marks=_slow), pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x2-prefill-128-uniform-1B", marks=_slow), pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x2-prefill-256-uniform-1B", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x2-prefill-512-uniform-1B", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x2-prefill-1024-uniform-1B", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x2-decode-32-uniform-1B", marks=_slow), pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x2-prefill-128-mixed-3B", marks=_slow), pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x2-prefill-256-mixed-3B", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x2-prefill-512-mixed-3B", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x2-prefill-1024-mixed-3B", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x2-decode-32-mixed-3B", marks=_slow), pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x2-prefill-128-uniform-3B", marks=_slow), pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x2-prefill-256-uniform-3B", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x2-prefill-512-uniform-3B", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x2-prefill-1024-uniform-3B", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x2-decode-32-uniform-3B", marks=_slow), pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x2-prefill-128-mixed-8B", marks=_slow), pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x2-prefill-128-uniform-8B", marks=_slow), pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x2-prefill-256-mixed-8B", marks=_slow), pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x2-prefill-256-uniform-8B", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x2-prefill-512-mixed-8B", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x2-prefill-512-uniform-8B", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x2-prefill-1024-mixed-8B", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x2-prefill-1024-uniform-8B", marks=_slow), pytest.param((1, 2), 1, 2048, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x2-prefill-2048-mixed-8B", marks=_slow), pytest.param((1, 2), 1, 2048, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x2-prefill-2048-uniform-8B", marks=_slow), pytest.param((1, 2), 1, 4096, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x2-prefill-4096-mixed-8B", marks=_slow), pytest.param((1, 2), 1, 4096, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x2-prefill-4096-uniform-8B", marks=_slow), pytest.param((1, 2), 1, 8192, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x2-prefill-8192-mixed-8B", marks=_slow), pytest.param((1, 2), 1, 8192, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x2-prefill-8192-uniform-8B", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x2-decode-32-mixed-8B", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x2-decode-32-uniform-8B", marks=_slow), pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_11B, 0.98, id="1x2-prefill-128-mixed-11B", marks=_slow), pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_11B, 0.98, id="1x2-prefill-256-mixed-11B", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_11B, 0.98, id="1x2-prefill-512-mixed-11B", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_11B, 0.98, id="1x2-prefill-1024-mixed-11B", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_11B, 0.98, id="1x2-decode-32-mixed-11B", marks=_slow), pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_11B, 0.99, id="1x2-prefill-256-uniform-11B", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_11B, 0.99, id="1x2-prefill-512-uniform-11B", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_11B, 0.99, id="1x2-prefill-1024-uniform-11B", marks=_slow), pytest.param((1, 1), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x1-prefill-128-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 1), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x1-prefill-256-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 1), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x1-prefill-512-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 1), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x1-decode-32-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 1), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x1-prefill-1024-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 1), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x1-prefill-128-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 1), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x1-prefill-256-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 1), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x1-prefill-512-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 1), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x1-decode-32-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x2-prefill-128-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x2-prefill-256-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x2-prefill-512-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x2-prefill-1024-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x2-decode-32-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x2-prefill-128-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x2-prefill-256-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x2-prefill-512-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x2-prefill-1024-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x2-decode-32-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 2), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN2_7B, 0.98, id="1x2-prefill-128-mixed-Qwen2-7B-Instruct", marks=_slow), pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN2_7B, 0.98, id="1x2-prefill-256-mixed-Qwen2-7B-Instruct", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN2_7B, 0.98, id="1x2-prefill-512-mixed-Qwen2-7B-Instruct", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN2_7B, 0.98, id="1x2-prefill-1024-mixed-Qwen2-7B-Instruct", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN2_7B, 0.98, id="1x2-decode-32-mixed-Qwen2-7B-Instruct", marks=_slow), pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, PHI_4, 0.99, id="1x2-prefill-256-uniform-phi-4", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, PHI_4, 0.99, id="1x2-prefill-512-uniform-phi-4", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, PHI_4, 0.99, id="1x2-prefill-1024-uniform-phi-4", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, PHI_4, 0.99, id="1x2-decode-32-uniform-phi-4", marks=_slow), pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x8-prefill-128-mixed-1B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x8-prefill-256-mixed-1B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x8-prefill-512-mixed-1B", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x8-prefill-1024-mixed-1B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_1B, 0.98, id="1x8-decode-32-mixed-1B", marks=_slow), pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x8-prefill-128-uniform-1B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x8-prefill-256-uniform-1B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x8-prefill-512-uniform-1B", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x8-prefill-1024-uniform-1B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_1B, 0.99, id="1x8-decode-32-uniform-1B", marks=_slow), pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x8-prefill-128-mixed-3B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x8-prefill-256-mixed-3B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x8-prefill-512-mixed-3B", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x8-prefill-1024-mixed-3B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_3B, 0.98, id="1x8-decode-32-mixed-3B", marks=_slow), pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x8-prefill-128-uniform-3B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x8-prefill-256-uniform-3B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x8-prefill-512-uniform-3B", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x8-prefill-1024-uniform-3B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_3B, 0.99, id="1x8-decode-32-uniform-3B", marks=_slow), pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x8-prefill-128-mixed-8B", marks=_slow), pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x8-prefill-128-uniform-8B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x8-prefill-256-mixed-8B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x8-prefill-256-uniform-8B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x8-prefill-512-mixed-8B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x8-prefill-512-uniform-8B", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x8-prefill-1024-mixed-8B", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x8-prefill-1024-uniform-8B", marks=_slow), pytest.param((1, 8), 1, 2048, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x8-prefill-2048-mixed-8B", marks=_slow), pytest.param((1, 8), 1, 2048, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x8-prefill-2048-uniform-8B", marks=_slow), pytest.param((1, 8), 1, 4096, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x8-prefill-4096-mixed-8B", marks=_slow), pytest.param((1, 8), 1, 4096, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x8-prefill-4096-uniform-8B", marks=_slow), pytest.param((1, 8), 1, 8192, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x8-prefill-8192-mixed-8B", marks=_slow), pytest.param((1, 8), 1, 8192, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x8-prefill-8192-uniform-8B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_8B, 0.98, id="1x8-decode-32-mixed-8B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_8B, 0.99, id="1x8-decode-32-uniform-8B", marks=_slow), pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_11B, 0.98, id="1x8-prefill-128-mixed-11B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_11B, 0.98, id="1x8-prefill-256-mixed-11B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_11B, 0.98, id="1x8-prefill-512-mixed-11B", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_11B, 0.98, id="1x8-prefill-1024-mixed-11B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, LLAMA_11B, 0.98, id="1x8-decode-32-mixed-11B", marks=_slow), pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_11B, 0.99, id="1x8-prefill-128-uniform-11B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_11B, 0.99, id="1x8-prefill-256-uniform-11B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_11B, 0.99, id="1x8-prefill-512-uniform-11B", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_11B, 0.99, id="1x8-prefill-1024-uniform-11B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, LLAMA_11B, 0.99, id="1x8-decode-32-uniform-11B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN3_32B, 0.98, id="1x8-prefill-256-mixed-Qwen3-32B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN3_32B, 0.98, id="1x8-prefill-512-mixed-Qwen3-32B", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN3_32B, 0.98, id="1x8-prefill-1024-mixed-Qwen3-32B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN3_32B, 0.98, id="1x8-decode-32-mixed-Qwen3-32B", marks=_slow), pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN3_32B, 0.99, id="1x8-prefill-128-uniform-Qwen3-32B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN3_32B, 0.99, id="1x8-prefill-256-uniform-Qwen3-32B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN3_32B, 0.99, id="1x8-prefill-512-uniform-Qwen3-32B", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN3_32B, 0.99, id="1x8-prefill-1024-uniform-Qwen3-32B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN3_32B, 0.99, id="1x8-decode-32-uniform-Qwen3-32B", marks=_slow), pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x8-prefill-128-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x8-prefill-256-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x8-prefill-512-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x8-prefill-1024-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, MISTRAL_7B, 0.98, id="1x8-decode-32-mixed-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x8-prefill-128-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x8-prefill-256-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x8-prefill-512-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x8-prefill-1024-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, MISTRAL_7B, 0.99, id="1x8-decode-32-uniform-Mistral-7B-Instruct-v0.3", marks=_slow), # --- New test cases from mlp_1d_performance.csv --- # Qwen2.5-7B on N300 (1x2) - uniform BF8 pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN25_7B, 0.99, id="1x2-prefill-256-uniform-Qwen2.5-7B", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN25_7B, 0.99, id="1x2-prefill-512-uniform-Qwen2.5-7B", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN25_7B, 0.99, id="1x2-prefill-1024-uniform-Qwen2.5-7B", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN25_7B, 0.99, id="1x2-decode-32-uniform-Qwen2.5-7B", marks=_slow), # Qwen2.5-72B on T3K (1x8) - mixed BF4/BF8 pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN25_72B, 0.98, id="1x8-prefill-128-mixed-Qwen2.5-72B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN25_72B, 0.98, id="1x8-prefill-256-mixed-Qwen2.5-72B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN25_72B, 0.98, id="1x8-prefill-512-mixed-Qwen2.5-72B", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN25_72B, 0.98, id="1x8-prefill-1024-mixed-Qwen2.5-72B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN25_72B, 0.98, id="1x8-decode-32-mixed-Qwen2.5-72B", marks=_slow), # DeepSeek-R1-Distill-Qwen-14B on N300 (1x2) - mixed BF4/BF8 pytest.param((1, 2), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, DEEPSEEK_R1_14B, 0.98, id="1x2-prefill-256-mixed-DeepSeek-R1-14B", marks=_slow), pytest.param((1, 2), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, DEEPSEEK_R1_14B, 0.98, id="1x2-prefill-512-mixed-DeepSeek-R1-14B", marks=_slow), pytest.param((1, 2), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, DEEPSEEK_R1_14B, 0.98, id="1x2-prefill-1024-mixed-DeepSeek-R1-14B", marks=_slow), pytest.param((1, 2), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, DEEPSEEK_R1_14B, 0.98, id="1x2-decode-32-mixed-DeepSeek-R1-14B", marks=_slow), # Qwen2.5-Coder-32B on T3K (1x8) - mixed BF4/BF8 pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN25_CODER_32B, 0.98, id="1x8-prefill-128-mixed-Qwen2.5-Coder-32B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN25_CODER_32B, 0.98, id="1x8-prefill-256-mixed-Qwen2.5-Coder-32B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN25_CODER_32B, 0.98, id="1x8-prefill-512-mixed-Qwen2.5-Coder-32B", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN25_CODER_32B, 0.98, id="1x8-prefill-1024-mixed-Qwen2.5-Coder-32B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat4_b, ttnn.bfloat8_b, ttnn.bfloat4_b, QWEN25_CODER_32B, 0.98, id="1x8-decode-32-mixed-Qwen2.5-Coder-32B", marks=_slow), # Qwen2.5-Coder-32B on T3K (1x8) - uniform BF8 pytest.param((1, 8), 1, 128, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN25_CODER_32B, 0.99, id="1x8-prefill-128-uniform-Qwen2.5-Coder-32B", marks=_slow), pytest.param((1, 8), 1, 256, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN25_CODER_32B, 0.99, id="1x8-prefill-256-uniform-Qwen2.5-Coder-32B", marks=_slow), pytest.param((1, 8), 1, 512, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN25_CODER_32B, 0.99, id="1x8-prefill-512-uniform-Qwen2.5-Coder-32B", marks=_slow), pytest.param((1, 8), 1, 1024, "prefill", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN25_CODER_32B, 0.99, id="1x8-prefill-1024-uniform-Qwen2.5-Coder-32B", marks=_slow), pytest.param((1, 8), 1, 32, "decode", ttnn.bfloat16, ttnn.bfloat8_b, ttnn.bfloat8_b, ttnn.bfloat8_b, QWEN25_CODER_32B, 0.99, id="1x8-decode-32-uniform-Qwen2.5-Coder-32B", marks=_slow), ] # [INFO] generate random tensor for every test case is too expensive; cache weights and reuse them across test cases # [INFO] separate out ttnn_mesh_device parameter allows for sharing the same mesh device across test cases @pytest.mark.parametrize( "ttnn_mesh_device", [(1, 1), (1, 2), (1, 8)], ids=["1x1", "1x2", "1x8"], indirect=True, ) @pytest.mark.parametrize( "mesh_shape,batch_size,seq_len,mode,act_dtype,w1_dtype,w2_dtype,w3_dtype,hf_model_name,pcc", _list_non_glx_test_cases() + _list_glx_test_cases(), ) def test_mlp_1d_vs_reference( ttnn_mesh_device: ttnn.MeshDevice, mesh_shape, batch_size, seq_len, mode, act_dtype, w1_dtype, w2_dtype, w3_dtype, hf_model_name, pcc, ): """ Test MLP1D constructed via direct APIs (MLP1DConfig) matches HF reference MLP. Configs pulled from perf sweep CSVs (b{batch_size}-DP-{dp}_{model}). """ # get reference model; generate and load deterministic, random weights into the reference model seed = 1234 torch.manual_seed(seed) # HF model (default small) for reference; skip global init to only seed MLP. config = AutoConfig.from_pretrained(hf_model_name) config.num_hidden_layers = 1 with no_init_weights(): hf_model = AutoModelForCausalLM.from_config(config, torch_dtype=torch.bfloat16) first_layer = hf_model.model.layers[0] reference_mlp = first_layer.mlp # Initialize only the MLP submodule deterministically (cached for speed). _get_or_init_mlp_weights(hf_model_name, reference_mlp) # Build MLP1D TT model and load the same weights as the reference model w1_torch, w2_torch, w3_torch = get_mlp_weights_from_ref_model(reference_mlp) dim = w1_torch.shape[0] torch_input = torch.randn(batch_size, 1, seq_len, dim, dtype=torch.bfloat16) # Create LazyWeights ttnn.SetDefaultDevice(ttnn_mesh_device) cache_dir = Path(os.getenv("TT_CACHE_PATH", "model_cache/mlp_1d")) lazy_w1 = LazyWeight(source=w1_torch, dtype=w1_dtype, cache_dir_weight_name=(cache_dir, "w1")) lazy_w2 = LazyWeight(source=w2_torch, dtype=w2_dtype, cache_dir_weight_name=(cache_dir, "w2")) lazy_w3 = LazyWeight(source=w3_torch, dtype=w3_dtype, cache_dir_weight_name=(cache_dir, "w3")) # Get model/device-specific prefill_len_cutoff prefill_len_cutoff = _get_prefill_len_cutoff(hf_model_name, mesh_shape) # Construct the MLP1D model if prefill_len_cutoff is None: # Use default prefill_len_cutoff, take the happy path of MLP1D tt_model = MLP1D(w1=lazy_w1, w2=lazy_w2, w3=lazy_w3) else: tt_model = MLP1D.from_config( MLP1DConfig( w1=lazy_w1, w2=lazy_w2, w3=lazy_w3, prefill_len_cutoff=prefill_len_cutoff, ) ) # Run TT model with the same input -- torch_input -- converted to ttnn tensor lazily on the fly # [INFO] we use LazyWeight on input for the benefit of faster testing (cached input); in production, the input is already a ttnn tensor. tt_input = LazyWeight( source=torch_input, dtype=act_dtype, # cache_dir_weight_name=(cache_dir, "input"), # todo)) needs better fingerprinting for input tensor to enable ) tt_output = tt_model.forward(tt_input, mode) tt_output_torch = to_torch_auto_compose(tt_output) ttnn.SetDefaultDevice(None) # Now both models are ready to go # Run reference model with torch.no_grad(): reference_output = reference_mlp(torch_input) passing, pcc_message = comp_pcc(reference_output, tt_output_torch, pcc) logger.info(comp_allclose(reference_output, tt_output_torch)) logger.info(f"MLP1D (direct API) vs HF reference: {pcc_message}") assert passing, f"MLP1D output does not meet PCC requirement {pcc}: {pcc_message}." logger.info(f"MLP1D (direct API) vs HF reference: PASSED for mode={mode}, seq_len={seq_len}") def _blackhole_mlp_kernel(*, fidelity: ttnn.MathFidelity) -> ttnn.DeviceComputeKernelConfig: """Materialize one TTTv1 Qwen3-32B performance MLP operation slot on BH.""" return ttnn.init_device_compute_kernel_config( ttnn.device.Arch.BLACKHOLE, math_fidelity=fidelity, math_approx_mode=False, fp32_dest_acc_en=False, packer_l1_acc=True, ) @pytest.mark.parametrize( "ttnn_mesh_device,mode,input_rows", [ pytest.param((1, 1), "decode", 32, id="p150-1x1-decode-batch32"), pytest.param((1, 1), "prefill", 512, id="p150-1x1-prefill-seq512-minimal-ff2"), pytest.param( {"mesh_shape": (1, 4), "fabric_config": ttnn.FabricConfig.FABRIC_1D_RING}, "decode", 32, id="p150x4-1x4-ring-decode-batch32", ), pytest.param( {"mesh_shape": (1, 4), "fabric_config": ttnn.FabricConfig.FABRIC_1D_RING}, "prefill", 512, id="p150x4-1x4-ring-prefill-seq512-minimal-ff2", ), ], indirect=["ttnn_mesh_device"], ) def test_mlp_1d_blackhole_common_config_correctness_cache_and_timing( request, ttnn_mesh_device, require_blackhole_mesh_device, mode, input_rows ): """Focused BH correctness/cache gate; synchronized timings are evidence, not a threshold.""" torch.manual_seed(2026) # A reduced Qwen-shaped 1:5 MLP keeps the P150x4 production 8x5 grids legal # while avoiding three full 32B-model weight allocations in this module gate. dim = 1280 hidden_dim = 6400 num_devices = ttnn_mesh_device.get_num_devices() assert num_devices in (1, 4) assert ttnn_mesh_device.dram_grid_size().x == 8, "This gate requires a P150 DRAM grid" w1 = torch.randn(dim, hidden_dim, dtype=torch.bfloat16) * 0.02 w2 = torch.randn(hidden_dim, dim, dtype=torch.bfloat16) * 0.02 w3 = torch.randn(dim, hidden_dim, dtype=torch.bfloat16) * 0.02 torch_input = torch.randn(1, 1, input_rows, dim, dtype=torch.bfloat16) with torch.no_grad(): reference = (torch.nn.functional.silu(torch_input @ w1) * (torch_input @ w3)) @ w2 common = MLP1DConfig( w1=LazyWeight(source=w1, dtype=ttnn.bfloat8_b), w2=LazyWeight(source=w2, dtype=ttnn.bfloat8_b), w3=LazyWeight(source=w3, dtype=ttnn.bfloat8_b), mesh_device=ttnn_mesh_device, dim=dim, hidden_dim=hidden_dim, max_batch_size=32, topology=ttnn.Topology.Ring if num_devices == 4 else None, prefill_w2_minimal_matmul=True, ) # TTTv1 Qwen3-32B performance recipe: FF1/FF3 use LoFi while FF2 uses # HiFi2 FP16 accumulation. Spell out all four operation slots so the gate # proves mode-specific wrapper routing rather than relying on defaults. common = replace( common, ff1_3_compute_kernel_cfg=_blackhole_mlp_kernel(fidelity=ttnn.MathFidelity.LoFi), ff2_compute_kernel_cfg=_blackhole_mlp_kernel(fidelity=ttnn.MathFidelity.HiFi2), decode_ff1_3_compute_kernel_cfg=_blackhole_mlp_kernel(fidelity=ttnn.MathFidelity.LoFi), decode_ff2_compute_kernel_cfg=_blackhole_mlp_kernel(fidelity=ttnn.MathFidelity.HiFi2), prefill_len_cutoff=512, prefill_dram_shard_grid_width=8, prefill_ff1_ff3_grid=(8, 5), prefill_ff2_grid=(8, 5), ) model = MLP1D.from_config(common) assert not hasattr(model, "arch_config") assert model.config.prefill_len_cutoff == 512 assert model.config.prefill_ff1_ff3_grid == (8, 5) assert model.config.prefill_ff2_grid == (8, 5) assert ( model.config.ff1_3_compute_kernel_cfg.math_fidelity, model.config.ff2_compute_kernel_cfg.math_fidelity, model.config.decode_ff1_3_compute_kernel_cfg.math_fidelity, model.config.decode_ff2_compute_kernel_cfg.math_fidelity, ) == ( ttnn.MathFidelity.LoFi, ttnn.MathFidelity.HiFi2, ttnn.MathFidelity.LoFi, ttnn.MathFidelity.HiFi2, ) assert model.config.use_minimal_w2_matmul(input_rows) is (mode == "prefill") ttnn_mesh_device.enable_program_cache() ttnn_mesh_device.clear_program_cache() request.addfinalizer(ttnn_mesh_device.disable_and_clear_program_cache) def run_once(): # MLP1D consumes and deallocates its device input. A fresh LazyWeight # prevents a warm replay from reusing a deallocated device tensor. fresh_input = LazyWeight(source=torch_input, dtype=ttnn.bfloat16) output = model.forward(fresh_input, mode=mode) ttnn.synchronize_device(ttnn_mesh_device) return output output = run_once() actual = to_torch_auto_compose(output) output.deallocate(True) passing, pcc_message = comp_pcc(reference, actual, 0.97) assert passing, f"Blackhole MLP1D PCC failed: {pcc_message}" cache_entries = ttnn_mesh_device.num_program_cache_entries() assert cache_entries > 0 timings_ms = [] for _ in range(3): start = time.perf_counter() output = run_once() timings_ms.append((time.perf_counter() - start) * 1000) assert ttnn_mesh_device.num_program_cache_entries() == cache_entries output.deallocate(True) logger.info( "BH MLP1D measurement mode={} mesh={} dim={} hidden_dim={}: warm-cache mean={:.3f} ms, samples={}", mode, tuple(ttnn_mesh_device.shape), dim, hidden_dim, sum(timings_ms) / len(timings_ms), timings_ms, ) @pytest.mark.parametrize( "ttnn_mesh_device,mode,input_rows", [ pytest.param((1, 1), "prefill", 512, id="n150-1x1-prefill-seq512-minimal-ff2"), pytest.param((1, 1), "decode", 32, id="n150-1x1-decode-batch32"), ], indirect=["ttnn_mesh_device"], ) def test_mlp_1d_wormhole_common_config_correctness_cache_and_timing(request, ttnn_mesh_device, mode, input_rows): """Focused WH correctness/cache gate using explicit common-config requests.""" torch.manual_seed(2026) dim = 1280 hidden_dim = 6400 w1 = torch.randn(dim, hidden_dim, dtype=torch.bfloat16) * 0.02 w2 = torch.randn(hidden_dim, dim, dtype=torch.bfloat16) * 0.02 w3 = torch.randn(dim, hidden_dim, dtype=torch.bfloat16) * 0.02 torch_input = torch.randn(1, 1, input_rows, dim, dtype=torch.bfloat16) with torch.no_grad(): reference = (torch.nn.functional.silu(torch_input @ w1) * (torch_input @ w3)) @ w2 common = MLP1DConfig( w1=LazyWeight(source=w1, dtype=ttnn.bfloat8_b), w2=LazyWeight(source=w2, dtype=ttnn.bfloat8_b), w3=LazyWeight(source=w3, dtype=ttnn.bfloat8_b), mesh_device=ttnn_mesh_device, dim=dim, hidden_dim=hidden_dim, max_batch_size=32, topology=None, prefill_w2_minimal_matmul=True, ) kernel = lambda: ttnn.init_device_compute_kernel_config( ttnn.device.Arch.WORMHOLE_B0, math_fidelity=ttnn.MathFidelity.HiFi2, math_approx_mode=False, fp32_dest_acc_en=False, packer_l1_acc=True, ) common = replace( common, ff1_3_compute_kernel_cfg=kernel(), ff2_compute_kernel_cfg=kernel(), decode_ff1_3_compute_kernel_cfg=kernel(), decode_ff2_compute_kernel_cfg=kernel(), prefill_len_cutoff=1024, prefill_dram_shard_grid_width=8, prefill_ff1_ff3_grid=(8, 5), prefill_ff2_grid=(8, 5), ) model = MLP1D.from_config(common) assert not hasattr(model, "arch_config") assert model.config.prefill_len_cutoff == 1024 assert ( len( { id(getattr(model.config, name)) for name in ( "ff1_3_compute_kernel_cfg", "ff2_compute_kernel_cfg", "decode_ff1_3_compute_kernel_cfg", "decode_ff2_compute_kernel_cfg", ) } ) == 4 ) assert model.config.use_minimal_w2_matmul(input_rows) is (mode == "prefill") ttnn_mesh_device.enable_program_cache() ttnn_mesh_device.clear_program_cache() request.addfinalizer(ttnn_mesh_device.disable_and_clear_program_cache) def run_once(): fresh_input = LazyWeight(source=torch_input, dtype=ttnn.bfloat16) output = model.forward(fresh_input, mode=mode) ttnn.synchronize_device(ttnn_mesh_device) return output output = run_once() actual = to_torch_auto_compose(output) output.deallocate(True) passing, pcc_message = comp_pcc(reference, actual, 0.97) assert passing, f"Wormhole MLP1D PCC failed: {pcc_message}" cache_entries = ttnn_mesh_device.num_program_cache_entries() assert cache_entries > 0 timings_ms = [] for _ in range(3): start = time.perf_counter() output = run_once() timings_ms.append((time.perf_counter() - start) * 1000) assert ttnn_mesh_device.num_program_cache_entries() == cache_entries output.deallocate(True) logger.info( "WH MLP1D measurement mode={} mesh={} dim={} hidden_dim={}: warm-cache mean={:.3f} ms, samples={}", mode, tuple(ttnn_mesh_device.shape), dim, hidden_dim, sum(timings_ms) / len(timings_ms), timings_ms, ) @pytest.mark.parametrize( "ttnn_mesh_device", [(1, 8)], ids=["1x8"], indirect=True, ) def test_mlp_1d_config_prefill_override(ttnn_mesh_device: ttnn.MeshDevice): """ Show how to override prefill_w2_prg_config with the MLP1DConfig API. Use MLP1D.from_config() for any customization beyond the simple 3-weight API. """ from models.common.modules.mlp.mlp_1d import _find_prefill_grid # Use Llama 8B config hf_model_name = "meta-llama/Llama-3.1-8B-Instruct" hf_config = AutoConfig.from_pretrained(hf_model_name) seq_len = 128 batch_size = 1 # Load HF model for reference weights hf_config.num_hidden_layers = 1 with no_init_weights(): hf_model = AutoModelForCausalLM.from_config(hf_config, torch_dtype=torch.bfloat16) reference_mlp = hf_model.model.layers[0].mlp # Generate random weights directly for this test # todo)) using _get_or_init_mlp_weights instead here would interfere with the test_mlp_1d_vs_reference test; this problem could be solved by provenance-based fingerprinting the torch.tensor inputs with torch.no_grad(): for param in reference_mlp.parameters(): param.copy_(torch.randn_like(param)) # Prepare weights w1_torch, w2_torch, w3_torch = get_mlp_weights_from_ref_model(reference_mlp) # Create LazyWeights (no disk cache) ttnn.SetDefaultDevice(ttnn_mesh_device) lazy_w1 = LazyWeight(source=w1_torch, dtype=ttnn.bfloat4_b) lazy_w2 = LazyWeight(source=w2_torch, dtype=ttnn.bfloat8_b) lazy_w3 = LazyWeight(source=w3_torch, dtype=ttnn.bfloat4_b) # Step 1: Create MLP1D with default config tt_model = MLP1D.from_config(MLP1DConfig(w1=lazy_w1, w2=lazy_w2, w3=lazy_w3)) # Step 2: Define custom prefill w2 config using resolved values from tt_model.config cfg = tt_model.config dim = cfg.dim hidden_dim = cfg.hidden_dim tile_size = TILE_SIZE prefill_len_cutoff = tt_model.config.prefill_len_cutoff @lru_cache def custom_prefill_w2_prg_config(seq_len: int): n_w2 = dim dram_shard_grid_width = 8 prefill_rows = 8 grid_size = _find_prefill_grid(prefill_rows, hidden_dim // tile_size) return _matmul_config( m=min(seq_len, prefill_len_cutoff), k=hidden_dim, n=n_w2, grid_size=grid_size, per_core_n=math.ceil(n_w2 / (tile_size * dram_shard_grid_width)), ) # Step 3: Override the prefill config on the existing model tt_model.config.prefill_w2_prg_config = custom_prefill_w2_prg_config # Verify the override was applied assert tt_model.config.prefill_w2_prg_config is custom_prefill_w2_prg_config # Run prefill forward torch_input = torch.randn(batch_size, 1, seq_len, dim, dtype=torch.bfloat16) tt_input = LazyWeight(source=torch_input, dtype=ttnn.bfloat8_b) tt_output = tt_model.forward(tt_input, mode="prefill") tt_output_torch = to_torch_auto_compose(tt_output) ttnn.SetDefaultDevice(None) # Verify output shape matches input shape (MLP is dim -> dim) assert tt_output_torch.shape == torch_input.shape, f"Expected {torch_input.shape}, got {tt_output_torch.shape}" # Verify numerical correctness against reference with torch.no_grad(): reference_output = reference_mlp(torch_input) passing, pcc_message = comp_pcc(reference_output, tt_output_torch, 0.98) assert passing, f"MLP1D with custom prefill config failed PCC: {pcc_message}" logger.info(f"test_mlp_1d_config_prefill_override: PASSED - {pcc_message}") # ============================================================================ # Integration Tests - Require device # ============================================================================ # [INFO] this test will retire once models/tt_transformers/tt/model_config.py retires @pytest.mark.parametrize( "ttnn_mesh_device", [ (1, 1), # single device (1, 2), # 1D mesh, 2 devices (1, 4), # 1D mesh, 4 devices (1, 8), # 1D mesh, 8 devices ], ids=["1x1", "1x2", "1x4", "1x8"], indirect=True, ) @pytest.mark.parametrize("seq_len", (512, 32)) def test_mlp_1d_vs_reference_from_model_args(ttnn_mesh_device: ttnn.MeshDevice, seq_len): """ Test that MLP1D class matches the HuggingFace/Meta reference model. """ from models.common.modules.mlp.mlp_1d import MLP1D from models.tt_transformers.tests.test_utils import get_ref_model_dype from models.tt_transformers.tt.ccl import TT_CCL from models.tt_transformers.tt.model_config import ModelArgs dtype = ttnn.bfloat8_b batch_size = 1 mode = "decode" if seq_len <= 32 else "prefill" model_args = ModelArgs(ttnn_mesh_device, max_batch_size=batch_size, max_seq_len=128, cache_hf=True) model_args.n_layers = 1 if model_args.is_galaxy: pytest.skip("MLP1D test only runs on non-TG devices") state_dict = model_args.load_state_dict() model_config = model_args.get_model_config() # Load reference model first_layer_prefix = model_args.get_state_dict_prefix("MLP", 0) partial_state_dict = { k[len(first_layer_prefix) + 1 :]: v for k, v in state_dict.items() if k.startswith(first_layer_prefix) } reference_model = model_args.reference_mlp() reference_model.load_state_dict(partial_state_dict) # Create MLP1D def topology_aware_cache_path(dtype): if model_args.instruct: return ( model_args.model_cache_path / { ttnn.bfloat16: f"tensor_cache_instruct_bf16_{ttnn_mesh_device.shape}", ttnn.bfloat8_b: f"tensor_cache_instruct_bfp8_{ttnn_mesh_device.shape}", }[dtype] ) else: return ( model_args.model_cache_path / { ttnn.bfloat16: f"tensor_cache_bf16_{ttnn_mesh_device.shape}", ttnn.bfloat8_b: f"tensor_cache_bfp8_{ttnn_mesh_device.shape}", }[dtype] ) tt_ccl = TT_CCL(ttnn_mesh_device) tt_model = MLP1D.from_model_args( mesh_device=ttnn_mesh_device, tt_ccl=tt_ccl, args=model_args, state_dict=state_dict, weight_cache_path=topology_aware_cache_path(dtype), layer_num=0, dtype=dtype, model_config=model_config, ) # Create input torch_input = torch.randn( 1, 1, seq_len, model_args.dim, dtype=get_ref_model_dype(reference_model, model_args.model_name) ) # Run reference reference_output = reference_model(torch_input) # Run TT model input_mem_config = model_args.get_mlp_input_mem_config(Mode(mode), None) tt_input = ttnn.from_torch( torch_input, device=ttnn_mesh_device, mesh_mapper=ttnn.ShardTensor2dMesh(ttnn_mesh_device, dims=(None, None), mesh_shape=model_args.cluster_shape), dtype=ttnn.bfloat8_b, memory_config=input_mem_config, layout=ttnn.TILE_LAYOUT, ) tt_output = tt_model.forward(tt_input, mode) tt_output_torch = ttnn.to_torch( tt_output, mesh_composer=ttnn.ConcatMesh2dToTensor(ttnn_mesh_device, dims=(1, 3), mesh_shape=model_args.cluster_shape), ) tt_output_torch = tt_output_torch[:, :1, :, :] # Compare pcc_required = 0.99 passing, pcc_message = comp_pcc(reference_output, tt_output_torch, pcc_required) logger.info(comp_allclose(reference_output, tt_output_torch)) logger.info(f"MLP1D vs reference: {pcc_message}") assert passing, f"MLP1D output does not meet PCC requirement {pcc_required}: {pcc_message}." logger.info(f"MLP1D vs reference: PASSED for mode={mode}, seq_len={seq_len}")