Download code/models/common/modules/moe/tt_moe_decode.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 47 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/modules/moe/tt_moe_decode.py
- Command line
-
hf download hf://tt-hous/clef/code/models/common/modules/moe/tt_moe_decode.py
-
curl -L -o tt_moe_decode.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/modules/moe/tt_moe_decode.py
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 | |
| 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") | |
| 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), | |
| ) | |
| 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, | |
| ) | |
| 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) | |
| 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) | |
| 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 | |
| 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 | |
| def _padded_pre_split_chunk(self) -> int: | |
| """Aligned per-fast-reduce-output width = config.reduce.split_size.""" | |
| return self.config.reduce.split_size | |
| 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] | |
| 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) | |