Download code/models/common/tests/modules/moe/test_tt_moe_decode.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 26.7 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/tests/modules/moe/test_tt_moe_decode.py
- Command line
-
hf download hf://tt-hous/clef/code/models/common/tests/modules/moe/test_tt_moe_decode.py
-
curl -L -o test_tt_moe_decode.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/tests/modules/moe/test_tt_moe_decode.py
26.7 kB
| # SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. | |
| # | |
| # SPDX-License-Identifier: Apache-2.0 | |
| """Integration test for `models.common.modules.moe.tt_moe_decode.TTMoEDecode`. | |
| Setup mirrors `test_optimized_moe_decode_block.py`: build torch weights / inputs, | |
| push them through the TTMoEDecode module, and verify the final output against a | |
| torch reference. Combine output verification is intentionally skipped — that | |
| intermediate buffer is exercised by the optimized-block test directly. | |
| Parametrized over every YAML model config in `models/common/modules/moe/configs/`. | |
| """ | |
| from __future__ import annotations | |
| import faulthandler | |
| import os | |
| import random | |
| import sys | |
| import traceback | |
| from pathlib import Path | |
| import pytest | |
| import torch | |
| from loguru import logger | |
| from ttnn.operations.ccl import MoEActivationFunction | |
| import ttnn | |
| from models.common.modules.moe.tt_moe_decode import TTMoEDecode | |
| from models.common.modules.moe.tt_moe_decode_config import TTMoEDecodeConfig | |
| from models.common.utility_functions import is_blackhole | |
| from models.demos.deepseek_v3.tests.fused_op_unit_tests.moe.test_optimized_moe_decode_block import ( | |
| create_torch_dispatch_input_expert_scores_tensor, | |
| create_torch_dispatch_input_tensor, | |
| verify_output, | |
| ) | |
| from tests.nightly.tg.ccl.moe.test_moe_compute_6U import _swiglu_reference | |
| faulthandler.enable() | |
| # occasionally running this test hangs due to conflicts with mesh device teardown. Enable this fixture if encountered | |
| def _hang_watchdog(): | |
| faulthandler.dump_traceback_later(300, exit=True) | |
| try: | |
| yield | |
| finally: | |
| faulthandler.cancel_dump_traceback_later() | |
| def _print_exception_and_fail(reason: str) -> None: | |
| """Print the active exception to the original (uncaptured) stderr and pytest.fail. | |
| pytest's stderr capture buffers `logger.exception()` output and only flushes it | |
| after the test exits. When the watchdog `_exit()`s the process (e.g. because | |
| a ttnn tensor `__repr__` was waiting on a wedged device), that buffer is lost. | |
| Writing to `sys.__stderr__` and flushing bypasses capture so the trace survives. | |
| Then pytest.fail(pytrace=False) skips pytest's own saferepr-driven traceback | |
| rendering, which is itself prone to deadlocking on hung device tensors. | |
| """ | |
| traceback.print_exc(file=sys.__stderr__) | |
| sys.__stderr__.flush() | |
| pytest.fail(reason, pytrace=False) | |
| MESH_GRAPH_DESC_16x1 = ( | |
| "tests/tt_metal/tt_fabric/custom_mesh_descriptors/single_galaxy_16x1_torus_graph_descriptor.textproto" | |
| ) | |
| MESH_GRAPH_DESC_BH_LB_8x1 = "tests/tt_metal/tt_fabric/custom_mesh_descriptors/bh_lb_8x1_line_graph_descriptor.textproto" | |
| def is_mesh_graph_descriptor_set(expected_path): | |
| """Check if TT_MESH_GRAPH_DESC_PATH is set to the expected path.""" | |
| return os.environ.get("TT_MESH_GRAPH_DESC_PATH") == expected_path | |
| # --------------------------------------------------------------------------- | |
| # torch reference helpers | |
| # | |
| # `_swiglu_reference`, `create_torch_dispatch_input_tensor`, | |
| # `create_torch_dispatch_input_expert_scores_tensor`, and `verify_output` are | |
| # imported above from the existing MoE tests — same logic, no need to duplicate. | |
| # Helpers that diverge (per-expert weight/bias init, activation/bias-aware | |
| # matmul, output-golden assembly) are defined locally. | |
| # --------------------------------------------------------------------------- | |
| torch.set_num_threads(max(1, os.cpu_count() or 1)) | |
| def _matmul_golden( | |
| token: torch.Tensor, | |
| w0: torch.Tensor, | |
| w1: torch.Tensor, | |
| w2: torch.Tensor, | |
| activation_type: MoEActivationFunction = MoEActivationFunction.SILU, | |
| b0: torch.Tensor | None = None, | |
| b1: torch.Tensor | None = None, | |
| b2: torch.Tensor | None = None, | |
| ) -> torch.Tensor: | |
| """MoE expert reference (num_layers=1 throughout). | |
| SILU: `silu(x @ w0 + b0) * (x @ w1 + b1) @ w2 + b2` | |
| SWIGLU: `(up + 1) * gate * sigmoid(alpha * gate) @ w2 + b2` with clamping (GPT-OSS), | |
| where `gate = x @ w0 + b0` and `up = x @ w1 + b1`. | |
| GELU: `gelu(x @ w0 + b0, tanh) * (x @ w1 + b1) @ w2 + b2` (tanh approximation | |
| matches the on-device kernel). | |
| Per-expert bias shapes: `b0`/`b1` are `[num_layers, 1, N]`, `b2` is | |
| `[num_layers, 1, hidden_size]`. `unsqueeze(-2)` broadcasts over the token dim. | |
| """ | |
| _orig_dtype = token.dtype | |
| token = token.float() | |
| w0 = w0.float() | |
| w1 = w1.float() | |
| w2 = w2.float() | |
| gate = token @ w0 | |
| if b0 is not None: | |
| gate = gate + b0.float().unsqueeze(-2) | |
| up = token @ w1 | |
| if b1 is not None: | |
| up = up + b1.float().unsqueeze(-2) | |
| if activation_type == MoEActivationFunction.SILU: | |
| intermediate = torch.nn.functional.silu(gate) * up | |
| elif activation_type == MoEActivationFunction.SWIGLU: | |
| intermediate = _swiglu_reference(gate, up) | |
| elif activation_type == MoEActivationFunction.GELU: | |
| intermediate = torch.nn.functional.gelu(gate, approximate="tanh") * up | |
| else: | |
| raise ValueError(f"Unsupported activation type: {activation_type}") | |
| output = intermediate @ w2 | |
| if b2 is not None: | |
| output = output + b2.float().unsqueeze(-2) | |
| return output.to(_orig_dtype) | |
| def _create_per_expert_weights(num_layers: int, num_experts: int, h: int, n: int) -> torch.Tensor: | |
| """Returns a [num_layers, num_experts, h, n] tensor of expert weights with calibrated scale. | |
| TLDR: stabilize output statistics with random weights. Aiming for 0.987 PCC and ATOL < 20 | |
| Weights are drawn from `U[-c, c]` with `c = sqrt(81/h)`. Reasoning: | |
| - Tokens are `U[-0.5, 0.5]` (Var = 1/12). `Var(matmul_out) = h * Var(token) * Var(w)`, | |
| so for output std ≈ 1.5 we want `Var(w) = 27/h`, i.e. `c = sqrt(81/h)` for uniform. | |
| - Why std ≈ 1.5 specifically: empirically the PCC sweet spot. Going smaller (std ≈ 1) | |
| pushes the bulk of gate/up values into the bf16 rounding floor → PCC drops. Going | |
| larger (std ≈ 2) compounds bf4 quantization noise through the three matmul cascade | |
| faster than it benefits from any silu-asymptotic stability → also drops. ~1.5 | |
| threads the needle. | |
| - Uniform (not normal) is critical for bf4_b: bf4 quantization uses a shared exponent | |
| per 16-element block set by the block's max-abs. Uniform draws produce nearly | |
| identical block max-abs across blocks → consistent quantization step everywhere. | |
| Normal draws give some blocks fat-tailed maxes that crush the precision of their | |
| smaller siblings, injecting position-dependent noise that tanks PCC. | |
| To take advantage of this calibration set c = (81.0 / h) ** 0.5 | |
| Note (AM): I have disabled this - c = 0.5 - until I am really confident that any PCC variance is benign | |
| `h` is the matmul input dim for all three of w0, w1, w2 (w2 is called with | |
| `h=intermediate_size`). | |
| """ | |
| c = 0.5 | |
| return ((torch.rand((num_layers, num_experts, h, n), dtype=torch.float32) - 0.5) * (2.0 * c)).to(torch.bfloat16) | |
| def _create_per_expert_biases(num_layers: int, num_experts: int, dim: int) -> torch.Tensor: | |
| """Returns a [num_layers, num_experts, dim] tensor of expert biases. | |
| Variance matches the bias init in `test_moe_compute_6U.py` (std=0.12), but draws are | |
| uniform `U[-c, c]` (`c = sqrt(3) * std`) rather than normal. Reason: the bias row is | |
| packed into the same bf4_b tile as the weights, and bf4's per-block shared exponent | |
| is set by the block's max-abs. Normal draws produce fat-tail blocks that crush the | |
| quantization of their smaller siblings — same issue we fixed for weights. Uniform | |
| draws keep block max-abs consistent and the per-element quantization error tight. | |
| Bias still adds ~8% of the matmul output at the current scale, so any extra | |
| position-dependent noise on the bias channel shows up directly in PCC. | |
| """ | |
| _bias_std = 0.12 | |
| c = (3.0**0.5) * _bias_std | |
| return ((torch.rand(num_layers, num_experts, dim, dtype=torch.float32) - 0.5) * (2.0 * c)).to(torch.bfloat16) | |
| def _create_expert_indices(batch: int, num_experts: int, select_k: int) -> torch.Tensor: | |
| """[batch, 1, 1, select_k] — random unique experts per token.""" | |
| out = torch.full((batch, 1, 1, select_k), -1, dtype=torch.int32) | |
| for b in range(batch): | |
| for k, e in enumerate(random.sample(range(num_experts), select_k)): | |
| out[b, 0, 0, k] = e | |
| return out | |
| def _gen_output_golden( | |
| tokens: torch.Tensor, | |
| expert_indices: torch.Tensor, | |
| expert_scores: torch.Tensor, | |
| w0_per_expert: list[torch.Tensor], | |
| w1_per_expert: list[torch.Tensor], | |
| w2_per_expert: list[torch.Tensor], | |
| batch: int, | |
| hidden_size: int, | |
| select_k: int, | |
| activation_type: MoEActivationFunction = MoEActivationFunction.SILU, | |
| b0_per_expert: list[torch.Tensor] | None = None, | |
| b1_per_expert: list[torch.Tensor] | None = None, | |
| b2_per_expert: list[torch.Tensor] | None = None, | |
| ) -> torch.Tensor: | |
| """[batch, 1, 1, hidden_size] — sum_k(score_k * matmul(token, expert_k)).""" | |
| out = torch.zeros((batch, 1, 1, hidden_size), dtype=torch.bfloat16) | |
| for t in range(batch): | |
| for k in range(select_k): | |
| e = expert_indices[t, 0, 0, k].item() | |
| contrib = _matmul_golden( | |
| tokens[t], | |
| w0_per_expert[e], | |
| w1_per_expert[e], | |
| w2_per_expert[e], | |
| activation_type, | |
| b0=b0_per_expert[e] if b0_per_expert is not None else None, | |
| b1=b1_per_expert[e] if b1_per_expert is not None else None, | |
| b2=b2_per_expert[e] if b2_per_expert is not None else None, | |
| ) | |
| out[t] = out[t] + expert_scores[t, 0, 0, k] * contrib | |
| return out | |
| def _create_shared_expert_weights( | |
| shared_expert_ids: list[int], num_layers: int, h: int, n: int, h2: int | |
| ) -> tuple[dict[int, torch.Tensor], dict[int, torch.Tensor], dict[int, torch.Tensor]]: | |
| """`shared_id -> [num_layers, 1, ...]` tensors for w0/w1/w2. | |
| Matches the format `_TTMoEDecodeExpertState` / `add_shared_expert_weights` expect: | |
| each shared expert is stored individually keyed by its global id. | |
| """ | |
| # Same calibrated uniform scaling as `_create_per_expert_weights`: bf4_b quantizes | |
| # uniform draws far more cleanly than normal draws (no per-block fat-tail outliers | |
| # → consistent quantization step). w0/w1 input dim is `h`, w2 input dim is `n`. | |
| # c_h = (81.0 / h) ** 0.5 | |
| # c_n = (81.0 / n) ** 0.5 | |
| # Note: I have disabled this, c = 0.5, until I am really confident that any PCC/ATOL variance is benign | |
| c_h = 0.5 # (81.0 / h) ** 0.5 | |
| c_n = 0.5 # (81.0 / n) ** 0.5 | |
| shared_w0 = { | |
| sid: ((torch.rand((num_layers, 1, h, n), dtype=torch.float32) - 0.5) * (2.0 * c_h)).to(torch.bfloat16) | |
| for sid in shared_expert_ids | |
| } | |
| shared_w1 = { | |
| sid: ((torch.rand((num_layers, 1, h, n), dtype=torch.float32) - 0.5) * (2.0 * c_h)).to(torch.bfloat16) | |
| for sid in shared_expert_ids | |
| } | |
| shared_w2 = { | |
| sid: ((torch.rand((num_layers, 1, n, h2), dtype=torch.float32) - 0.5) * (2.0 * c_n)).to(torch.bfloat16) | |
| for sid in shared_expert_ids | |
| } | |
| return shared_w0, shared_w1, shared_w2 | |
| def _add_shared_experts_to_golden( | |
| out: torch.Tensor, | |
| tokens: torch.Tensor, | |
| batch: int, | |
| shared_w0: dict[int, torch.Tensor], | |
| shared_w1: dict[int, torch.Tensor], | |
| shared_w2: dict[int, torch.Tensor], | |
| shared_expert_scale: float, | |
| activation_type: MoEActivationFunction, | |
| ) -> torch.Tensor: | |
| """Every token sees every shared expert; contributions add with a fixed scalar scale. | |
| Mirrors `deepseek_moe_fast_reduce_nc_fused`'s shared-expert behavior: no per-token | |
| score, just `shared_expert_scale` applied uniformly. | |
| """ | |
| for sid in shared_w0: | |
| w0, w1, w2 = shared_w0[sid], shared_w1[sid], shared_w2[sid] | |
| for t in range(batch): | |
| contrib = _matmul_golden(tokens[t], w0, w1, w2, activation_type) | |
| out[t] = out[t] + shared_expert_scale * contrib | |
| return out | |
| def _add_shared_experts_to_golden_tp( | |
| out: torch.Tensor, | |
| tokens: torch.Tensor, | |
| batch: int, | |
| shared_w0: dict[int, torch.Tensor], | |
| shared_w1: dict[int, torch.Tensor], | |
| shared_w2: dict[int, torch.Tensor], | |
| shared_expert_scale: float, | |
| activation_type: MoEActivationFunction, | |
| num_tp: int, | |
| ) -> torch.Tensor: | |
| """Shared-expert golden that mimics the tensor-parallel device path step-for-step. | |
| Each shared expert's intermediate dim `N` is partitioned into `num_tp` contiguous | |
| chunks — one per device along the TP axis (`1 - cluster_axis`). Each device's chunk is | |
| zero-padded back to full `N` (front block `[0:N/num_tp]` real, rest zero — exactly the | |
| layout `add_shared_expert_weights` produces: W0/W1 padded on the intermediate dim, W2 | |
| on its row dim), the full-`N` FFN is run on the padded weights to get that device's | |
| partial, and the partials are summed (the reduce-scatter), then scaled by | |
| `shared_expert_scale`. | |
| Because SiLU/SwiGLU/GELU are column-separable and each device's W2 rows are zero outside | |
| its chunk, the summed partials equal the full FFN — so this matches | |
| `_add_shared_experts_to_golden` exactly (modulo bf16 accumulation order). It's written | |
| this way on purpose: it tracks what each device actually computes, so if the device | |
| output diverges from this golden the fault is in the kernel's handling of the | |
| zero-padded TP layout, not the decomposition. No `num_replicated` factor — the sum of | |
| disjoint partials is one full copy, not `num_tp` copies. | |
| """ | |
| for sid in shared_w0: | |
| w0, w1, w2 = shared_w0[sid], shared_w1[sid], shared_w2[sid] | |
| n = w0.shape[-1] | |
| assert n % num_tp == 0, f"shared expert intermediate dim {n} not divisible by num_tp {num_tp}" | |
| chunk = n // num_tp | |
| for t in range(batch): | |
| partial_sum = torch.zeros_like(out[t]) | |
| for d in range(num_tp): | |
| lo, hi = d * chunk, (d + 1) * chunk | |
| # Device d: its chunk, front-block zero-padded to full N. W0/W1 partition the | |
| # intermediate (last) dim; W2 partitions its row (second-to-last) dim. | |
| w0_d = torch.zeros_like(w0) | |
| w1_d = torch.zeros_like(w1) | |
| w2_d = torch.zeros_like(w2) | |
| w0_d[..., :chunk] = w0[..., lo:hi] | |
| w1_d[..., :chunk] = w1[..., lo:hi] | |
| w2_d[..., :chunk, :] = w2[..., lo:hi, :] | |
| partial_sum = partial_sum + _matmul_golden(tokens[t], w0_d, w1_d, w2_d, activation_type) | |
| out[t] = out[t] + shared_expert_scale * partial_sum | |
| return out | |
| CONFIGS_DIR = Path(__file__).resolve().parents[3] / "modules" / "moe" / "configs" | |
| CONFIG_PATHS = sorted(CONFIGS_DIR.glob("*.yaml")) | |
| assert CONFIG_PATHS, f"no YAML configs found in {CONFIGS_DIR}" | |
| def _config_id(path: Path) -> str: | |
| return path.stem | |
| # --------------------------------------------------------------------------- | |
| # test | |
| # --------------------------------------------------------------------------- | |
| # known failures | |
| # Note: it would be better to test all of these and let them fail but some cause hard crashes and derail the test | |
| SKIP_LIST = [ | |
| "ling_1t.yaml", | |
| "mistral_large_3.yaml", | |
| "deepseek_v4_pro.yaml", | |
| ] | |
| def test_tt_moe_decode( | |
| mesh_device: ttnn.MeshDevice, | |
| device_params: dict, | |
| config_path: Path, | |
| num_iterations: int, | |
| ): | |
| torch.manual_seed(2005) | |
| random.seed(2005) | |
| mesh_shape = tuple(mesh_device.shape) | |
| fabric_config = device_params["fabric_config"] | |
| is_line_fabric = fabric_config == ttnn.FabricConfig.FABRIC_1D | |
| if is_line_fabric and mesh_shape != (8, 1): | |
| pytest.skip("FABRIC_1D only valid for (8,1) BH LB mesh") | |
| if not is_line_fabric and mesh_shape == (8, 1): | |
| pytest.skip("(8,1) BH LB requires FABRIC_1D (line topology)") | |
| if str(config_path.name) in SKIP_LIST: | |
| pytest.skip(f"{config_path} is a known failure") | |
| topology = ttnn.Topology.Ring if fabric_config == ttnn.FabricConfig.FABRIC_1D_RING else ttnn.Topology.Linear | |
| config = TTMoEDecodeConfig.from_yaml(config_path.read_text(), topology=topology) | |
| if config.mesh_shape != mesh_shape: | |
| try: | |
| config = config.with_mesh_shape(mesh_shape) | |
| except ValueError as e: | |
| pytest.skip(f"config mesh_shape {config.mesh_shape} can't slice to device mesh_shape {mesh_shape}: {e}") | |
| logger.info(f"Sliced config mesh_shape to {mesh_shape}; num_routed_experts={config.num_routed_experts}") | |
| # --- derived sizes (all from config) --- | |
| cluster_axis = config.cluster_axis | |
| routed_experts = config.num_routed_experts | |
| hidden_size = config.hidden_size | |
| intermediate_size = config.compute.intermediate_size | |
| select_experts_k = config.select_experts_k | |
| batches_per_device = config.batch_per_device | |
| num_devices = mesh_shape[0] * mesh_shape[1] | |
| num_dispatch_devices = mesh_shape[cluster_axis] | |
| batch = batches_per_device * num_dispatch_devices | |
| shard_dim = 0 | |
| shard_dims = (shard_dim, None) if cluster_axis == 0 else (None, shard_dim) | |
| logger.info( | |
| f"Setup [{config_path.stem}]: mesh_shape={mesh_shape} cluster_axis={cluster_axis} " | |
| f"num_devices={num_devices} batch={batch} hidden={hidden_size} N={intermediate_size} " | |
| f"routed_experts={routed_experts} select_experts_k={select_experts_k} " | |
| f"has_bias={config.has_bias} activation={config.compute.activation_type.name}" | |
| ) | |
| # --- weights: [num_layers=1, routed_experts, H/N, N/H] --- | |
| num_layers = 1 | |
| torch_w0 = _create_per_expert_weights(num_layers, routed_experts, hidden_size, intermediate_size) | |
| torch_w1 = _create_per_expert_weights(num_layers, routed_experts, hidden_size, intermediate_size) | |
| torch_w2 = _create_per_expert_weights(num_layers, routed_experts, intermediate_size, hidden_size) | |
| w0_per_expert = [torch_w0[:, e : e + 1, ...] for e in range(routed_experts)] | |
| w1_per_expert = [torch_w1[:, e : e + 1, ...] for e in range(routed_experts)] | |
| w2_per_expert = [torch_w2[:, e : e + 1, ...] for e in range(routed_experts)] | |
| # --- biases (optional): [num_layers=1, routed_experts, N or hidden_size] --- | |
| torch_b0 = torch_b1 = torch_b2 = None | |
| b0_per_expert = b1_per_expert = b2_per_expert = None | |
| if config.has_bias: | |
| torch_b0 = _create_per_expert_biases(num_layers, routed_experts, intermediate_size) | |
| torch_b1 = _create_per_expert_biases(num_layers, routed_experts, intermediate_size) | |
| torch_b2 = _create_per_expert_biases(num_layers, routed_experts, hidden_size) | |
| b0_per_expert = [torch_b0[:, e : e + 1, :] for e in range(routed_experts)] | |
| b1_per_expert = [torch_b1[:, e : e + 1, :] for e in range(routed_experts)] | |
| b2_per_expert = [torch_b2[:, e : e + 1, :] for e in range(routed_experts)] | |
| # --- shared experts (optional): id -> [num_layers, 1, ...] weight dicts --- | |
| shared_id_to_torch_w0 = shared_id_to_torch_w1 = shared_id_to_torch_w2 = None | |
| if config.num_shared_experts > 0: | |
| if config.has_bias: | |
| pytest.skip("TTMoEDecode does not yet support has_bias=True with shared experts") | |
| shared_expert_ids = sorted(config.experts.shared_expert_ids_to_devices.keys()) | |
| shared_id_to_torch_w0, shared_id_to_torch_w1, shared_id_to_torch_w2 = _create_shared_expert_weights( | |
| shared_expert_ids, num_layers, hidden_size, intermediate_size, hidden_size | |
| ) | |
| logger.info( | |
| f"Shared experts: {len(shared_expert_ids)} ids={shared_expert_ids} " | |
| f"scale={config.reduce.shared_expert_scale}" | |
| ) | |
| # --- build module --- | |
| # Wrap in try/except + pytest.fail(pytrace=False) — pytest's pretty-traceback | |
| # saferepr'ing the deeply nested config/mesh args takes long enough that the | |
| # 300s faulthandler watchdog _exit()s the process before any output is shown. | |
| try: | |
| decode = TTMoEDecode( | |
| mesh_device=mesh_device, | |
| config=config, | |
| torch_w0=torch_w0, | |
| torch_w1=torch_w1, | |
| torch_w2=torch_w2, | |
| torch_b0=torch_b0, | |
| torch_b1=torch_b1, | |
| torch_b2=torch_b2, | |
| shared_id_to_torch_w0=shared_id_to_torch_w0, | |
| shared_id_to_torch_w1=shared_id_to_torch_w1, | |
| shared_id_to_torch_w2=shared_id_to_torch_w2, | |
| ) | |
| except Exception as e: | |
| _print_exception_and_fail(f"TTMoEDecode init failed: {type(e).__name__}") | |
| logger.info("Module Setup complete") | |
| # --- per-iteration inputs + goldens --- | |
| tt_dispatch_inputs = [] | |
| tt_dispatch_indices = [] | |
| tt_dispatch_scores = [] | |
| output_goldens = [] | |
| for _ in range(num_iterations): | |
| tokens = create_torch_dispatch_input_tensor(batch, 1, hidden_size, ttnn.bfloat16) | |
| indices = _create_expert_indices(batch, routed_experts, select_experts_k) | |
| scores = create_torch_dispatch_input_expert_scores_tensor(batch, 1, select_experts_k, ttnn.bfloat16) | |
| tt_dispatch_inputs.append( | |
| ttnn.from_torch( | |
| tokens, | |
| device=mesh_device, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| dtype=ttnn.bfloat16, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| mesh_mapper=ttnn.ShardTensor2dMesh(mesh_device, dims=shard_dims, mesh_shape=mesh_shape), | |
| ) | |
| ) | |
| tt_dispatch_indices.append( | |
| ttnn.from_torch( | |
| indices, | |
| device=mesh_device, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| dtype=ttnn.uint16, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| mesh_mapper=ttnn.ShardTensor2dMesh(mesh_device, dims=shard_dims, mesh_shape=mesh_shape), | |
| ) | |
| ) | |
| tt_dispatch_scores.append( | |
| ttnn.from_torch( | |
| scores, | |
| device=mesh_device, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| dtype=ttnn.bfloat16, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| mesh_mapper=ttnn.ShardTensor2dMesh(mesh_device, dims=shard_dims, mesh_shape=mesh_shape), | |
| ) | |
| ) | |
| golden = _gen_output_golden( | |
| tokens, | |
| indices, | |
| scores, | |
| w0_per_expert, | |
| w1_per_expert, | |
| w2_per_expert, | |
| batch, | |
| hidden_size, | |
| select_experts_k, | |
| activation_type=config.compute.activation_type, | |
| b0_per_expert=b0_per_expert, | |
| b1_per_expert=b1_per_expert, | |
| b2_per_expert=b2_per_expert, | |
| ) | |
| if shared_id_to_torch_w0 is not None: | |
| num_tp = mesh_shape[1 - cluster_axis] | |
| golden = _add_shared_experts_to_golden_tp( | |
| golden, | |
| tokens, | |
| batch, | |
| shared_id_to_torch_w0, | |
| shared_id_to_torch_w1, | |
| shared_id_to_torch_w2, | |
| shared_expert_scale=config.reduce.shared_expert_scale, | |
| activation_type=config.compute.activation_type, | |
| num_tp=num_tp, | |
| ) | |
| output_goldens.append(golden) | |
| logger.info("Goldens computed") | |
| # --- run + collect outputs --- | |
| logger.info("Running forward iterations") | |
| tt_outputs = [] | |
| for it in range(num_iterations): | |
| try: | |
| output = decode.forward( | |
| tt_x=tt_dispatch_inputs[it], | |
| tt_scores=tt_dispatch_scores[it], | |
| tt_indices=tt_dispatch_indices[it], | |
| layer_id=0, | |
| ) | |
| if output.memory_config() != ttnn.DRAM_MEMORY_CONFIG: | |
| final_output = ttnn.to_memory_config(output, ttnn.DRAM_MEMORY_CONFIG) | |
| ttnn.deallocate(output) | |
| else: | |
| final_output = output | |
| tt_outputs.append(final_output) | |
| ttnn.synchronize_device(mesh_device, sub_device_ids=[ttnn.SubDeviceId(0)]) | |
| except Exception as e: | |
| _print_exception_and_fail(f"forward iteration {it} failed: {type(e).__name__}") | |
| logger.info(f"Op iteration {it} complete") | |
| # --- verify --- | |
| logger.info("Verifying outputs") | |
| all_passed = True | |
| for it in range(num_iterations): | |
| # ATOL is a tad high, only observed this large for big models (deepseek) might improve with col reduction | |
| # might just need to use calibrated values (see above) but I am still fairly sure it is benign. | |
| if not verify_output(it, mesh_device, mesh_shape, tt_outputs[it], output_goldens[it], atol_threshold=800): | |
| all_passed = False | |
| assert all_passed, f"TTMoEDecode output verification failed for {config_path.stem}" | |