Download code/models/common/tests/modules/mlp/test_mlp_1d.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 71.5 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/tests/modules/mlp/test_mlp_1d.py
- Command line
-
hf download hf://tt-hous/clef/code/models/common/tests/modules/mlp/test_mlp_1d.py
-
curl -L -o test_mlp_1d.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/tests/modules/mlp/test_mlp_1d.py
71.5 kB
| # 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 | |
| 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, | |
| ) | |
| 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, | |
| ) | |
| 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, | |
| ) | |
| 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 | |
| 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 | |
| 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}") | |