clef / code /models /common /modules /moe /tt_moe_decode.py
tt-hous's picture
Add files using upload-large-folder tool
2415c4c verified
Raw History Blame Contribute Delete
47 kB
# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
"""Configurable Mixture-of-Experts decode block (`TTMoEDecode`).
A single forward step of a decode-time MoE layer on a 2D device mesh, wrapping the
ttnn `all_to_all_dispatch_metadata` → `moe_compute` → `deepseek_moe_fast_reduce_nc_fused`
→ reduce-scatter pipeline. All op kwargs (memory configs, cluster topology, splits,
shared-expert plumbing) are driven by `TTMoEDecodeConfig`, which derives sane defaults
from a minimal YAML per model. This module's job is just to wire the configured pieces
together and expose a `forward(x, scores, indices)` interface.
Pipeline overview (one decode step):
1. `all_to_all_dispatch_metadata`: route each token to its `select_experts_k` chosen
routed experts (cross-cluster send) plus all shared experts (local broadcast).
2. `moe_compute`: per-device per-token-slot matmul `x @ w0`, `x @ w1`, activation
(SiLU/SWIGLU/GELU), `intermediate @ w2`, optionally with bias.
3. `deepseek_moe_fast_reduce_nc_fused`: score-weighted combine of the per-expert
outputs back to per-token results, with a fixed scalar `shared_expert_scale`
applied to shared-expert contributions.
4. Reduce-scatter across the replicated mesh axis to produce per-device output
chunks of width `hidden_size / num_replicated`.
Weight ownership: routed experts are sharded across the dispatch axis; shared experts
are replicated. Weights upload as `bfloat4_b` (with bf16 intermediate tiles for bias).
The two private classes `_TTMoEDecodeExpertState` and `_TTMoEDecodeBuffers` exist to
keep weight/mapping init and per-iteration scratch separated from the forward logic.
"""
from __future__ import annotations
import torch
from loguru import logger
from ttnn.experimental.moe_compute_utils import (
auto_output_width_shard_dim,
effective_matmul_ring_size,
map_shared_experts,
)
import ttnn
from models.common.modules.moe.tt_moe_decode_config import TTMoEDecodeConfig
def _tt_to_torch_dtype(tt_dtype):
"""Map a ttnn dtype to the closest host torch dtype for buffer allocation.
Only the dtypes this module actually uses are handled — `bfloat8_b` falls back to
`torch.bfloat16` since torch has no native 8-bit float; the host tensor is just a
placeholder that gets reinterpreted at upload time.
"""
if tt_dtype == ttnn.bfloat16 or tt_dtype == ttnn.bfloat8_b:
return torch.bfloat16
if tt_dtype == ttnn.float32:
return torch.float32
if tt_dtype == ttnn.uint16:
return torch.uint16
raise ValueError(f"Unsupported tt dtype: {tt_dtype}")
class _TTMoEDecodeExpertState:
"""Owns per-layer routed + shared expert weights and the global expert-mapping table.
Holds three uploaded ttnn tensors after init:
- `tt_expert_mapping`: `[num_devices, num_experts]` lookup of which linearized mesh
coord owns each expert, replicated to every device. Used by both `dispatch` and
`fast_reduce` ops.
- `tt_w0_w1`: interleaved-and-tile-reordered w0/w1 weights for `moe_compute`'s
consumption, sharded across mesh devices along the expert dim. `bfloat4_b`.
- `tt_w2`: same idea, but with the ring-rotated N-tile layout w2 needs. `bfloat4_b`.
Bias support: when `has_bias=True`, biases are folded into the same prepared weight
tensors via `prepare_w0_w1_tensor_with_bias` / `prepare_w2_tensor_with_bias` (one
extra K/N tile each). Bias + shared experts together is not supported (raises).
"""
def _load_weights(self):
# TODO (AM) eventually support loading weights from path
pass
@staticmethod
def _validate(
torch_w0: "torch.Tensor",
torch_w1: "torch.Tensor",
torch_w2: "torch.Tensor",
mesh_device: ttnn.MeshDevice,
mesh_shape: tuple[int, int],
cluster_axis: int,
has_bias: bool,
num_routed_experts: int,
expert_mapping: list[int],
num_shared_experts: int,
shared_expert_ids_to_devices: dict[int, list[int]] | None,
shared_id_to_torch_w0: dict[int, "torch.Tensor"] | None,
shared_id_to_torch_w1: dict[int, "torch.Tensor"] | None,
shared_id_to_torch_w2: dict[int, "torch.Tensor"] | None,
torch_b0: "torch.Tensor" | None,
torch_b1: "torch.Tensor" | None,
torch_b2: "torch.Tensor" | None,
) -> None:
"""Fail fast on (weights, biases, shared experts) ↔ config mismatch.
Cross-checks tensor shapes against each other (w0 is the source of truth for
`L`, `H`, `N`) and against the config-derived counts. Catches the common
misconfigurations that otherwise surface as opaque shape errors deep inside
the bf4 preparers or kernels.
"""
# --- mesh / topology sanity ---
if cluster_axis not in (0, 1):
raise ValueError(f"cluster_axis must be 0 or 1, got {cluster_axis}")
num_devices = mesh_device.get_num_devices()
if mesh_shape[0] * mesh_shape[1] != num_devices:
raise ValueError(
f"mesh_shape {mesh_shape} (= {mesh_shape[0] * mesh_shape[1]} devices) does "
f"not match mesh_device.get_num_devices() = {num_devices}"
)
if num_routed_experts % num_devices != 0:
raise ValueError(
f"num_routed_experts ({num_routed_experts}) must be divisible by num_devices ({num_devices})"
)
# --- routed weight shape cross-checks (w0 is the source of truth for L, H, N) ---
if torch_w0.ndim != 4:
raise ValueError(f"torch_w0 must be 4D [L, E, H, N], got shape {tuple(torch_w0.shape)}")
L, E, H, N = torch_w0.shape
if E != num_routed_experts:
raise ValueError(f"torch_w0.shape[1] = {E} does not match num_routed_experts ({num_routed_experts})")
if tuple(torch_w1.shape) != (L, E, H, N):
raise ValueError(f"torch_w1 shape {tuple(torch_w1.shape)} must match torch_w0 shape ({L}, {E}, {H}, {N})")
if tuple(torch_w2.shape) != (L, E, N, H):
raise ValueError(f"torch_w2 shape {tuple(torch_w2.shape)} must be (L, E, N, H) = ({L}, {E}, {N}, {H})")
# --- expert_mapping ---
if len(expert_mapping) != num_routed_experts:
raise ValueError(
f"expert_mapping length {len(expert_mapping)} != num_routed_experts ({num_routed_experts})"
)
bad = [(e, d) for e, d in enumerate(expert_mapping) if not (0 <= d < num_devices)]
if bad:
raise ValueError(f"expert_mapping has out-of-range device ids (0..{num_devices - 1}): {bad[:5]}")
# --- bias presence and shapes ---
bias_tensors = (torch_b0, torch_b1, torch_b2)
any_bias = any(b is not None for b in bias_tensors)
all_bias = all(b is not None for b in bias_tensors)
if has_bias and not all_bias:
missing = [name for name, b in zip(("b0", "b1", "b2"), bias_tensors) if b is None]
raise ValueError(f"has_bias=True but {missing} not provided")
if not has_bias and any_bias:
raise ValueError("has_bias=False but one or more of torch_b0/b1/b2 was provided")
if has_bias:
if tuple(torch_b0.shape) != (L, E, N) or tuple(torch_b1.shape) != (L, E, N):
raise ValueError(
f"b0/b1 must be (L, E, N) = ({L}, {E}, {N}); got "
f"{tuple(torch_b0.shape)} and {tuple(torch_b1.shape)}"
)
if tuple(torch_b2.shape) != (L, E, H):
raise ValueError(f"b2 must be (L, E, H) = ({L}, {E}, {H}); got {tuple(torch_b2.shape)}")
# --- shared experts ---
if num_shared_experts > 0:
if shared_expert_ids_to_devices is None:
raise ValueError(f"num_shared_experts={num_shared_experts} but shared_expert_ids_to_devices is None")
if len(shared_expert_ids_to_devices) != num_shared_experts:
raise ValueError(
f"shared_expert_ids_to_devices has {len(shared_expert_ids_to_devices)} entries "
f"but num_shared_experts={num_shared_experts}"
)
shared_experts_per_device = [0] * num_devices
for edl in shared_expert_ids_to_devices.values():
for d in edl:
shared_experts_per_device[d] += 1
if len(set(shared_experts_per_device)) > 1 or 0 in shared_experts_per_device:
raise ValueError(
"Every device, should have the same number of, and at least 1 shared expert:"
f" {shared_experts_per_device=}"
)
expected_ids = set(range(num_routed_experts, num_routed_experts + num_shared_experts))
if set(shared_expert_ids_to_devices.keys()) != expected_ids:
raise ValueError(
f"shared expert ids must be contiguous after routed: expected {sorted(expected_ids)}, "
f"got {sorted(shared_expert_ids_to_devices.keys())}"
)
shared_dicts = (shared_id_to_torch_w0, shared_id_to_torch_w1, shared_id_to_torch_w2)
if any(d is None for d in shared_dicts):
raise ValueError("num_shared_experts > 0 but shared_id_to_torch_w0/w1/w2 not all provided")
for name, d in zip(("w0", "w1", "w2"), shared_dicts):
if set(d.keys()) != expected_ids:
raise ValueError(
f"shared_id_to_torch_{name} keys {sorted(d.keys())} != "
f"shared expert ids {sorted(expected_ids)}"
)
expected_shape = (L, 1, N, H) if name == "w2" else (L, 1, H, N)
for sid, t in d.items():
if tuple(t.shape) != expected_shape:
raise ValueError(
f"shared_id_to_torch_{name}[{sid}] shape {tuple(t.shape)} != expected {expected_shape}"
)
if has_bias:
# Mirrors the runtime check in __init__; surface it here too so it fails
# before any device upload.
raise NotImplementedError("bias + shared experts is not yet supported")
else:
if shared_expert_ids_to_devices:
raise ValueError(
f"num_shared_experts=0 but shared_expert_ids_to_devices is non-empty: "
f"{shared_expert_ids_to_devices}"
)
if any(d is not None for d in (shared_id_to_torch_w0, shared_id_to_torch_w1, shared_id_to_torch_w2)):
raise ValueError("num_shared_experts=0 but one or more shared_id_to_torch_w* dicts were provided")
@staticmethod
def _init_expert_mapping(
torch_expert_mapping: "torch.Tensor",
shared_expert_ids_to_devices: dict[int, list[int]] | None,
mesh_device: ttnn.MeshDevice,
mesh_shape: tuple[int, int],
cluster_axis: int,
):
"""Build and upload the `[num_devices, num_experts]` expert-mapping table.
`mapping[d, e]` = linearized mesh coord of the device that owns expert `e`.
Routed experts have the same value across all source rows `d` (ownership is
global), so we just `repeat` the 1D input. Shared experts let different source
devices pick different replicas based on cluster distance — `map_shared_experts`
rewrites those columns to the nearest replica per source row.
Matches the "new format" used by `test_all_to_all_dispatch_metadata_6U.py` and
`gen_expert_mapping` in `test_moe_compute_6U.py`. Replicated to every device
because both dispatch and fast-reduce need the full lookup locally.
"""
if torch_expert_mapping.ndim != 1:
raise ValueError(
f"expected 1D expert_mapping (length=num_experts), got shape {tuple(torch_expert_mapping.shape)}"
)
num_devices = mesh_device.get_num_devices()
mapping_2d = torch_expert_mapping.to(torch.int32).unsqueeze(0).repeat(num_devices, 1)
if shared_expert_ids_to_devices is not None:
mapping_2d = map_shared_experts(mapping_2d, shared_expert_ids_to_devices, mesh_shape, cluster_axis)
return ttnn.from_torch(
mapping_2d,
device=mesh_device,
layout=ttnn.ROW_MAJOR_LAYOUT,
dtype=ttnn.uint16,
memory_config=ttnn.DRAM_MEMORY_CONFIG,
mesh_mapper=ttnn.ReplicateTensorToMesh(mesh_device),
)
@staticmethod
def _device_reorder_weights(
torch_expert_mapping: "torch.Tensor",
torch_w0: "torch.Tensor",
torch_w1: "torch.Tensor",
torch_w2: "torch.Tensor",
torch_b0: "torch.Tensor" | None,
torch_b1: "torch.Tensor" | None,
torch_b2: "torch.Tensor" | None,
):
"""Permute the expert dim so that `ShardTensorToMesh(dim=experts)` lands each
expert on its assigned device.
The host weights come in routed-expert-id order (`[L, num_experts, ...]`), but
sharding by the expert dim splits contiguously — so without reordering, device 0
gets experts [0..E/D), device 1 gets [E/D..2E/D), etc. `expert_mapping[e]` tells
us the *target* device for expert `e`; `argsort` (stable) groups experts that
share a target device into contiguous chunks in the right order. Same permutation
applies to biases when present.
"""
perm = torch.argsort(torch_expert_mapping, stable=True)
mapped_tensors = [t[:, perm, :, :] for t in (torch_w0, torch_w1, torch_w2)]
if torch_b0 is not None:
mapped_tensors += [t[:, perm, :] for t in (torch_b0, torch_b1, torch_b2)]
else:
mapped_tensors += [None] * 3
return tuple(mapped_tensors)
def __init__(
self,
mesh_device: ttnn.MeshDevice,
torch_w0: "torch.Tensor",
torch_w1: "torch.Tensor",
torch_w2: "torch.Tensor",
*,
mesh_shape: tuple[int, int],
cluster_axis: int,
has_bias: bool,
num_routed_experts: int,
expert_mapping: list[int],
num_shared_experts: int,
shared_expert_ids_to_devices: dict[int, list[int]] | None = None,
shared_id_to_torch_w0: dict[int, "torch.Tensor"] | None = None,
shared_id_to_torch_w1: dict[int, "torch.Tensor"] | None = None,
shared_id_to_torch_w2: dict[int, "torch.Tensor"] | None = None,
torch_b0: "torch.Tensor" | None = None,
torch_b1: "torch.Tensor" | None = None,
torch_b2: "torch.Tensor" | None = None,
):
"""Prepare and upload all expert state to the mesh.
Pipeline: argsort-permute routed weights on host to match the device assignment
(`_device_reorder_weights`), upload them to the mesh sharded on the experts dim
(dim 1) as bf16, optionally splice in shared experts on-device so each device holds
`routed_per_device + shared_per_device` slots (`ttnn.experimental.add_shared_expert_weights`),
then run the C++ on-device packers (`ttnn.experimental.prepare_*`) and quantize to
`bfloat4_b` (`ttnn.experimental.quantize_weights_via_host`) — see
`_init_total_expert_weights_impl`.
Routed weight shapes (host, post-permute): `w0/w1 = [L, num_routed, H, N]`,
`w2 = [L, num_routed, N, H]`. Shared weights are dicts keyed by global expert id;
each value has shape `[L, 1, ...]` matching the routed layout.
Biases match the routed shape minus the matmul-output dim (`[L, num_routed, N]`
for `b0/b1`, `[L, num_routed, H]` for `b2`). Shared experts + bias is not yet
supported because the shared splice doesn't carry bias rows — would need a parallel API.
"""
self._validate(
torch_w0,
torch_w1,
torch_w2,
mesh_device,
mesh_shape,
cluster_axis,
has_bias,
num_routed_experts,
expert_mapping,
num_shared_experts,
shared_expert_ids_to_devices,
shared_id_to_torch_w0,
shared_id_to_torch_w1,
shared_id_to_torch_w2,
torch_b0,
torch_b1,
torch_b2,
)
# An empty mapping ({} when num_shared_experts==0) means "no shared experts" —
# normalize to None so the shared-expert paths below are skipped entirely.
if not shared_expert_ids_to_devices:
shared_expert_ids_to_devices = None
num_routed = torch_w0.shape[1]
logger.info(
f"Initializing expert state: routed_experts={num_routed} num_shared={num_shared_experts} "
f"mesh_shape={mesh_shape} cluster_axis={cluster_axis} has_bias={has_bias}"
)
torch_expert_mapping = torch.tensor(expert_mapping, dtype=torch.int)
(
mapped_torch_w0,
mapped_torch_w1,
mapped_torch_w2,
mapped_torch_b0,
mapped_torch_b1,
mapped_torch_b2,
) = self._device_reorder_weights(
torch_expert_mapping, torch_w0, torch_w1, torch_w2, torch_b0, torch_b1, torch_b2
)
num_devices = mesh_device.get_num_devices()
num_layers = torch_w0.shape[0]
hidden_size = torch_w0.shape[-2]
intermediate_size = torch_w0.shape[-1]
routed_per_device = num_routed // num_devices
# Upload the (reordered) routed weights to the mesh sharded on the experts dim (dim 1),
# bf16 / ROW_MAJOR, so the C++ `ttnn.experimental.prepare_*` helpers can do the layout
# transform on-device. Each device receives its assigned contiguous expert chunk.
def _shard_experts(t):
return ttnn.from_torch(
t,
device=mesh_device,
layout=ttnn.ROW_MAJOR_LAYOUT,
dtype=ttnn.bfloat16,
mesh_mapper=ttnn.ShardTensorToMesh(mesh_device, dim=1),
)
tt_w0 = _shard_experts(mapped_torch_w0)
tt_w1 = _shard_experts(mapped_torch_w1)
tt_w2 = _shard_experts(mapped_torch_w2)
if shared_expert_ids_to_devices is not None:
if has_bias:
# add_shared_expert_weights only handles weights; extending it to bias
# requires a parallel API and per-shared-expert bias tensors.
raise NotImplementedError("bias + shared experts is not yet supported")
logger.info(f"Adding shared expert weights for {len(shared_expert_ids_to_devices)} shared experts")
# The C++ add_shared_expert_weights consumes pre-arranged shared *device* tensors
# (per device: its assigned shared experts in sorted-id order; then concatenated
# across devices). Arrange on host, upload sharded on dim 1, splice on-device.
tt_shared_w0 = _shard_experts(
self._arrange_shared(shared_id_to_torch_w0, shared_expert_ids_to_devices, num_devices)
)
tt_shared_w1 = _shard_experts(
self._arrange_shared(shared_id_to_torch_w1, shared_expert_ids_to_devices, num_devices)
)
tt_shared_w2 = _shard_experts(
self._arrange_shared(shared_id_to_torch_w2, shared_expert_ids_to_devices, num_devices)
)
routed_w0, routed_w1, routed_w2 = tt_w0, tt_w1, tt_w2
tt_w0, tt_w1, tt_w2 = ttnn.experimental.add_shared_expert_weights(
routed_w0,
routed_w1,
routed_w2,
tt_shared_w0,
tt_shared_w1,
tt_shared_w2,
cluster_axis=cluster_axis,
)
for t in (routed_w0, routed_w1, routed_w2, tt_shared_w0, tt_shared_w1, tt_shared_w2):
ttnn.deallocate(t)
total_shared_slots = sum(len(devs) for devs in shared_expert_ids_to_devices.values())
experts_per_device = (num_routed + total_shared_slots) // num_devices
else:
experts_per_device = routed_per_device
tt_b0 = tt_b1 = tt_b2 = None
if has_bias:
tt_b0 = _shard_experts(mapped_torch_b0)
tt_b1 = _shard_experts(mapped_torch_b1)
tt_b2 = _shard_experts(mapped_torch_b2)
self.tt_expert_mapping = self._init_expert_mapping(
torch_expert_mapping, shared_expert_ids_to_devices, mesh_device, mesh_shape, cluster_axis
)
logger.info("Uploaded expert mapping to mesh")
self.tt_w0_w1, self.tt_w2 = self._init_total_expert_weights_impl(
tt_w0,
tt_w1,
tt_w2,
mesh_device,
num_layers,
experts_per_device,
hidden_size,
intermediate_size,
has_bias,
tt_b0,
tt_b1,
tt_b2,
)
@staticmethod
def _arrange_shared(
shared_dict: dict[int, "torch.Tensor"],
shared_expert_ids_to_devices: dict[int, list[int]],
num_devices: int,
) -> "torch.Tensor":
"""Arrange a `{shared_expert_id: tensor}` dict into the per-device-stacked layout the
C++ `add_shared_expert_weights` expects.
For each device (in order), concatenate its assigned shared experts (in sorted global-id
order) along the experts dim (dim 1), then concatenate across devices. Sharding the
result on dim 1 then lands each device's shared experts on it, in slot order — matching
how the device-side splice appends shared experts after routed ones per device.
"""
device_to_shared: list[list[int]] = [[] for _ in range(num_devices)]
for sid in sorted(shared_expert_ids_to_devices):
for d in shared_expert_ids_to_devices[sid]:
device_to_shared[d].append(sid)
per_device = [torch.cat([shared_dict[sid] for sid in device_to_shared[d]], dim=1) for d in range(num_devices)]
return torch.cat(per_device, dim=1)
@staticmethod
def _init_total_expert_weights_impl(
tt_w0: ttnn.Tensor,
tt_w1: ttnn.Tensor,
tt_w2: ttnn.Tensor,
mesh_device: ttnn.MeshDevice,
num_layers: int,
experts_per_device: int,
hidden_size: int,
intermediate_size: int,
has_bias: bool,
tt_b0: ttnn.Tensor | None = None,
tt_b1: ttnn.Tensor | None = None,
tt_b2: ttnn.Tensor | None = None,
) -> tuple[ttnn.Tensor, ttnn.Tensor]:
"""Pack and quantize the combined routed+shared weight device tensors to `bfloat4_b`.
Inputs are multi-device tensors sharded on the experts dim (dim 1), bf16 / ROW_MAJOR.
Runs the C++ on-device packers (`ttnn.experimental.prepare_*`, with the `_with_bias`
variants when bias rows are folded in), fetches the DRAM-sharded mem configs from
`ttnn.experimental.get_weight_mem_configs`, then quantizes to `bfloat4_b` onto those
configs via `ttnn.experimental.quantize_weights_via_host` (host round-trip).
`experts_per_device` is the *combined* (routed + shared) per-device expert count — the
experts dim of the inputs, divided across the mesh. Returns the two device tensors
`(tt_w0_w1, tt_w2)` ready to feed `moe_compute`.
"""
logger.info(
f"Preparing expert weights on device: per_device={experts_per_device} "
f"hidden={hidden_size} intermediate={intermediate_size} has_bias={has_bias}"
)
if has_bias:
tt_w0_w1_prepped = ttnn.experimental.prepare_w0_w1_tensor_with_bias(
tt_w0, tt_w1, tt_b0, tt_b1, L=num_layers, E=experts_per_device, K=hidden_size, N=intermediate_size
)
tt_w2_prepped = ttnn.experimental.prepare_w2_tensor_with_bias(
tt_w2, tt_b2, L=num_layers, E=experts_per_device, N=intermediate_size, K=hidden_size
)
else:
tt_w0_w1_prepped = ttnn.experimental.prepare_w0_w1_tensor_for_moe_compute(
tt_w0, tt_w1, L=num_layers, E=experts_per_device, K=hidden_size, N=intermediate_size
)
tt_w2_prepped = ttnn.experimental.prepare_w2_tensor_for_moe_compute(
tt_w2, L=num_layers, E=experts_per_device, N=intermediate_size, K=hidden_size
)
# Raw uploaded inputs are consumed by the packers; free them before the host round-trip.
for t in (tt_w0, tt_w1, tt_w2):
ttnn.deallocate(t)
if has_bias:
for t in (tt_b0, tt_b1, tt_b2):
ttnn.deallocate(t)
# has_bias grows the padded K / N by a tile to accommodate the bias row.
weight_mem_configs = ttnn.experimental.get_weight_mem_configs(
mesh_device,
num_layers=num_layers,
experts_per_device=experts_per_device,
hidden_size=hidden_size,
intermediate_size=intermediate_size,
has_bias=has_bias,
)
tt_w0_w1 = ttnn.experimental.quantize_weights_via_host(
tt_w0_w1_prepped, dtype=ttnn.bfloat4_b, memory_config=weight_mem_configs.w0_w1
)
ttnn.deallocate(tt_w0_w1_prepped)
tt_w2 = ttnn.experimental.quantize_weights_via_host(
tt_w2_prepped, dtype=ttnn.bfloat4_b, memory_config=weight_mem_configs.w2
)
ttnn.deallocate(tt_w2_prepped)
logger.info("Prepared and quantized w0/w1 and w2 to bfloat4_b on mesh")
return tt_w0_w1, tt_w2
class _TTMoEDecodeBuffers:
"""Persistent buffers and semaphores for the MoE decode pipeline.
Allocates the dispatch output triple (sparse buffer, expert indices, expert
scores), the dispatch/combine cross-device semaphores, and the combine
output buffer once in __init__, for reuse across forward() calls.
Shapes and memory configs mirror those used in test_optimized_moe_decode_block.py.
"""
SPARSE_BUFFER_DTYPE = ttnn.bfloat16
INDICES_DTYPE = ttnn.uint16
SCORES_DTYPE = ttnn.bfloat16
COMBINE_OUTPUT_DTYPE = ttnn.bfloat16
def __init__(
self,
mesh_device: ttnn.MeshDevice,
*,
mesh_shape: tuple[int, int],
cluster_axis: int,
batch_per_device: int,
hidden_size: int,
effective_experts_k: int,
shard_dim: int,
compute_tilize_drain_core: ttnn.CoreCoord,
):
"""Allocate the persistent buffers and semaphores reused across `forward()` calls.
Allocates four ttnn tensors and two global semaphores:
- `dispatch_global_semaphore` / `combine_global_semaphore`: cross-device sync
points. Single-use per forward (no double buffering needed — combine syncs
after fully reading the dispatch output, dispatch syncs at end of pipeline).
- `tt_dispatch_output_tensors` triple: sparse buffer (DRAM, hidden-wide token slots),
indices and scores (both L1 height-sharded on the drain core, narrow).
- `tt_combine_output`: DRAM `[effective_experts_k, batch_per_device, hidden_size]`
intermediate after `moe_compute`, before the post-combine tilize.
`effective_experts_k = select_experts_k + num_shared_experts` — the K dimension of
the per-token expert-output stack. Sized from config; passed in to keep this class
agnostic of where the value came from.
"""
# --- derived sizes (seq=1 for decode) ---
num_dispatch_devices = mesh_shape[cluster_axis]
total_tokens = batch_per_device * num_dispatch_devices
tokens_per_device = batch_per_device
shard_dims = (shard_dim, None) if cluster_axis == 0 else (None, shard_dim)
# --- global semaphores (one each, no double buffering required —
# combine syncs after reading dispatch output, dispatch syncs at end) ---
compute_grid_size = mesh_device.compute_with_storage_grid_size()
worker_cores = ttnn.CoreRangeSet(
{ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(compute_grid_size.x - 1, compute_grid_size.y - 1))}
)
self.dispatch_global_semaphore = ttnn.create_global_semaphore(mesh_device, worker_cores, 0)
self.combine_global_semaphore = ttnn.create_global_semaphore(mesh_device, worker_cores, 0)
# --- dispatch output buffers ---
# Sparse buffer: DRAM, row-major, sharded along cluster axis
sparse_buffer = ttnn.from_torch(
torch.zeros(
[num_dispatch_devices, total_tokens, hidden_size],
dtype=_tt_to_torch_dtype(self.SPARSE_BUFFER_DTYPE),
),
device=mesh_device,
layout=ttnn.ROW_MAJOR_LAYOUT,
dtype=self.SPARSE_BUFFER_DTYPE,
memory_config=ttnn.DRAM_MEMORY_CONFIG,
mesh_mapper=ttnn.ShardTensor2dMesh(mesh_device, dims=shard_dims, mesh_shape=mesh_shape),
)
# Indices / scores share an L1 height-sharded mem config on the drain core
shard_spec = ttnn.ShardSpec(
ttnn.CoreRangeSet({ttnn.CoreRange(compute_tilize_drain_core, compute_tilize_drain_core)}),
[total_tokens, effective_experts_k],
ttnn.ShardOrientation.ROW_MAJOR,
)
l1_height_sharded = ttnn.MemoryConfig(
ttnn.TensorMemoryLayout.HEIGHT_SHARDED,
ttnn.BufferType.L1,
shard_spec,
)
indices = ttnn.from_torch(
torch.zeros(
[num_dispatch_devices, total_tokens, effective_experts_k],
dtype=_tt_to_torch_dtype(self.INDICES_DTYPE),
),
device=mesh_device,
layout=ttnn.ROW_MAJOR_LAYOUT,
dtype=self.INDICES_DTYPE,
memory_config=l1_height_sharded,
mesh_mapper=ttnn.ShardTensor2dMesh(mesh_device, dims=shard_dims, mesh_shape=mesh_shape),
)
scores = ttnn.from_torch(
torch.zeros(
[num_dispatch_devices, total_tokens, effective_experts_k],
dtype=_tt_to_torch_dtype(self.SCORES_DTYPE),
),
device=mesh_device,
layout=ttnn.ROW_MAJOR_LAYOUT,
dtype=self.SCORES_DTYPE,
memory_config=l1_height_sharded,
mesh_mapper=ttnn.ShardTensor2dMesh(mesh_device, dims=shard_dims, mesh_shape=mesh_shape),
)
self.tt_dispatch_output_tensors = (sparse_buffer, indices, scores)
self.tt_combine_output = ttnn.from_torch(
torch.zeros(
[effective_experts_k, tokens_per_device, hidden_size],
dtype=_tt_to_torch_dtype(self.COMBINE_OUTPUT_DTYPE),
),
device=mesh_device,
layout=ttnn.ROW_MAJOR_LAYOUT,
dtype=self.COMBINE_OUTPUT_DTYPE,
memory_config=ttnn.DRAM_MEMORY_CONFIG,
)
class TTMoEDecode:
"""MoE decode block: token dispatch → expert compute → score-weighted combine → reduce-scatter.
Constructed once per layer (or per shared layer slot); `forward()` drives one
decode step. All shape / topology / memory-config decisions live in `config`
(`TTMoEDecodeConfig`) — this class only orchestrates the ttnn op sequence.
Two RS branches are auto-selected from `mesh_shape[1 - cluster_axis]`:
- `== DEEPSEEK_RS_DP_DIM (8)`: fused `deepseek_moe_reduce_scatter` consuming the
pre-split list of N outputs from `fast_reduce_nc_fused`.
- `== SKIP_RS_DP_DIM (1)`: no replication, RS is a no-op.
- else: generic `ttnn.reduce_scatter` over the single fast-reduce output.
"""
DEEPSEEK_RS_DP_DIM: int = 8
SKIP_RS_DP_DIM: int = 1
def __init__(
self,
mesh_device: ttnn.MeshDevice,
config: TTMoEDecodeConfig,
torch_w0: torch.Tensor,
torch_w1: torch.Tensor,
torch_w2: torch.Tensor,
shared_id_to_torch_w0: dict[int, torch.Tensor] | None = None,
shared_id_to_torch_w1: dict[int, torch.Tensor] | None = None,
shared_id_to_torch_w2: dict[int, torch.Tensor] | None = None,
torch_b0: torch.Tensor | None = None,
torch_b1: torch.Tensor | None = None,
torch_b2: torch.Tensor | None = None,
):
"""Upload weights / biases / shared experts to the mesh and allocate scratch buffers.
Routed weight shapes: `w0/w1 = [L, num_routed_experts, hidden_size, intermediate_size]`,
`w2 = [L, num_routed_experts, intermediate_size, hidden_size]`. Shared weights are
dicts keyed by global expert id (in `[num_routed, num_routed + num_shared)`), each
value `[L, 1, ...]` matching the routed layout.
Bias shapes (only when `config.has_bias`): `b0/b1 = [L, num_routed_experts, intermediate_size]`,
`b2 = [L, num_routed_experts, hidden_size]`. Bias + shared experts together raises
`NotImplementedError`.
"""
self.config = config
self.expert_state = _TTMoEDecodeExpertState(
mesh_device,
torch_w0,
torch_w1,
torch_w2,
**config.experts.model_dump(),
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,
torch_b0=torch_b0,
torch_b1=torch_b1,
torch_b2=torch_b2,
)
buffers_dict = config.buffers.model_dump()
if buffers_dict.get("compute_tilize_drain_core") is None:
matmul_ring_size = effective_matmul_ring_size(mesh_device)
buffers_dict["compute_tilize_drain_core"] = ttnn.experimental.get_moe_tilize_drain_core(
mesh_device,
config.compute.output_height_shard_dim,
auto_output_width_shard_dim(config.hidden_size, matmul_ring_size=matmul_ring_size),
config.hidden_size,
mux_core_range_set=config.compute.mux_core_range_set,
)
else:
raise ValueError(
"compute_tilize_drain_core is not user-configurable; omit it to resolve dynamically at runtime"
)
self.buffers = _TTMoEDecodeBuffers(mesh_device, **buffers_dict)
@property
def _num_fast_reduce_outputs(self) -> int:
"""Number of outputs fast_reduce_nc_fused will produce — N for the deepseek
RS-list path, 1 otherwise (downstream ttnn.reduce_scatter takes a single tensor)."""
return self.config.num_fast_reduce_outputs
@property
def _pre_split_chunk(self) -> int:
"""Logical per-fast-reduce-output width (hidden_size / num_fast_reduce_outputs).
For the single-output path this is just hidden_size."""
return self.config.hidden_size // self._num_fast_reduce_outputs
@property
def _padded_pre_split_chunk(self) -> int:
"""Aligned per-fast-reduce-output width = config.reduce.split_size."""
return self.config.reduce.split_size
@property
def _post_rs_logical_chunk(self) -> int:
"""Logical per-device width after RS — what the model expects downstream."""
return self.config.hidden_size // self.config.mesh_shape[1 - self.config.cluster_axis]
@property
def _needs_fast_reduce_padding(self) -> bool:
return self._pre_split_chunk != self._padded_pre_split_chunk
def _pad_for_fast_reduce(self, tt_x: ttnn.Tensor) -> ttnn.Tensor:
"""Interleave-pad the H dim so each post-RS per-device chunk is tile-aligned.
Downstream RS splits the H dim evenly into `num_replicated` chunks, so padding
must be inserted at each device-chunk boundary — not just appended at the end —
or device d>0 ends up with data shifted by `d * (padded - logical)` positions.
Layout produced: `[chunk_0, pad_0, ..., chunk_{R-1}, pad_{R-1}]`, R=num_replicated.
Each `chunk_d` is `hidden / R` wide; each `pad_d` brings it up to TILE_SIZE alignment.
"""
if not self._needs_fast_reduce_padding:
return tt_x
num_replicated = self.config.mesh_shape[1 - self.config.cluster_axis]
chunk = self._post_rs_logical_chunk
padded_chunk = self._padded_pre_split_chunk * self._num_fast_reduce_outputs // num_replicated
shape = list(tt_x.shape)
reshaped = ttnn.reshape(tt_x, shape[:-1] + [num_replicated, chunk])
padded = ttnn.pad(
reshaped,
padding=[(0, 0)] * len(shape) + [(0, padded_chunk - chunk)],
value=0.0,
)
return ttnn.reshape(padded, shape[:-1] + [num_replicated * padded_chunk])
def _unpad_after_reduce_scatter(self, tt_final: ttnn.Tensor) -> ttnn.Tensor:
"""Slice the trailing padding off each device's post-RS tensor.
After RS each device holds at least `_post_rs_logical_chunk` valid columns followed
by zero padding (inserted before tilize). No-op when the original hidden splits
evenly without padding.
"""
if not self._needs_fast_reduce_padding:
return tt_final
chunk = self._post_rs_logical_chunk
shape = list(tt_final.shape)
start = [0] * len(shape)
end = shape[:-1] + [chunk]
return ttnn.slice(tt_final, start, end)
def _format_dispatch_inputs(
self,
tt_x: ttnn.Tensor,
tt_indices: ttnn.Tensor,
tt_scores: ttnn.Tensor,
):
"""Coerce each dispatch input into the memory config the dispatch op needs.
For each of (`tt_x`, `tt_indices`, `tt_scores`) returns a `(tensor, dealloc_flag)`
pair: if a `to_memory_config` was necessary, the returned tensor is a fresh
allocation that `forward()` should deallocate after the dispatch op; if the input
already matched, the original is returned with `dealloc=False` to leave caller
ownership intact.
"""
if tt_x.memory_config() != self.config.dispatch_input_memory_config:
tt_dispatch_input_tensor_bundle = (
ttnn.to_memory_config(tt_x, memory_config=self.config.dispatch_input_memory_config),
True,
)
else:
tt_dispatch_input_tensor_bundle = tt_x, False
if tt_indices.memory_config() != self.config.dispatch_input_memory_config:
tt_dispatch_input_expert_indices_tensor_bundle = (
ttnn.to_memory_config(
tt_indices,
memory_config=self.config.dispatch_input_memory_config,
),
True,
)
else:
tt_dispatch_input_expert_indices_tensor_bundle = tt_indices, False
if tt_scores.memory_config() != self.config.dispatch_input_expert_scores_memory_config:
tt_dispatch_input_expert_scores_tensor_bundle = (
ttnn.to_memory_config(
tt_scores,
memory_config=self.config.dispatch_input_expert_scores_memory_config,
),
True,
)
else:
tt_dispatch_input_expert_scores_tensor_bundle = tt_scores, False
return (
tt_dispatch_input_tensor_bundle,
tt_dispatch_input_expert_indices_tensor_bundle,
tt_dispatch_input_expert_scores_tensor_bundle,
)
def forward(
self, tt_x: ttnn.Tensor, tt_scores: ttnn.Tensor, tt_indices: ttnn.Tensor, layer_id: int = 0
) -> ttnn.Tensor:
"""Run one decode-step MoE forward.
Inputs (sharded along the dispatch axis):
- `tt_x`: `[1, batch_per_device, 1, hidden_size]` activations per device.
- `tt_indices`: `[batch_per_device, 1, 1, select_experts_k]` chosen routed experts
per token (uint16).
- `tt_scores`: `[batch_per_device, 1, 1, select_experts_k]` per-(token, k) score
for the routed combine; shared experts use `config.reduce.shared_expert_scale`
uniformly, not this tensor.
- `layer_id`: which slice of the layered weight tensors to use (currently
assumed `0` since the rest of the test/module pipeline is `num_layers=1`).
Output: `[1, 1, batch_per_device, hidden_size / num_replicated]` per device,
i.e. each device holds its post-reduce-scatter chunk of the combined hidden dim.
Pipeline matches the reference test_optimized_moe_decode_block:
1. `all_to_all_dispatch_metadata` → per-device sparse buffer of inbound tokens.
2. `moe_compute` → per-(k, token) expert output stack, optionally with bias.
3. Tilize (`deepseek_moe_post_combine_tilize` when batch_per_device == TILE_SIZE
and an NdShard config is available; else `tilize_with_val_padding` fallback).
4. `deepseek_moe_fast_reduce_nc_fused` → score-weighted sum over k, with shared
experts scaled by `shared_expert_scale`.
5. Reduce-scatter across the replicated axis (3 variants — see class docstring).
6. Strip the per-device-chunk padding inserted before tilize, if any.
"""
(
(tt_dispatch_input_tensor, dealloc_input),
(tt_dispatch_input_expert_indices_tensor, dealloc_indices),
(tt_dispatch_input_expert_scores_tensor, dealloc_scores),
) = self._format_dispatch_inputs(tt_x, tt_indices, tt_scores)
(
tt_dispatch_output_sparse_buffer,
tt_dispatch_output_expert_indices,
tt_dispatch_output_expert_scores,
) = ttnn.experimental.all_to_all_dispatch_metadata(
tt_dispatch_input_tensor,
tt_dispatch_input_expert_indices_tensor,
tt_dispatch_input_expert_scores_tensor,
self.expert_state.tt_expert_mapping,
**self.config.dispatch.model_dump(),
# shared_expert_ids
# cluster_axi
# num_links
# drain_sync_tilizer_core
# worker_mode
# dispatch_algorithm
output_tensors=self.buffers.tt_dispatch_output_tensors,
cross_device_semaphore=self.buffers.dispatch_global_semaphore,
)
if dealloc_input:
ttnn.deallocate(tt_dispatch_input_tensor)
if dealloc_scores:
ttnn.deallocate(tt_dispatch_input_expert_scores_tensor)
(
_,
_,
_,
tt_l1_compute_output,
_,
tt_combine_output,
) = ttnn.experimental.moe_compute(
tt_dispatch_output_sparse_buffer,
tt_dispatch_output_expert_indices,
tt_dispatch_output_expert_scores,
self.expert_state.tt_expert_mapping,
self.expert_state.tt_w0_w1,
self.expert_state.tt_w2,
layer_id=layer_id,
# output_height_shard_dim
# cluster_axis
# mux_core_range_set
# has_bias
# activation_type
**self.config.compute.model_dump(),
optional_output_tensor=self.buffers.tt_combine_output,
optional_cross_device_semaphore=self.buffers.combine_global_semaphore,
)
ttnn.deallocate(tt_l1_compute_output)
# unsqueeze
# [select_experts_k, tokens_per_device, hidden_size] -> [select_experts_k, 1, tokens_per_device, hidden_size]
# Note: this does not reallocate, aliases tt_combine_output so don't dealloc
tt_unsqueezed_output = ttnn.unsqueeze(tt_combine_output, dim=1)
# When hidden / num_replicated isn't tile-aligned, interleave-pad each per-device
# chunk up to split_size so fast_reduce_nc_fused's 128-divisibility check passes.
# No-op when already aligned.
tt_unsqueezed_output = self._pad_for_fast_reduce(tt_unsqueezed_output)
if self.config.use_post_combine_tilize:
tt_tilized_compute_output = ttnn.experimental.deepseek_moe_post_combine_tilize(
tt_unsqueezed_output,
# output_memory_config,
**self.config.post_combine_tilize.model_dump(),
)
else:
output_tensor_shape = list(tt_unsqueezed_output.shape)
output_tensor_shape[2] = ((output_tensor_shape[2] + ttnn.TILE_SIZE - 1) // ttnn.TILE_SIZE) * ttnn.TILE_SIZE
tt_tilized_compute_output = ttnn.tilize_with_val_padding(
tt_unsqueezed_output,
output_tensor_shape=output_tensor_shape,
pad_value=0.0,
# memory_config
**self.config.tilize_with_val_padding.model_dump(),
)
# scale with scores and accumulate
tt_fast_reduce_output_tensors = ttnn.experimental.deepseek_moe_fast_reduce_nc_fused(
tt_tilized_compute_output,
tt_dispatch_input_expert_indices_tensor,
self.expert_state.tt_expert_mapping,
# reduce_dim
# cluster_axis
# split_size
# output_memory_config
# num_shared_experts
# shared_expert_scale
**self.config.reduce.model_dump(),
scores_tensor=tt_scores,
)
ttnn.deallocate(tt_tilized_compute_output)
if dealloc_indices:
ttnn.deallocate(tt_dispatch_input_expert_indices_tensor)
# [select_experts_k, tokens_per_device, hidden_size // num_replicated_devices] final per device shape
if (
self.config.mesh_shape[1 - self.config.cluster_axis] == self.DEEPSEEK_RS_DP_DIM
and self.config.topology == ttnn.Topology.Ring
):
tt_final_output = ttnn.experimental.deepseek_moe_reduce_scatter(
tt_fast_reduce_output_tensors,
# output_memory_config
# dim
# num_links
# topology
# cluster_axis
**self.config.deepseek_moe_reduce_scatter.model_dump(),
)
for t in tt_fast_reduce_output_tensors:
ttnn.deallocate(t)
# note: in this path the output is L1 sharded as if set up for deepseek_moe_reduce_scatter. Likely better to
# switch this to something more generic.
elif self.config.mesh_shape[1 - self.config.cluster_axis] == self.SKIP_RS_DP_DIM:
tt_final_output = tt_fast_reduce_output_tensors[0]
else:
tt_final_output = ttnn.reduce_scatter(
tt_fast_reduce_output_tensors[0], **self.config.reduce_scatter.model_dump()
)
for t in tt_fast_reduce_output_tensors:
ttnn.deallocate(t)
# Strip the per-chunk padding we inserted before tilize. No-op when no padding was applied.
return self._unpad_after_reduce_scatter(tt_final_output)