From 104719d69c7509bdf2f624b524f259010f8ae71e Mon Sep 17 00:00:00 2001 From: yh0903 Date: Mon, 24 Aug 2026 16:56:26 -0700 Subject: [PATCH 1/9] Add an opt-in fused local token engine for AutoEP The AutoEP path moves routed rows around the two expert all-to-alls with general-purpose tensor ops. It materializes a padded copy of the token matrix, gathers through an advanced index, scatters expert outputs into a zero-filled buffer, and builds a [tokens, top_k, hidden] FP32 intermediate only to apply routing weights and reduce over top-k. Every one of those steps costs a full pass over the routed activations, and they repeat in every MoE layer. Add expert_parallel.local_token_backend. It defaults to "eager", which leaves the current implementation untouched. Setting it to "fused" swaps the reorder and the weighted restore for Triton kernels that touch each row once. Writing the reorder out as four passes shows that all of them are the same row gather, where perm maps an expert-major slot to its source row, inv is its inverse, and a negative index reads as zero: permute forward out[i] = tokens[perm[i]] gather by perm permute backward dtok[j] = dout[inv[j]] gather by inv unpermute forward out[j] = expert_out[inv[j]] gather by inv unpermute backward dexp[i] = dout[perm[i]] gather by perm So one kernel serves all four. perm is injective on real rows, so no pass needs atomics, and carrying inv costs 4 bytes per row against the 2 * hidden bytes per row that the padded copy and the zero-filled scatter buffer cost today. The restore goes straight from [tokens * top_k, hidden] to [batch, seq, hidden], weighting each row and reducing over top-k in registers. It keeps the eager dtype discipline: FP32 product and accumulation, one cast on the way out. Its backward produces both the expert-output gradient and the routing-score gradient, reducing the score gradient over the hidden dimension per token so that no cross-program atomic is needed. The collectives, the router and the grouped GEMM are unchanged, so a measured difference between the two backends belongs to the local token engine alone. "fused" is rejected rather than quietly ignored wherever it would have nothing to replace or would change semantics: folded tensor parallelism, which restores combined tokens from assignment metadata; an explicit combine_impl="legacy_bmm"; a resolved score_apply other than "post"; and non-CUDA, non-Triton or non-bf16/fp16 execution, which is checked once before any collective so that ranks fail together instead of stalling. A run that asked for "fused" and silently got "eager" would otherwise report the difference between a backend and itself. The alignment and index generation that both backends share moves to generate_local_expert_permute_indices, so the two cannot drift apart. Signed-off-by: yh0903 --- deepspeed/module_inject/auto_ep_config.py | 32 ++ deepspeed/module_inject/auto_ep_layer.py | 86 ++-- .../module_inject/auto_ep_presets/base.py | 1 + deepspeed/moe/autoep_fused_token_ops.py | 455 ++++++++++++++++++ deepspeed/moe/ep_kernels.py | 62 +++ docs/_pages/config-json.md | 6 + docs/code-docs/source/autoep.rst | 42 ++ .../v1/moe/test_autoep_fused_token_ops.py | 181 +++++++ tests/unit/v1/moe/test_autoep_unit.py | 43 ++ 9 files changed, 860 insertions(+), 48 deletions(-) create mode 100644 deepspeed/moe/autoep_fused_token_ops.py create mode 100644 tests/unit/v1/moe/test_autoep_fused_token_ops.py diff --git a/deepspeed/module_inject/auto_ep_config.py b/deepspeed/module_inject/auto_ep_config.py index 2841b5317f8c..c068f4f73cbd 100644 --- a/deepspeed/module_inject/auto_ep_config.py +++ b/deepspeed/module_inject/auto_ep_config.py @@ -58,6 +58,7 @@ def parse_autoep_config(param_dict: dict) -> AutoEPConfig: config.route_scale = param_dict.get("route_scale", 1.0) config.score_apply = param_dict.get("score_apply", "auto") config.combine_impl = param_dict.get("combine_impl", "auto") + config.local_token_backend = param_dict.get("local_token_backend", "eager") config.num_expert_groups = param_dict.get("num_expert_groups", None) config.num_limited_groups = param_dict.get("num_limited_groups", None) config.score_func = param_dict.get("score_func", "auto") @@ -118,6 +119,28 @@ def validate_autoep_config( if not config.enabled: return + # Validate local_token_backend + valid_local_token_backend = ("eager", "fused") + if config.local_token_backend not in valid_local_token_backend: + raise ValueError(f"local_token_backend must be one of {valid_local_token_backend}, " + f"got '{config.local_token_backend}'") + + # The fused engine only replaces the reorder and weighted restore that the + # standard expert-parallel path runs. Where it has nothing to replace, say so + # instead of running eager under a config that asked for fused: a benchmark + # that believes it measured the fused path would otherwise report noise. + if config.local_token_backend == "fused": + if config.expert_tensor_parallel_size > 1: + raise ValueError('local_token_backend="fused" does not support folded tensor parallelism ' + f"(expert_tensor_parallel_size={config.expert_tensor_parallel_size}), which restores " + "combined tokens from assignment metadata instead of the weighted reduction the fused " + 'engine implements. Set expert_tensor_parallel_size to 1, or local_token_backend to ' + '"eager".') + if config.combine_impl == "legacy_bmm": + raise ValueError('local_token_backend="fused" implements the weighted-sum reduction, so it cannot honor ' + 'combine_impl="legacy_bmm". Leave combine_impl unset, or set local_token_backend to ' + '"eager" to keep the legacy reduction for model-family verification.') + folding_spec = build_folding_spec( world_size=world_size, pp_size=pp_size, @@ -272,6 +295,15 @@ def validate_autoep_post_detection( return for spec in specs: + # The fused weighted restore folds the routing weight into the top-k + # reduction, which only exists when scores are applied after the experts. + if config.local_token_backend == "fused": + resolved_score_apply = config.score_apply if config.score_apply != "auto" else spec.score_apply + if resolved_score_apply != "post": + raise ValueError(f'local_token_backend="fused" requires score_apply="post", but layer ' + f"'{spec.moe_module_name}' resolved score_apply=\"{resolved_score_apply}\". " + 'Set local_token_backend to "eager".') + # ep_size must not exceed num_experts if config.autoep_size > spec.num_experts: valid_divisors = _divisors(spec.num_experts) diff --git a/deepspeed/module_inject/auto_ep_layer.py b/deepspeed/module_inject/auto_ep_layer.py index 1cb7696064d6..4639df2053e2 100644 --- a/deepspeed/module_inject/auto_ep_layer.py +++ b/deepspeed/module_inject/auto_ep_layer.py @@ -22,6 +22,7 @@ from deepspeed.module_inject.auto_ep_config import AutoEPConfig, MoELayerSpec, resolve_autoep_config_defaults from deepspeed.module_inject.auto_ep_folding import mark_autoep_folding_router_parameter from deepspeed.utils import logger +from deepspeed.moe import autoep_fused_token_ops as fused_token_ops from deepspeed.moe.ep_router import TokenChoiceTopKRouter from deepspeed.moe.ep_count import count_tokens_per_expert from deepspeed.moe.ep_experts import GroupedExperts @@ -241,44 +242,14 @@ def permute_by_local_expert( aligned_counts: [E_local] aligned token counts per expert (for expert computation) n_tokens: original token count before padding (for unpermute) """ - from deepspeed.moe.ep_kernels import generate_permute_indices, TOKEN_GROUP_ALIGN_SIZE_M - - if local_counts.ndim == 1: - # [E_local]: already aggregated over sources (ep_degree=1) - ep_degree = 1 - num_local_experts = local_counts.shape[0] - local_counts_flat = local_counts - elif local_counts.ndim == 2: - # [ep_size, E_local]: preserve per-source layout for correct regrouping - ep_degree, num_local_experts = local_counts.shape - local_counts_flat = local_counts.reshape(-1) - else: - raise ValueError( - f"local_counts must have shape [E_local] or [ep_degree, E_local], got {tuple(local_counts.shape)}") + from deepspeed.moe.ep_kernels import generate_local_expert_permute_indices n_tokens = tokens.shape[0] - alignment = TOKEN_GROUP_ALIGN_SIZE_M - - # Compute padded max length - x_padded_per_expert = n_tokens + num_local_experts * alignment - padded_max_len = ((x_padded_per_expert + alignment - 1) // alignment) * alignment - - # Use the pure-PyTorch path for host tensors. The CPU accelerator reports - # CPU tensors as "on accelerator", but Triton still requires a GPU driver. - use_cpu = tokens.device.type == "cpu" - counts_for_permute = local_counts_flat.cpu() if use_cpu else local_counts_flat - with torch.no_grad(): - permuted_indices, m_sizes, _offsets = generate_permute_indices( - counts_for_permute, - num_local_experts, - ep_degree, - padded_max_len, - alignment, - use_cpu=use_cpu, - ) - if not use_cpu: - permuted_indices = permuted_indices.to(tokens.device) - m_sizes = m_sizes.to(tokens.device) + permuted_indices, m_sizes = generate_local_expert_permute_indices( + n_tokens=n_tokens, + local_counts=local_counts, + device=tokens.device, + ) # Add padding row for out-of-bounds indices (index n_tokens -> zero row) tokens_padded = torch.vstack((tokens, tokens.new_zeros((tokens.shape[-1], )))) @@ -377,6 +348,8 @@ def __init__( self.top_k = spec.top_k self.score_apply = resolve_score_apply_mode(spec, config.score_apply) self.combine_impl = resolve_combine_impl(config.combine_impl) + self.local_token_backend = config.local_token_backend + self._fused_backend_checked = False route_norm = spec.route_norm if config.route_norm is None else config.route_norm self.ep_size = ep_size self.ep_rank = ep_rank @@ -550,6 +523,10 @@ def set_deepspeed_parallelism( if folding_group_handles is not None: self.folding_group_handles = folding_group_handles + if self.local_token_backend == "fused" and folding_group_handles.spec.tp_size > 1: + raise ValueError('local_token_backend="fused" does not support folded tensor parallelism ' + f"(expert_tensor_parallel_size={folding_group_handles.spec.tp_size}). Set " + 'expert_tensor_parallel_size to 1, or local_token_backend to "eager".') self.ep_group_name = folding_group_handles.ep_group_name self.ep_group = folding_group_handles.ep_group self.tp_group = folding_group_handles.tp_group @@ -572,6 +549,17 @@ def set_deepspeed_parallelism( ) self.ep_group = groups._get_expert_parallel_group(self.ep_group_name) + def _run_local_experts(self, rows: torch.Tensor, local_counts: torch.Tensor) -> torch.Tensor: + """Group rows by local expert, run the grouped GEMM, and undo the grouping.""" + if self.local_token_backend == "fused": + reordered, reorder_context = fused_token_ops.fused_permute_by_local_expert(rows, local_counts) + expert_output = self.experts(reordered, reorder_context.aligned_counts) + return fused_token_ops.fused_unpermute_by_local_expert(expert_output, reorder_context) + + reordered, perm_indices, aligned_counts, n_tokens = permute_by_local_expert(rows, local_counts) + expert_output = self.experts(reordered, aligned_counts) + return unpermute_by_local_expert(expert_output, perm_indices, n_tokens) + def forward( self, hidden_states: torch.Tensor, @@ -588,6 +576,10 @@ def forward( bsz, seqlen, hdim = hidden_states.shape x = hidden_states.reshape(-1, hdim) # [T, H] + if self.local_token_backend == "fused" and not self._fused_backend_checked: + fused_token_ops.assert_supported(x, score_apply=self.score_apply) + self._fused_backend_checked = True + # Router ro: RouterOutput = RouterOutput(*self.router(x, self.expert_bias)) @@ -652,12 +644,7 @@ def forward( if self.ep_size == 1: # No AllToAll needed - local computation only - local_counts = ro.num_tokens_per_expert - - routed_input_permuted, perm_indices, aligned_counts, n_tokens = permute_by_local_expert( - routed_input, local_counts) - expert_output = self.experts(routed_input_permuted, aligned_counts) - expert_output = unpermute_by_local_expert(expert_output, perm_indices, n_tokens) + expert_output = self._run_local_experts(routed_input, ro.num_tokens_per_expert) else: # EP dispatch/compute/combine if folded_tp: @@ -679,12 +666,7 @@ def forward( ) routed_input = _AllToAllV.apply(self.ep_group, routed_input, plan.input_splits, plan.output_splits) - - routed_input, perm_indices, aligned_counts, n_tokens = permute_by_local_expert( - routed_input, plan.local_counts_by_source) - expert_output = self.experts(routed_input, aligned_counts) - expert_output = unpermute_by_local_expert(expert_output, perm_indices, n_tokens) - + expert_output = self._run_local_experts(routed_input, plan.local_counts_by_source) expert_output = _AllToAllV.apply(self.ep_group, expert_output, plan.output_splits, plan.input_splits) if folded_tp: @@ -693,6 +675,14 @@ def forward( tp_group=self.tp_group, validate_coverage=self.validate_folding_routing).reshape(bsz, seqlen, hdim) self._last_folding_dispatch_counters = dispatch_counters(restore_ctx) + elif self.local_token_backend == "fused": + output = fused_token_ops.fused_weighted_restore( + expert_output, + top_scores=ro.top_scores, + token_indices_sorted=token_indices_sorted, + top_k=self.top_k, + shape=(bsz, seqlen, hdim), + ) else: output = combine_from_routed( expert_output, diff --git a/deepspeed/module_inject/auto_ep_presets/base.py b/deepspeed/module_inject/auto_ep_presets/base.py index 7f214210aed9..429f2f7ee2d7 100644 --- a/deepspeed/module_inject/auto_ep_presets/base.py +++ b/deepspeed/module_inject/auto_ep_presets/base.py @@ -110,6 +110,7 @@ class AutoEPConfig: route_scale: float = 1.0 score_apply: Literal["auto", "pre", "post"] = "auto" combine_impl: Literal["auto", "weighted_sum", "legacy_bmm"] = "auto" + local_token_backend: Literal["eager", "fused"] = "eager" num_expert_groups: int | None = None num_limited_groups: int | None = None score_func: Literal["auto", "softmax", "sigmoid"] = "auto" diff --git a/deepspeed/moe/autoep_fused_token_ops.py b/deepspeed/moe/autoep_fused_token_ops.py new file mode 100644 index 000000000000..4af8adc046b2 --- /dev/null +++ b/deepspeed/moe/autoep_fused_token_ops.py @@ -0,0 +1,455 @@ +# Copyright (c) DeepSpeed Team. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team +"""Fused GPU-local token movement for AutoEP. + +The eager path moves routed rows around the two expert all-to-alls with +general-purpose tensor ops: it materializes a padded copy of the token matrix, +gathers through an advanced index, scatters expert outputs into a zero-filled +buffer, and builds a ``[T, K, H]`` FP32 intermediate only to apply routing +weights and reduce over top-k. + +This module replaces that sequence with Triton kernels that touch each row once. +It deliberately leaves the collectives, the router and the grouped GEMM alone, so +that a measured difference is attributable to the local token engine. + +The reorder is expressed entirely as row gathers. Writing the four passes out +shows why only one kernel is needed, where ``perm`` maps an expert-major slot to +its source row and ``inv`` is its inverse: + + permute forward out[i] = tokens[perm[i]] gather by perm + permute backward dtok[j] = dout[inv[j]] gather by inv + unpermute forward out[j] = expert_out[inv[j]] gather by inv + unpermute backward dexp[i] = dout[perm[i]] gather by perm + +``perm`` is injective on real rows, so no pass needs atomics, and carrying +``inv`` (4 bytes per row) removes both the padded input copy and the zero-filled +scatter buffer (``2 * H`` bytes per row) that the eager path allocates. +""" + +from __future__ import annotations + +from typing import NamedTuple + +import torch + +try: + import triton + import triton.language as tl + + _TRITON_AVAILABLE = True +except ImportError: + _TRITON_AVAILABLE = False + +# The grouped GEMM consumes the reordered rows, so the fused path supports the +# dtypes it is built for rather than silently widening them. +SUPPORTED_ROW_DTYPES = (torch.bfloat16, torch.float16) + +_MAX_BLOCK_HIDDEN = 512 +_INVERT_INDEX_BLOCK = 256 +# The restore kernels hold a [slots, BLOCK_H] FP32 block live, so the hidden tile +# shrinks as top-k grows to keep that block in registers instead of spilling. +_MAX_BLOCK_ELEMENTS = 2048 + +if _TRITON_AVAILABLE: + + @triton.jit + def _gather_rows_kernel( + source_ptr, + index_ptr, + out_ptr, + hidden, + source_stride, + out_stride, + BLOCK_H: tl.constexpr, + ): + out_row = tl.program_id(0) + hidden_offsets = tl.program_id(1) * BLOCK_H + tl.arange(0, BLOCK_H) + hidden_mask = hidden_offsets < hidden + + source_row = tl.load(index_ptr + out_row).to(tl.int64) + # A negative index marks an alignment-padding slot. The eager path spelled + # that as an appended zero row; here it is a masked-off load. + row_valid = source_row >= 0 + safe_row = tl.where(row_valid, source_row, 0) + + values = tl.load( + source_ptr + safe_row * source_stride + hidden_offsets, + mask=hidden_mask & row_valid, + other=0.0, + ) + tl.store(out_ptr + out_row * out_stride + hidden_offsets, values, mask=hidden_mask) + + @triton.jit + def _invert_index_kernel( + index_ptr, + inverse_ptr, + num_indices, + num_inverse_rows, + BLOCK: tl.constexpr, + ): + offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + in_range = offsets < num_indices + + targets = tl.load(index_ptr + offsets, mask=in_range, other=-1).to(tl.int64) + writable = in_range & (targets >= 0) & (targets < num_inverse_rows) + tl.store(inverse_ptr + tl.where(writable, targets, 0), offsets.to(tl.int32), mask=writable) + + @triton.jit + def _weighted_restore_forward_kernel( + rows_ptr, + inverse_ptr, + scores_ptr, + out_ptr, + hidden, + rows_stride, + scores_stride, + out_stride, + TOP_K: tl.constexpr, + K_PADDED: tl.constexpr, + BLOCK_H: tl.constexpr, + ): + token = tl.program_id(0) + hidden_offsets = tl.program_id(1) * BLOCK_H + tl.arange(0, BLOCK_H) + hidden_mask = hidden_offsets < hidden + + slots = tl.arange(0, K_PADDED) + slot_mask = slots < TOP_K + + source_rows = tl.load(inverse_ptr + token * TOP_K + slots, mask=slot_mask, other=-1).to(tl.int64) + row_valid = slot_mask & (source_rows >= 0) + safe_rows = tl.where(row_valid, source_rows, 0) + + scores = tl.load(scores_ptr + token * scores_stride + slots, mask=slot_mask, other=0.0).to(tl.float32) + + block_mask = row_valid[:, None] & hidden_mask[None, :] + values = tl.load( + rows_ptr + safe_rows[:, None] * rows_stride + hidden_offsets[None, :], + mask=block_mask, + other=0.0, + ).to(tl.float32) + + # FP32 product and reduction with a single cast on the way out, matching + # the dtype discipline of the eager weighted sum. + weighted = tl.sum(values * scores[:, None], axis=0) + tl.store( + out_ptr + token * out_stride + hidden_offsets, + weighted.to(out_ptr.dtype.element_ty), + mask=hidden_mask, + ) + + @triton.jit + def _weighted_restore_backward_kernel( + grad_out_ptr, + rows_ptr, + inverse_ptr, + scores_ptr, + grad_rows_ptr, + grad_scores_ptr, + hidden, + grad_out_stride, + rows_stride, + scores_stride, + grad_rows_stride, + grad_scores_stride, + TOP_K: tl.constexpr, + K_PADDED: tl.constexpr, + BLOCK_H: tl.constexpr, + ): + token = tl.program_id(0) + + slots = tl.arange(0, K_PADDED) + slot_mask = slots < TOP_K + + source_rows = tl.load(inverse_ptr + token * TOP_K + slots, mask=slot_mask, other=-1).to(tl.int64) + row_valid = slot_mask & (source_rows >= 0) + safe_rows = tl.where(row_valid, source_rows, 0) + scores = tl.load(scores_ptr + token * scores_stride + slots, mask=slot_mask, other=0.0).to(tl.float32) + + grad_rows_dtype = grad_rows_ptr.dtype.element_ty + # One token per program, so the score gradient reduces over the hidden + # dimension in registers instead of through a cross-program atomic. + score_partials = tl.zeros([K_PADDED, BLOCK_H], dtype=tl.float32) + + for hidden_start in range(0, hidden, BLOCK_H): + hidden_offsets = hidden_start + tl.arange(0, BLOCK_H) + hidden_mask = hidden_offsets < hidden + block_mask = row_valid[:, None] & hidden_mask[None, :] + + upstream = tl.load( + grad_out_ptr + token * grad_out_stride + hidden_offsets, + mask=hidden_mask, + other=0.0, + ).to(tl.float32) + + values = tl.load( + rows_ptr + safe_rows[:, None] * rows_stride + hidden_offsets[None, :], + mask=block_mask, + other=0.0, + ).to(tl.float32) + score_partials += values * upstream[None, :] + + tl.store( + grad_rows_ptr + safe_rows[:, None] * grad_rows_stride + hidden_offsets[None, :], + (upstream[None, :] * scores[:, None]).to(grad_rows_dtype), + mask=block_mask, + ) + + grad_scores = tl.sum(score_partials, axis=1) + tl.store( + grad_scores_ptr + token * grad_scores_stride + slots, + grad_scores.to(grad_scores_ptr.dtype.element_ty), + mask=slot_mask, + ) + + +class FusedReorderContext(NamedTuple): + """Index metadata shared by the reorder forward and backward passes.""" + + permutation: torch.Tensor # [N_padded] int32; -1 marks an alignment-padding slot + inverse: torch.Tensor # [n_tokens] int32; -1 marks a row no slot claimed + aligned_counts: torch.Tensor # [E_local] int32 row counts for the grouped GEMM + n_tokens: int + + +def is_available() -> bool: + """Whether this build can run the fused local token engine at all.""" + return _TRITON_AVAILABLE + + +def assert_supported(rows: torch.Tensor, *, score_apply: str) -> None: + """Reject configurations the fused path does not implement. + + Checked before any collective runs: a rank that raised while its peers + proceeded would turn a clear error into a hang. + """ + if not _TRITON_AVAILABLE: + raise RuntimeError('expert_parallel.local_token_backend="fused" needs Triton, which is not installed in ' + 'this environment. Install Triton, or set local_token_backend to "eager".') + if rows.device.type != "cuda": + raise RuntimeError('expert_parallel.local_token_backend="fused" runs CUDA kernels but this layer is on ' + f'device "{rows.device.type}". Set local_token_backend to "eager" to run here.') + if rows.dtype not in SUPPORTED_ROW_DTYPES: + raise RuntimeError('expert_parallel.local_token_backend="fused" supports bfloat16 and float16 rows, got ' + f'{rows.dtype}. Set local_token_backend to "eager", or train in bf16/fp16.') + if score_apply != "post": + raise RuntimeError('expert_parallel.local_token_backend="fused" implements the post-expert weighted ' + f'restore, but this layer resolved score_apply="{score_apply}". Set local_token_backend ' + 'to "eager".') + + +def _block_hidden(hidden: int, slots: int = 1) -> int: + """Pick a power-of-two hidden tile that fits alongside ``slots`` rows of FP32. + + The floor keeps the budget honest for top-k values far wider than any real + router, so the tile shrinks rather than overrunning the element budget. + """ + budget = max(16, _MAX_BLOCK_ELEMENTS // slots) + return min(_MAX_BLOCK_HIDDEN, budget, max(16, triton.next_power_of_2(hidden))) + + +def _padded_top_k(top_k: int) -> int: + """Round top-k up to a power of two, which ``tl.arange`` requires.""" + return max(2, triton.next_power_of_2(top_k)) + + +def _gather_rows(source: torch.Tensor, index: torch.Tensor, num_out_rows: int) -> torch.Tensor: + """Gather the ``source`` rows named by ``index``, reading a negative index as zero.""" + hidden = source.shape[-1] + out = torch.empty((num_out_rows, hidden), dtype=source.dtype, device=source.device) + if num_out_rows == 0 or hidden == 0: + return out + + source = source.contiguous() + block_hidden = _block_hidden(hidden) + grid = (num_out_rows, triton.cdiv(hidden, block_hidden)) + _gather_rows_kernel[grid]( + source, + index, + out, + hidden, + source.stride(0), + out.stride(0), + BLOCK_H=block_hidden, + ) + return out + + +def _invert_index(index: torch.Tensor, num_inverse_rows: int) -> torch.Tensor: + """Invert a row permutation, leaving -1 wherever no slot claimed a row.""" + inverse = torch.full((num_inverse_rows, ), -1, dtype=torch.int32, device=index.device) + num_indices = index.numel() + if num_indices == 0 or num_inverse_rows == 0: + return inverse + + grid = (triton.cdiv(num_indices, _INVERT_INDEX_BLOCK), ) + _invert_index_kernel[grid]( + index.contiguous(), + inverse, + num_indices, + num_inverse_rows, + BLOCK=_INVERT_INDEX_BLOCK, + ) + return inverse + + +class _FusedReorder(torch.autograd.Function): + """Gather routed rows into expert-major, alignment-padded order.""" + + @staticmethod + def forward(ctx, tokens, permutation, inverse, n_padded): + ctx.save_for_backward(inverse) + ctx.n_tokens = tokens.shape[0] + return _gather_rows(tokens, permutation, n_padded) + + @staticmethod + def backward(ctx, grad_out): + inverse, = ctx.saved_tensors + return _gather_rows(grad_out.contiguous(), inverse, ctx.n_tokens), None, None, None + + +class _FusedInverseReorder(torch.autograd.Function): + """Restore source-major order and drop the alignment padding.""" + + @staticmethod + def forward(ctx, expert_output, permutation, inverse): + ctx.save_for_backward(permutation) + ctx.n_padded = expert_output.shape[0] + return _gather_rows(expert_output, inverse, inverse.numel()) + + @staticmethod + def backward(ctx, grad_out): + permutation, = ctx.saved_tensors + return _gather_rows(grad_out.contiguous(), permutation, ctx.n_padded), None, None + + +class _FusedWeightedRestore(torch.autograd.Function): + """Weight rows by their routing score and reduce over top-k in one pass.""" + + @staticmethod + def forward(ctx, combined_rows, top_scores, inverse, top_k): + combined_rows = combined_rows.contiguous() + n_tokens, hidden = top_scores.shape[0], combined_rows.shape[-1] + output = torch.empty((n_tokens, hidden), dtype=combined_rows.dtype, device=combined_rows.device) + + ctx.save_for_backward(combined_rows, top_scores, inverse) + ctx.top_k = top_k + if n_tokens == 0 or hidden == 0: + return output + + k_padded = _padded_top_k(top_k) + block_hidden = _block_hidden(hidden, slots=k_padded) + grid = (n_tokens, triton.cdiv(hidden, block_hidden)) + _weighted_restore_forward_kernel[grid]( + combined_rows, + inverse, + top_scores, + output, + hidden, + combined_rows.stride(0), + top_scores.stride(0), + output.stride(0), + TOP_K=top_k, + K_PADDED=k_padded, + BLOCK_H=block_hidden, + ) + return output + + @staticmethod + def backward(ctx, grad_output): + combined_rows, top_scores, inverse = ctx.saved_tensors + grad_output = grad_output.contiguous() + + # Every row is claimed by exactly one slot, which fused_weighted_restore + # checks by shape before building the inverse, so both gradients are + # written in full and neither buffer needs pre-zeroing. + grad_rows = torch.empty_like(combined_rows) + grad_scores = torch.empty_like(top_scores) + + n_tokens, hidden = top_scores.shape[0], combined_rows.shape[-1] + if n_tokens == 0 or hidden == 0: + # Nothing is reduced, so the score gradient is zero rather than + # whatever the uninitialized buffer happened to hold. + return grad_rows, torch.zeros_like(top_scores), None, None + + k_padded = _padded_top_k(ctx.top_k) + _weighted_restore_backward_kernel[(n_tokens, )]( + grad_output, + combined_rows, + inverse, + top_scores, + grad_rows, + grad_scores, + hidden, + grad_output.stride(0), + combined_rows.stride(0), + top_scores.stride(0), + grad_rows.stride(0), + grad_scores.stride(0), + TOP_K=ctx.top_k, + K_PADDED=k_padded, + BLOCK_H=_block_hidden(hidden, slots=k_padded), + ) + return grad_rows, grad_scores, None, None + + +def fused_permute_by_local_expert( + tokens: torch.Tensor, + local_counts: torch.Tensor, +) -> tuple[torch.Tensor, FusedReorderContext]: + """Reorder routed rows into expert-contiguous, alignment-padded order. + + Fused counterpart of ``permute_by_local_expert``. It shares that function's + index generation and produces the same rows, without materializing the padded + copy of ``tokens`` that the eager advanced index needs. + """ + from deepspeed.moe.ep_kernels import generate_local_expert_permute_indices + + permutation, aligned_counts = generate_local_expert_permute_indices( + n_tokens=tokens.shape[0], + local_counts=local_counts, + device=tokens.device, + ) + inverse = _invert_index(permutation, tokens.shape[0]) + context = FusedReorderContext( + permutation=permutation, + inverse=inverse, + aligned_counts=aligned_counts, + n_tokens=tokens.shape[0], + ) + permuted = _FusedReorder.apply(tokens, permutation, inverse, permutation.numel()) + return permuted, context + + +def fused_unpermute_by_local_expert( + expert_output: torch.Tensor, + context: FusedReorderContext, +) -> torch.Tensor: + """Reverse :func:`fused_permute_by_local_expert` and strip the padding.""" + return _FusedInverseReorder.apply(expert_output.contiguous(), context.permutation, context.inverse) + + +def fused_weighted_restore( + combined_rows: torch.Tensor, + top_scores: torch.Tensor, + token_indices_sorted: torch.Tensor, + top_k: int, + shape: tuple[int, int, int], +) -> torch.Tensor: + """Weight combined rows by their routing scores and reduce over top-k. + + Fused counterpart of ``combine_from_routed`` for ``score_apply="post"``. It + goes straight from ``[T * K, H]`` to ``[B, S, H]``, so neither the scattered + assignment buffer nor the ``[T, K, H]`` FP32 intermediate is allocated. + """ + bsz, seqlen, hidden = shape + n_tokens = bsz * seqlen + expected_rows = n_tokens * top_k + if combined_rows.shape[0] != expected_rows: + raise RuntimeError(f"fused weighted restore expects one row per assignment: {expected_rows} rows for " + f"{n_tokens} tokens at top_k={top_k}, got {combined_rows.shape[0]}.") + + inverse = _invert_index(token_indices_sorted, expected_rows) + output = _FusedWeightedRestore.apply(combined_rows, top_scores.contiguous(), inverse, top_k) + return output.reshape(bsz, seqlen, hidden) diff --git a/deepspeed/moe/ep_kernels.py b/deepspeed/moe/ep_kernels.py index 6b3da7335e88..6818d858c81a 100644 --- a/deepspeed/moe/ep_kernels.py +++ b/deepspeed/moe/ep_kernels.py @@ -250,6 +250,68 @@ def generate_permute_indices( return permuted_indices, m_sizes, m_offsets.to(torch.int32) +# =================================================================== +# generate_local_expert_permute_indices +# =================================================================== + + +def generate_local_expert_permute_indices( + n_tokens: int, + local_counts: torch.Tensor, + device: torch.device, +) -> tuple: + """Build the expert-major row order for ``n_tokens`` received rows. + + Args: + n_tokens: Number of received rows, before alignment padding. + local_counts: ``(E_local,)`` when the counts are already aggregated over + sources, or ``(ep_degree, E_local)`` to keep the per-source layout + that correct regrouping depends on. + device: Device the returned tensors should live on. + + Returns: + Tuple of: + - permuted_indices: expert-major slot -> source row, ``-1`` where a + slot is alignment padding rather than a routed row. + - aligned_counts: per-expert row counts for the grouped GEMM. + """ + if local_counts.ndim == 1: + # [E_local]: already aggregated over sources (ep_degree=1) + ep_degree = 1 + num_local_experts = local_counts.shape[0] + local_counts_flat = local_counts + elif local_counts.ndim == 2: + # [ep_size, E_local]: preserve per-source layout for correct regrouping + ep_degree, num_local_experts = local_counts.shape + local_counts_flat = local_counts.reshape(-1) + else: + raise ValueError( + f"local_counts must have shape [E_local] or [ep_degree, E_local], got {tuple(local_counts.shape)}") + + alignment = TOKEN_GROUP_ALIGN_SIZE_M + x_padded_per_expert = n_tokens + num_local_experts * alignment + padded_max_len = _round_up(x_padded_per_expert, alignment) + + # Use the pure-PyTorch path for host tensors. The CPU accelerator reports + # CPU tensors as "on accelerator", but Triton still requires a GPU driver. + use_cpu = device.type == "cpu" + counts_for_permute = local_counts_flat.cpu() if use_cpu else local_counts_flat + with torch.no_grad(): + permuted_indices, m_sizes, _offsets = generate_permute_indices( + counts_for_permute, + num_local_experts, + ep_degree, + padded_max_len, + alignment, + use_cpu=use_cpu, + ) + if not use_cpu: + permuted_indices = permuted_indices.to(device) + m_sizes = m_sizes.to(device) + + return permuted_indices, m_sizes + + # =================================================================== # _permute / _unpermute / indices_padding_wrapper # =================================================================== diff --git a/docs/_pages/config-json.md b/docs/_pages/config-json.md index ffc8d66e2dd7..6914986512f5 100644 --- a/docs/_pages/config-json.md +++ b/docs/_pages/config-json.md @@ -994,6 +994,12 @@ smoke coverage used for this AutoEP surface produced the following version gates | -------------------------------------------------------------------------------------------------------------- | -------- | | When to apply router scores: `"pre"` (before experts), `"post"` (during combine), or `"auto"` (from preset). | `"auto"` | +***local_token_backend***: [string] + +| Description | Default | +| -------------------------------------------------------------------------------------------------------------- | -------- | +| GPU-local token movement backend. `"eager"` keeps the general-purpose reorder and combine. `"fused"` is experimental and replaces them with Triton kernels that group rows by expert, undo that grouping, and apply router scores while reducing over top-k, without materializing the padded token copy or the `[tokens, top_k, hidden]` FP32 intermediate. Collectives, routing and the grouped GEMM are unchanged. `"fused"` requires CUDA, Triton, bfloat16/float16 activations, `expert_tensor_parallel_size=1`, and a resolved `score_apply="post"`; it is rejected rather than silently ignored when any of those does not hold. | `"eager"` | + ***route_norm***: [boolean] | Description | Default | diff --git a/docs/code-docs/source/autoep.rst b/docs/code-docs/source/autoep.rst index b7b7fa293d82..e960a7fc2c32 100644 --- a/docs/code-docs/source/autoep.rst +++ b/docs/code-docs/source/autoep.rst @@ -84,6 +84,48 @@ Weights-only/module-only Universal Checkpoint loads use the converted 4. Expert parameters are marked for expert-data-parallel gradient reduction; router and shared-expert parameters use standard data-parallel reduction. +**Fused local token engine (experimental):** + +``expert_parallel.local_token_backend`` selects how routed rows are moved around +the two expert all-to-alls. ``"eager"`` (default) keeps the existing +general-purpose tensor ops. ``"fused"`` replaces them with Triton kernels: + +.. code-block:: json + + { + "expert_parallel": { + "enabled": true, + "autoep_size": 16, + "preset_model": "qwen3_moe", + "local_token_backend": "fused" + } + } + +The fused engine groups rows by local expert, undoes that grouping, and applies +router scores while reducing over top-k. It reads and writes each row once, so +the padded copy of the token matrix, the zero-filled scatter buffer, and the +``[tokens, top_k, hidden]`` FP32 intermediate are never allocated. Routing +weights are still accumulated in FP32 and cast once, matching the eager result to +within the reduction order. + +The collectives, the router and the grouped GEMM are untouched, so a measured +difference between the two backends belongs to the local token engine alone. + +``"fused"`` is rejected, rather than quietly ignored, when it would have nothing +to replace or would change semantics: + +- ``expert_tensor_parallel_size`` greater than 1, which restores combined tokens + from assignment metadata instead of the weighted reduction; +- an explicit ``combine_impl="legacy_bmm"``, which selects the legacy reduction + kept for model-family verification; +- a resolved ``score_apply`` other than ``"post"``; +- activations that are not bfloat16 or float16, a non-CUDA device, or a build + without Triton. + +Failing fast matters for measurement: a run that asked for ``"fused"`` and +silently got ``"eager"`` would report the difference between a backend and +itself. + **Constraints:** - ``autoep_size`` must divide ``num_experts`` for all detected MoE layers. diff --git a/tests/unit/v1/moe/test_autoep_fused_token_ops.py b/tests/unit/v1/moe/test_autoep_fused_token_ops.py new file mode 100644 index 000000000000..6d376656563b --- /dev/null +++ b/tests/unit/v1/moe/test_autoep_fused_token_ops.py @@ -0,0 +1,181 @@ +# Copyright (c) DeepSpeed Team. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team +"""The fused local token engine against the eager reorder and weighted restore. + +The eager implementations are the reference: the fused engine is only worth +having if it is indistinguishable from them, so every assertion here compares the +two directly rather than against hand-written expectations. +""" + +import pytest +import torch + +from deepspeed.accelerator import get_accelerator +from deepspeed.moe import autoep_fused_token_ops as fused_ops +from deepspeed.module_inject.auto_ep_layer import ( + combine_from_routed, + permute_by_local_expert, + unpermute_by_local_expert, +) + + +def _fused_engine_available(): + accelerator = get_accelerator() + return (accelerator.is_available() and accelerator.device_name().startswith("cuda") and fused_ops.is_available()) + + +pytestmark = pytest.mark.skipif(not _fused_engine_available(), + reason="the fused local token engine needs CUDA and Triton") + +# Row counts per local expert, or per [source rank, local expert] where nested. +# They are the only thing that decides the reorder, so they carry every shape the +# engine has to survive: idle experts, one expert taking everything, and sources +# that contribute nothing to a given expert. +REORDER_CASES = { + "balanced": [8, 8, 8, 8], + "empty_experts": [0, 12, 0, 4], + "extreme_skew": [40, 0, 0, 0], + "per_source": [[3, 5], [7, 1]], + "ragged_per_source": [[0, 9], [5, 0], [2, 2]], +} + + +def _device(): + return get_accelerator().current_device_name() + + +def _counts(case): + return torch.tensor(REORDER_CASES[case], dtype=torch.int32, device=_device()) + + +@pytest.mark.parametrize("case", sorted(REORDER_CASES)) +@pytest.mark.parametrize("hidden", [64, 130]) +def test_fused_reorder_places_the_same_rows_as_eager(case, hidden): + counts = _counts(case) + tokens = torch.randn(int(counts.sum()), hidden, device=_device(), dtype=torch.bfloat16) + + eager_rows, _permutation, eager_counts, _n_tokens = permute_by_local_expert(tokens, counts) + fused_rows, context = fused_ops.fused_permute_by_local_expert(tokens, counts) + + # Pure data movement on both sides, so anything short of equality is a bug. + assert torch.equal(fused_rows, eager_rows) + assert torch.equal(context.aligned_counts, eager_counts) + + +def test_fused_reorder_handles_a_batch_no_expert_claimed(): + counts = torch.zeros(4, dtype=torch.int32, device=_device()) + tokens = torch.randn(0, 32, device=_device(), dtype=torch.bfloat16) + + eager_rows, _permutation, eager_counts, _n_tokens = permute_by_local_expert(tokens, counts) + fused_rows, context = fused_ops.fused_permute_by_local_expert(tokens, counts) + + assert torch.equal(fused_rows, eager_rows) + assert torch.equal(context.aligned_counts, eager_counts) + assert not fused_rows.any() + + +@pytest.mark.parametrize("case", sorted(REORDER_CASES)) +def test_fused_reorder_round_trip_matches_eager_including_gradients(case): + hidden = 96 + counts = _counts(case) + n_tokens = int(counts.sum()) + tokens = torch.randn(n_tokens, hidden, device=_device(), dtype=torch.bfloat16) + upstream = torch.randn(n_tokens, hidden, device=_device(), dtype=torch.bfloat16) + + eager_tokens = tokens.clone().requires_grad_(True) + eager_rows, permutation, _counts_out, n = permute_by_local_expert(eager_tokens, counts) + # Scaling by a power of two keeps the comparison exact while still putting a + # real op between the reorder and its inverse. + eager_output = unpermute_by_local_expert(eager_rows * 2.0, permutation, n) + + fused_tokens = tokens.clone().requires_grad_(True) + fused_rows, context = fused_ops.fused_permute_by_local_expert(fused_tokens, counts) + fused_output = fused_ops.fused_unpermute_by_local_expert(fused_rows * 2.0, context) + + assert torch.equal(fused_output, eager_output) + + eager_output.backward(upstream) + fused_output.backward(upstream) + assert torch.equal(fused_tokens.grad, eager_tokens.grad) + + +@pytest.mark.parametrize("top_k", [2, 4, 6, 8]) +@pytest.mark.parametrize("hidden", [128, 130]) +@pytest.mark.parametrize("score_dtype", [torch.float32, torch.bfloat16]) +def test_fused_weighted_restore_matches_eager_including_gradients(top_k, hidden, score_dtype): + device = _device() + num_tokens, num_experts = 24, 8 + generator = torch.Generator(device=device).manual_seed(20260824) + + selected_experts = torch.randint(0, num_experts, (num_tokens, top_k), device=device, generator=generator) + token_indices_sorted = torch.argsort(selected_experts.view(-1), stable=True) + # A restore that only had to undo the identity would not exercise anything. + assert not torch.equal(token_indices_sorted, torch.arange(num_tokens * top_k, device=device)) + + rows = torch.randn(num_tokens * top_k, hidden, device=device, dtype=torch.bfloat16, generator=generator) + scores = torch.rand(num_tokens, top_k, device=device, dtype=score_dtype, generator=generator) + upstream = torch.randn(1, num_tokens, hidden, device=device, dtype=torch.bfloat16, generator=generator) + + eager_rows = rows.clone().requires_grad_(True) + eager_scores = scores.clone().requires_grad_(True) + eager_output = combine_from_routed( + eager_rows, + top_scores=eager_scores, + token_indices_sorted=token_indices_sorted, + top_k=top_k, + score_apply="post", + combine_impl="weighted_sum", + shape=(1, num_tokens, hidden), + ) + + fused_rows = rows.clone().requires_grad_(True) + fused_scores = scores.clone().requires_grad_(True) + fused_output = fused_ops.fused_weighted_restore( + fused_rows, + top_scores=fused_scores, + token_indices_sorted=token_indices_sorted, + top_k=top_k, + shape=(1, num_tokens, hidden), + ) + + torch.testing.assert_close(fused_output, eager_output) + + eager_output.backward(upstream) + fused_output.backward(upstream) + + torch.testing.assert_close(fused_rows.grad, eager_rows.grad) + # The score gradient reduces over the hidden dimension, so the fused and eager + # summation orders differ even though both accumulate in FP32. That shows up + # in FP32 scores; a bfloat16 score rounds the difference away, and asking for + # FP32 precision there would fail on a rounding boundary rather than on a bug. + score_tolerance = {"rtol": 1e-4, "atol": 1e-5} if score_dtype == torch.float32 else {} + torch.testing.assert_close(fused_scores.grad, eager_scores.grad, **score_tolerance) + + +def test_fused_weighted_restore_requires_one_row_per_assignment(): + device = _device() + with pytest.raises(RuntimeError, match="one row per assignment"): + fused_ops.fused_weighted_restore( + torch.randn(10, 16, device=device, dtype=torch.bfloat16), + top_scores=torch.rand(4, 2, device=device), + token_indices_sorted=torch.arange(8, device=device), + top_k=2, + shape=(1, 4, 16), + ) + + +def test_fused_engine_names_what_it_cannot_run(): + device = _device() + supported = torch.randn(8, 16, device=device, dtype=torch.bfloat16) + fused_ops.assert_supported(supported, score_apply="post") + + with pytest.raises(RuntimeError, match="bfloat16 and float16"): + fused_ops.assert_supported(torch.randn(8, 16, device=device, dtype=torch.float32), score_apply="post") + + with pytest.raises(RuntimeError, match='score_apply="post"'): + fused_ops.assert_supported(supported, score_apply="pre") + + with pytest.raises(RuntimeError, match="CUDA kernels"): + fused_ops.assert_supported(torch.randn(8, 16, dtype=torch.bfloat16), score_apply="post") diff --git a/tests/unit/v1/moe/test_autoep_unit.py b/tests/unit/v1/moe/test_autoep_unit.py index d28eb96c8036..642bd11b7476 100644 --- a/tests/unit/v1/moe/test_autoep_unit.py +++ b/tests/unit/v1/moe/test_autoep_unit.py @@ -234,6 +234,49 @@ def test_validate_folding_routing_requires_boolean(self): tp_size=1, sp_size=1) + def test_local_token_backend_defaults_to_eager(self): + assert parse_autoep_config({}).local_token_backend == "eager" + assert parse_autoep_config({"enabled": True}).local_token_backend == "eager" + + def test_local_token_backend_rejects_unknown_value(self): + config = parse_autoep_config({"enabled": True, "local_token_backend": "triton"}) + with pytest.raises(ValueError, match="local_token_backend must be one of"): + validate_autoep_config(config, world_size=1, pp_size=1, tp_size=1, sp_size=1) + + def test_fused_local_token_backend_rejects_folded_tensor_parallelism(self): + config = parse_autoep_config({ + "enabled": True, + "autoep_size": 2, + "expert_tensor_parallel_size": 2, + "local_token_backend": "fused", + }) + with pytest.raises(ValueError, match="folded tensor parallelism"): + validate_autoep_config(config, world_size=4, pp_size=1, tp_size=2, sp_size=1) + + def test_fused_local_token_backend_rejects_explicit_legacy_bmm(self): + config = parse_autoep_config({ + "enabled": True, + "local_token_backend": "fused", + "combine_impl": "legacy_bmm", + }) + with pytest.raises(ValueError, match="legacy_bmm"): + validate_autoep_config(config, world_size=1, pp_size=1, tp_size=1, sp_size=1) + + @pytest.mark.parametrize("score_apply, spec_score_apply", [("auto", "pre"), ("pre", "post")]) + def test_fused_local_token_backend_requires_post_score_apply(self, score_apply, spec_score_apply): + config = parse_autoep_config({ + "enabled": True, + "local_token_backend": "fused", + "score_apply": score_apply, + }) + with pytest.raises(ValueError, match='requires score_apply="post"'): + validate_autoep_post_detection(config, [_make_spec(score_apply=spec_score_apply)]) + + def test_fused_local_token_backend_accepts_the_standard_path(self): + config = parse_autoep_config({"enabled": True, "autoep_size": 2, "local_token_backend": "fused"}) + validate_autoep_config(config, world_size=2, pp_size=1, tp_size=1, sp_size=1) + validate_autoep_post_detection(config, [_make_spec(num_experts=4, score_apply="post")]) + @pytest.mark.parametrize("value", UNSUPPORTED_LOAD_BALANCE_VALUES) def test_load_balance_coeff_rejected_at_parse(self, value): with pytest.raises(ValueError) as exc_info: From 0c93d4db532c7f8e724c5f1736187e6737e9ba79 Mon Sep 17 00:00:00 2001 From: yh0903 Date: Mon, 24 Aug 2026 17:18:49 -0700 Subject: [PATCH 2/9] Compare the fused and eager local token backends through a full step The kernel tests cover the reorder and the weighted restore in isolation. They cannot show that a step taken through the fused backend trains the same model, which is the property that decides whether the backend is safe to select. Add a parity test that runs one step through each backend from the same initial state and the same batch, then compares the loss, the block output, the input gradient, every trainable gradient by name, and the parameter delta the optimizer produced. Router and expert gradients reach the comparison by different routes through the restore, so they are asserted to be present rather than left to a bulk comparison that would pass on an empty set. The benchmarked configuration recomputes each MoE block in backward, so the expert-parallel case runs with and without activation checkpointing. A local case with autoep_size=1 covers the branch that skips the all-to-alls. Also fix the fail-fast assertion in the kernel tests: it matched on the required score_apply rather than the rejected one, so it failed against a correct message. Validated on an H100 node: 139 passed, none skipped, for the kernel and config tests, and all three parity cases passed. Signed-off-by: yh0903 --- tests/unit/v1/moe/test_autoep_fused_parity.py | 201 ++++++++++++++++++ .../v1/moe/test_autoep_fused_token_ops.py | 2 +- 2 files changed, 202 insertions(+), 1 deletion(-) create mode 100644 tests/unit/v1/moe/test_autoep_fused_parity.py diff --git a/tests/unit/v1/moe/test_autoep_fused_parity.py b/tests/unit/v1/moe/test_autoep_fused_parity.py new file mode 100644 index 000000000000..47ddcab5a60b --- /dev/null +++ b/tests/unit/v1/moe/test_autoep_fused_parity.py @@ -0,0 +1,201 @@ +# Copyright (c) DeepSpeed Team. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team +"""End-to-end parity between the eager and fused local token backends. + +The fused engine is only a layout change, so a step taken through it has to +produce the same loss, the same gradients on every trainable tensor, and the same +parameter update as the eager path it replaces. +""" + +import functools + +import deepspeed +import pytest +import torch + +from deepspeed.accelerator import get_accelerator +from deepspeed.moe import autoep_fused_token_ops as fused_ops +from deepspeed.module_inject.auto_ep_layer import AutoEPMoELayer +from deepspeed.utils import safe_get_full_grad +from unit.common import DistributedTest +from unit.v1.moe.autoep_test_utils import ( + MockMoETransformer, + engine_input_dtype, + mixed_precision_config, + seed_everything, +) + +HIDDEN_SIZE = 64 +SEQ_LEN = 16 +NUM_EXPERTS = 4 + +# Both backends run the same collectives, the same router and the same grouped +# GEMM; they differ only in the order the top-k reduction accumulates. That +# survives a full step as a last-few-bits difference, not a structural one, so +# the tolerance stays far tighter than a wrong permutation could hide behind. +PARITY_TOLERANCE = {"rtol": 1e-2, "atol": 1e-3} + + +def _fused_engine_available(): + accelerator = get_accelerator() + return (accelerator.is_available() and accelerator.device_name().startswith("cuda") and fused_ops.is_available()) + + +pytestmark = pytest.mark.skipif(not _fused_engine_available(), + reason="the fused local token engine needs CUDA and Triton") + + +def _config(local_token_backend, ep_size): + return { + **mixed_precision_config(), + "train_micro_batch_size_per_gpu": 1, + "gradient_clipping": 0.0, + "optimizer": { + "type": "AdamW", + "params": { + "lr": 1e-3, + "betas": [0.9, 0.999], + "eps": 1e-8, + }, + }, + "expert_parallel": { + "enabled": True, + "autoep_size": ep_size, + "preset_model": "mixtral", + "load_balance_coeff": None, + "local_token_backend": local_token_backend, + }, + } + + +def _build_engine(local_token_backend, ep_size, reference_state, seed): + seed_everything(seed) + model = MockMoETransformer(num_layers=2, + num_experts=NUM_EXPERTS, + hidden_size=HIDDEN_SIZE, + intermediate_size=2 * HIDDEN_SIZE) + model.load_state_dict(reference_state) + engine, _, _, _ = deepspeed.initialize(model=model, config=_config(local_token_backend, ep_size)) + return engine + + +def _checkpoint_moe_layers(engine): + """Recompute each MoE block in backward, as the benchmarked runs do.""" + for module in engine.module.modules(): + if isinstance(module, AutoEPMoELayer): + module.forward = functools.partial(torch.utils.checkpoint.checkpoint, module.forward, use_reentrant=False) + + +def _named_gradients(engine): + gradients = {} + for name, param in engine.module.named_parameters(): + if not param.requires_grad: + continue + grad = safe_get_full_grad(param) + if grad is not None: + gradients[name] = grad.detach().float().cpu().clone() + return gradients + + +def _parameters(engine): + return { + name: param.detach().float().cpu().clone() + for name, param in engine.module.named_parameters() if param.requires_grad + } + + +def _take_one_step(engine, seed, *, checkpoint_activations): + if checkpoint_activations: + _checkpoint_moe_layers(engine) + + generator = torch.Generator().manual_seed(seed) + batch = torch.randn((1, SEQ_LEN, HIDDEN_SIZE), generator=generator, dtype=torch.float32) + batch = batch.to(engine.device, dtype=engine_input_dtype(engine)).requires_grad_(True) + + before = _parameters(engine) + output = engine(batch) + loss = output.float().pow(2).mean() + engine.backward(loss) + + gradients = _named_gradients(engine) + input_grad = batch.grad.detach().float().cpu().clone() + engine.step() + + delta = {name: _parameters(engine)[name] - value for name, value in before.items()} + return { + "loss": loss.detach().float().cpu().clone(), + "output": output.detach().float().cpu().clone(), + "input_grad": input_grad, + "gradients": gradients, + "delta": delta, + } + + +def _assert_step_matches(fused, eager): + torch.testing.assert_close(fused["loss"], eager["loss"], **PARITY_TOLERANCE) + torch.testing.assert_close(fused["output"], eager["output"], **PARITY_TOLERANCE) + torch.testing.assert_close(fused["input_grad"], eager["input_grad"], **PARITY_TOLERANCE) + + assert fused["gradients"], "no gradients were captured, so the comparison would be vacuous" + assert set(fused["gradients"]) == set(eager["gradients"]) + # Router and expert gradients travel different routes through the fused + # restore, so they are named rather than left to a bulk comparison. + assert any(".router." in name for name in fused["gradients"]), "no router gradient was captured" + assert any(".experts.w" in name for name in fused["gradients"]), "no expert gradient was captured" + + for name in sorted(eager["gradients"]): + torch.testing.assert_close(fused["gradients"][name], + eager["gradients"][name], + msg=lambda formatted, name=name: f"gradient mismatch for {name}\n{formatted}", + **PARITY_TOLERANCE) + + for name in sorted(eager["delta"]): + torch.testing.assert_close(fused["delta"][name], + eager["delta"][name], + msg=lambda formatted, name=name: f"parameter update mismatch for {name}\n" + f"{formatted}", + **PARITY_TOLERANCE) + + assert any(value.abs().sum() > 0 for value in eager["delta"].values()), "the optimizer step changed nothing" + + +class TestAutoEPFusedParityExpertParallel(DistributedTest): + world_size = 2 + + @pytest.mark.parametrize("checkpoint_activations", [True, False]) + def test_fused_matches_eager_through_a_full_step(self, checkpoint_activations): + seed = 4321 + seed_everything(seed) + reference_state = MockMoETransformer(num_layers=2, + num_experts=NUM_EXPERTS, + hidden_size=HIDDEN_SIZE, + intermediate_size=2 * HIDDEN_SIZE).state_dict() + + eager_engine = _build_engine("eager", 2, reference_state, seed) + eager = _take_one_step(eager_engine, seed, checkpoint_activations=checkpoint_activations) + + fused_engine = _build_engine("fused", 2, reference_state, seed) + assert all(module.local_token_backend == "fused" for module in fused_engine.module.modules() + if isinstance(module, AutoEPMoELayer)), "the fused backend was not actually selected" + fused = _take_one_step(fused_engine, seed, checkpoint_activations=checkpoint_activations) + + _assert_step_matches(fused, eager) + + +class TestAutoEPFusedParityLocalExperts(DistributedTest): + world_size = 1 + + def test_fused_matches_eager_without_expert_parallelism(self): + seed = 8765 + seed_everything(seed) + reference_state = MockMoETransformer(num_layers=2, + num_experts=NUM_EXPERTS, + hidden_size=HIDDEN_SIZE, + intermediate_size=2 * HIDDEN_SIZE).state_dict() + + eager = _take_one_step(_build_engine("eager", 1, reference_state, seed), seed, checkpoint_activations=False) + fused = _take_one_step(_build_engine("fused", 1, reference_state, seed), seed, checkpoint_activations=False) + + _assert_step_matches(fused, eager) diff --git a/tests/unit/v1/moe/test_autoep_fused_token_ops.py b/tests/unit/v1/moe/test_autoep_fused_token_ops.py index 6d376656563b..87df1c8b4d3c 100644 --- a/tests/unit/v1/moe/test_autoep_fused_token_ops.py +++ b/tests/unit/v1/moe/test_autoep_fused_token_ops.py @@ -174,7 +174,7 @@ def test_fused_engine_names_what_it_cannot_run(): with pytest.raises(RuntimeError, match="bfloat16 and float16"): fused_ops.assert_supported(torch.randn(8, 16, device=device, dtype=torch.float32), score_apply="post") - with pytest.raises(RuntimeError, match='score_apply="post"'): + with pytest.raises(RuntimeError, match='resolved score_apply="pre"'): fused_ops.assert_supported(supported, score_apply="pre") with pytest.raises(RuntimeError, match="CUDA kernels"): From 7871a4671e72b987e6413443fcbafbc1d8765d0a Mon Sep 17 00:00:00 2001 From: yh0903 Date: Tue, 25 Aug 2026 10:28:25 -0700 Subject: [PATCH 3/9] Narrow the fused path to the weighted restore, which is where the win is Timing the two fused ops against the eager ones they replaced, on an H100 at the canonical shape, separates them cleanly: weighted restore forward 2.98x fwd+bwd 1.87x saves 0.215 ms per layer expert reorder forward 0.89x fwd+bwd 1.01x saves 0.004 ms per layer The reorder is a wash, and its forward is slower than the eager one. Three attempts to fix that -- a wider hidden tile, more warps, and several rows per program -- all landed within noise of the same number. PyTorch's advanced-index gather and scatter are already close to optimal for this access pattern, and the fused version additionally has to build an inverse index the eager one does not need. So it is removed: it carried a kernel, an autograd pair, an index buffer and a shared-helper refactor, and returned nothing. An earlier measurement had the reorder at 1.23x. That was an artifact of timing one iteration at a time and synchronising after each, which charges every iteration the launch latency a training step hides by queueing work ahead of the GPU. Timing a batch of iterations between one pair of events measures the regime these ops actually run in, and the reorder's apparent win disappeared. What remains is the weighted restore, so it is spelled as what it is: another implementation of the combine, selected by combine_impl="fused_weighted_sum" alongside the existing weighted_sum and legacy_bmm, rather than by a second config key describing a "local token backend" that now moves no tokens. This also dissolves the question of what a fused backend should do when someone asks for legacy_bmm: they are alternatives in one enum and cannot both be chosen. It is still rejected rather than quietly ignored where it has nothing to replace: folded tensor parallelism, a resolved score_apply other than "post", and non-CUDA, non-Triton or non-bf16/fp16 execution, checked once before any collective so ranks fail together instead of stalling. Validated on an H100 node: 121 kernel and config tests pass, and all three eager-versus-fused parity cases pass, comparing loss, output, input gradient, every named gradient, and the optimizer's parameter delta. Signed-off-by: yh0903 --- deepspeed/module_inject/auto_ep_config.py | 43 ++-- deepspeed/module_inject/auto_ep_layer.py | 91 +++++--- .../module_inject/auto_ep_presets/base.py | 3 +- deepspeed/moe/autoep_fused_token_ops.py | 206 ++++-------------- deepspeed/moe/ep_kernels.py | 62 ------ docs/_pages/config-json.md | 4 +- docs/code-docs/source/autoep.rst | 42 ++-- tests/unit/v1/moe/test_autoep_fused_parity.py | 34 +-- .../v1/moe/test_autoep_fused_token_ops.py | 83 +------ tests/unit/v1/moe/test_autoep_unit.py | 31 +-- 10 files changed, 174 insertions(+), 425 deletions(-) diff --git a/deepspeed/module_inject/auto_ep_config.py b/deepspeed/module_inject/auto_ep_config.py index c068f4f73cbd..42c8f370cdda 100644 --- a/deepspeed/module_inject/auto_ep_config.py +++ b/deepspeed/module_inject/auto_ep_config.py @@ -58,7 +58,6 @@ def parse_autoep_config(param_dict: dict) -> AutoEPConfig: config.route_scale = param_dict.get("route_scale", 1.0) config.score_apply = param_dict.get("score_apply", "auto") config.combine_impl = param_dict.get("combine_impl", "auto") - config.local_token_backend = param_dict.get("local_token_backend", "eager") config.num_expert_groups = param_dict.get("num_expert_groups", None) config.num_limited_groups = param_dict.get("num_limited_groups", None) config.score_func = param_dict.get("score_func", "auto") @@ -119,27 +118,15 @@ def validate_autoep_config( if not config.enabled: return - # Validate local_token_backend - valid_local_token_backend = ("eager", "fused") - if config.local_token_backend not in valid_local_token_backend: - raise ValueError(f"local_token_backend must be one of {valid_local_token_backend}, " - f"got '{config.local_token_backend}'") - - # The fused engine only replaces the reorder and weighted restore that the - # standard expert-parallel path runs. Where it has nothing to replace, say so - # instead of running eager under a config that asked for fused: a benchmark - # that believes it measured the fused path would otherwise report noise. - if config.local_token_backend == "fused": - if config.expert_tensor_parallel_size > 1: - raise ValueError('local_token_backend="fused" does not support folded tensor parallelism ' - f"(expert_tensor_parallel_size={config.expert_tensor_parallel_size}), which restores " - "combined tokens from assignment metadata instead of the weighted reduction the fused " - 'engine implements. Set expert_tensor_parallel_size to 1, or local_token_backend to ' - '"eager".') - if config.combine_impl == "legacy_bmm": - raise ValueError('local_token_backend="fused" implements the weighted-sum reduction, so it cannot honor ' - 'combine_impl="legacy_bmm". Leave combine_impl unset, or set local_token_backend to ' - '"eager" to keep the legacy reduction for model-family verification.') + # The fused reduction only replaces the weighted sum the standard + # expert-parallel path runs. Where it has nothing to replace, say so instead + # of running the eager reduction under a config that asked for the fused + # one: a benchmark believing it measured the fused path would report noise. + if config.combine_impl == "fused_weighted_sum" and config.expert_tensor_parallel_size > 1: + raise ValueError('combine_impl="fused_weighted_sum" does not support folded tensor parallelism ' + f"(expert_tensor_parallel_size={config.expert_tensor_parallel_size}), which restores " + "combined tokens from assignment metadata instead of the weighted reduction it " + 'implements. Set expert_tensor_parallel_size to 1, or leave combine_impl unset.') folding_spec = build_folding_spec( world_size=world_size, @@ -175,7 +162,7 @@ def validate_autoep_config( f"got '{config.score_apply}'") # Validate combine_impl - valid_combine_impl = ("auto", "weighted_sum", "legacy_bmm") + valid_combine_impl = ("auto", "weighted_sum", "fused_weighted_sum", "legacy_bmm") if config.combine_impl not in valid_combine_impl: raise ValueError(f"combine_impl must be one of {valid_combine_impl}, " f"got '{config.combine_impl}'") @@ -295,14 +282,14 @@ def validate_autoep_post_detection( return for spec in specs: - # The fused weighted restore folds the routing weight into the top-k - # reduction, which only exists when scores are applied after the experts. - if config.local_token_backend == "fused": + # The fused reduction folds the routing weight into the top-k reduction, + # which only exists when scores are applied after the experts. + if config.combine_impl == "fused_weighted_sum": resolved_score_apply = config.score_apply if config.score_apply != "auto" else spec.score_apply if resolved_score_apply != "post": - raise ValueError(f'local_token_backend="fused" requires score_apply="post", but layer ' + raise ValueError(f'combine_impl="fused_weighted_sum" requires score_apply="post", but layer ' f"'{spec.moe_module_name}' resolved score_apply=\"{resolved_score_apply}\". " - 'Set local_token_backend to "eager".') + "Leave combine_impl unset.") # ep_size must not exceed num_experts if config.autoep_size > spec.num_experts: diff --git a/deepspeed/module_inject/auto_ep_layer.py b/deepspeed/module_inject/auto_ep_layer.py index 4639df2053e2..9b1320965ca3 100644 --- a/deepspeed/module_inject/auto_ep_layer.py +++ b/deepspeed/module_inject/auto_ep_layer.py @@ -62,7 +62,8 @@ def resolve_score_apply_mode( def resolve_combine_impl( - config_override: Literal["auto", "weighted_sum", "legacy_bmm"], ) -> Literal["weighted_sum", "legacy_bmm"]: + config_override: Literal["auto", "weighted_sum", "fused_weighted_sum", "legacy_bmm"], +) -> Literal["weighted_sum", "fused_weighted_sum", "legacy_bmm"]: """Resolve combine implementation from config override or default.""" if config_override != "auto": return config_override @@ -242,14 +243,44 @@ def permute_by_local_expert( aligned_counts: [E_local] aligned token counts per expert (for expert computation) n_tokens: original token count before padding (for unpermute) """ - from deepspeed.moe.ep_kernels import generate_local_expert_permute_indices + from deepspeed.moe.ep_kernels import generate_permute_indices, TOKEN_GROUP_ALIGN_SIZE_M + + if local_counts.ndim == 1: + # [E_local]: already aggregated over sources (ep_degree=1) + ep_degree = 1 + num_local_experts = local_counts.shape[0] + local_counts_flat = local_counts + elif local_counts.ndim == 2: + # [ep_size, E_local]: preserve per-source layout for correct regrouping + ep_degree, num_local_experts = local_counts.shape + local_counts_flat = local_counts.reshape(-1) + else: + raise ValueError( + f"local_counts must have shape [E_local] or [ep_degree, E_local], got {tuple(local_counts.shape)}") n_tokens = tokens.shape[0] - permuted_indices, m_sizes = generate_local_expert_permute_indices( - n_tokens=n_tokens, - local_counts=local_counts, - device=tokens.device, - ) + alignment = TOKEN_GROUP_ALIGN_SIZE_M + + # Compute padded max length + x_padded_per_expert = n_tokens + num_local_experts * alignment + padded_max_len = ((x_padded_per_expert + alignment - 1) // alignment) * alignment + + # Use the pure-PyTorch path for host tensors. The CPU accelerator reports + # CPU tensors as "on accelerator", but Triton still requires a GPU driver. + use_cpu = tokens.device.type == "cpu" + counts_for_permute = local_counts_flat.cpu() if use_cpu else local_counts_flat + with torch.no_grad(): + permuted_indices, m_sizes, _offsets = generate_permute_indices( + counts_for_permute, + num_local_experts, + ep_degree, + padded_max_len, + alignment, + use_cpu=use_cpu, + ) + if not use_cpu: + permuted_indices = permuted_indices.to(tokens.device) + m_sizes = m_sizes.to(tokens.device) # Add padding row for out-of-bounds indices (index n_tokens -> zero row) tokens_padded = torch.vstack((tokens, tokens.new_zeros((tokens.shape[-1], )))) @@ -348,8 +379,7 @@ def __init__( self.top_k = spec.top_k self.score_apply = resolve_score_apply_mode(spec, config.score_apply) self.combine_impl = resolve_combine_impl(config.combine_impl) - self.local_token_backend = config.local_token_backend - self._fused_backend_checked = False + self._fused_combine_checked = False route_norm = spec.route_norm if config.route_norm is None else config.route_norm self.ep_size = ep_size self.ep_rank = ep_rank @@ -523,10 +553,13 @@ def set_deepspeed_parallelism( if folding_group_handles is not None: self.folding_group_handles = folding_group_handles - if self.local_token_backend == "fused" and folding_group_handles.spec.tp_size > 1: - raise ValueError('local_token_backend="fused" does not support folded tensor parallelism ' + if self.combine_impl == "fused_weighted_sum" and folding_group_handles.spec.tp_size > 1: + # Folded TP restores combined tokens from assignment metadata and + # never reaches the weighted reduction this replaces, so the + # request would otherwise be silently ignored. + raise ValueError('combine_impl="fused_weighted_sum" does not support folded tensor parallelism ' f"(expert_tensor_parallel_size={folding_group_handles.spec.tp_size}). Set " - 'expert_tensor_parallel_size to 1, or local_token_backend to "eager".') + 'expert_tensor_parallel_size to 1, or leave combine_impl unset.') self.ep_group_name = folding_group_handles.ep_group_name self.ep_group = folding_group_handles.ep_group self.tp_group = folding_group_handles.tp_group @@ -549,17 +582,6 @@ def set_deepspeed_parallelism( ) self.ep_group = groups._get_expert_parallel_group(self.ep_group_name) - def _run_local_experts(self, rows: torch.Tensor, local_counts: torch.Tensor) -> torch.Tensor: - """Group rows by local expert, run the grouped GEMM, and undo the grouping.""" - if self.local_token_backend == "fused": - reordered, reorder_context = fused_token_ops.fused_permute_by_local_expert(rows, local_counts) - expert_output = self.experts(reordered, reorder_context.aligned_counts) - return fused_token_ops.fused_unpermute_by_local_expert(expert_output, reorder_context) - - reordered, perm_indices, aligned_counts, n_tokens = permute_by_local_expert(rows, local_counts) - expert_output = self.experts(reordered, aligned_counts) - return unpermute_by_local_expert(expert_output, perm_indices, n_tokens) - def forward( self, hidden_states: torch.Tensor, @@ -576,9 +598,12 @@ def forward( bsz, seqlen, hdim = hidden_states.shape x = hidden_states.reshape(-1, hdim) # [T, H] - if self.local_token_backend == "fused" and not self._fused_backend_checked: + # Checked once, ahead of the router and of every collective, so an + # unsupported configuration fails on all ranks together rather than + # stalling the ones that carried on. + if self.combine_impl == "fused_weighted_sum" and not self._fused_combine_checked: fused_token_ops.assert_supported(x, score_apply=self.score_apply) - self._fused_backend_checked = True + self._fused_combine_checked = True # Router ro: RouterOutput = RouterOutput(*self.router(x, self.expert_bias)) @@ -644,7 +669,12 @@ def forward( if self.ep_size == 1: # No AllToAll needed - local computation only - expert_output = self._run_local_experts(routed_input, ro.num_tokens_per_expert) + local_counts = ro.num_tokens_per_expert + + routed_input_permuted, perm_indices, aligned_counts, n_tokens = permute_by_local_expert( + routed_input, local_counts) + expert_output = self.experts(routed_input_permuted, aligned_counts) + expert_output = unpermute_by_local_expert(expert_output, perm_indices, n_tokens) else: # EP dispatch/compute/combine if folded_tp: @@ -666,7 +696,12 @@ def forward( ) routed_input = _AllToAllV.apply(self.ep_group, routed_input, plan.input_splits, plan.output_splits) - expert_output = self._run_local_experts(routed_input, plan.local_counts_by_source) + + routed_input, perm_indices, aligned_counts, n_tokens = permute_by_local_expert( + routed_input, plan.local_counts_by_source) + expert_output = self.experts(routed_input, aligned_counts) + expert_output = unpermute_by_local_expert(expert_output, perm_indices, n_tokens) + expert_output = _AllToAllV.apply(self.ep_group, expert_output, plan.output_splits, plan.input_splits) if folded_tp: @@ -675,7 +710,7 @@ def forward( tp_group=self.tp_group, validate_coverage=self.validate_folding_routing).reshape(bsz, seqlen, hdim) self._last_folding_dispatch_counters = dispatch_counters(restore_ctx) - elif self.local_token_backend == "fused": + elif self.combine_impl == "fused_weighted_sum": output = fused_token_ops.fused_weighted_restore( expert_output, top_scores=ro.top_scores, diff --git a/deepspeed/module_inject/auto_ep_presets/base.py b/deepspeed/module_inject/auto_ep_presets/base.py index 429f2f7ee2d7..33a996ea518d 100644 --- a/deepspeed/module_inject/auto_ep_presets/base.py +++ b/deepspeed/module_inject/auto_ep_presets/base.py @@ -109,8 +109,7 @@ class AutoEPConfig: route_norm: bool | None = None route_scale: float = 1.0 score_apply: Literal["auto", "pre", "post"] = "auto" - combine_impl: Literal["auto", "weighted_sum", "legacy_bmm"] = "auto" - local_token_backend: Literal["eager", "fused"] = "eager" + combine_impl: Literal["auto", "weighted_sum", "fused_weighted_sum", "legacy_bmm"] = "auto" num_expert_groups: int | None = None num_limited_groups: int | None = None score_func: Literal["auto", "softmax", "sigmoid"] = "auto" diff --git a/deepspeed/moe/autoep_fused_token_ops.py b/deepspeed/moe/autoep_fused_token_ops.py index 4af8adc046b2..b061589dd0d9 100644 --- a/deepspeed/moe/autoep_fused_token_ops.py +++ b/deepspeed/moe/autoep_fused_token_ops.py @@ -2,36 +2,28 @@ # SPDX-License-Identifier: Apache-2.0 # DeepSpeed Team -"""Fused GPU-local token movement for AutoEP. - -The eager path moves routed rows around the two expert all-to-alls with -general-purpose tensor ops: it materializes a padded copy of the token matrix, -gathers through an advanced index, scatters expert outputs into a zero-filled -buffer, and builds a ``[T, K, H]`` FP32 intermediate only to apply routing -weights and reduce over top-k. - -This module replaces that sequence with Triton kernels that touch each row once. -It deliberately leaves the collectives, the router and the grouped GEMM alone, so -that a measured difference is attributable to the local token engine. - -The reorder is expressed entirely as row gathers. Writing the four passes out -shows why only one kernel is needed, where ``perm`` maps an expert-major slot to -its source row and ``inv`` is its inverse: - - permute forward out[i] = tokens[perm[i]] gather by perm - permute backward dtok[j] = dout[inv[j]] gather by inv - unpermute forward out[j] = expert_out[inv[j]] gather by inv - unpermute backward dexp[i] = dout[perm[i]] gather by perm - -``perm`` is injective on real rows, so no pass needs atomics, and carrying -``inv`` (4 bytes per row) removes both the padded input copy and the zero-filled -scatter buffer (``2 * H`` bytes per row) that the eager path allocates. +"""Fused weighted token restoration for AutoEP. + +After the combine all-to-all, the eager path returns one row per routed +assignment and turns it back into one row per token in general-purpose steps: it +scatters the rows into a zero-filled ``[tokens * top_k, hidden]`` buffer, views +that as ``[tokens, top_k, hidden]``, widens it to FP32 to apply routing weights, +and reduces over top-k. The FP32 intermediate alone is 64 MiB at the canonical +shape, and every step of that sequence costs a full pass over the routed +activations, in every MoE layer, on every step. + +This module does the same arithmetic in one pass. Each program owns one token +and one slice of the hidden dimension, walks its top-k rows in registers, +accumulates in FP32 and writes the token's output once, so neither the scattered +assignment buffer nor the FP32 intermediate is ever allocated. + +Only the reduction is replaced. The collectives, the router, the grouped GEMM +and the expert-major reorder are all untouched, so a measured difference belongs +to the reduction alone. """ from __future__ import annotations -from typing import NamedTuple - import torch try: @@ -42,45 +34,18 @@ except ImportError: _TRITON_AVAILABLE = False -# The grouped GEMM consumes the reordered rows, so the fused path supports the -# dtypes it is built for rather than silently widening them. +# The grouped GEMM produces the rows this consumes, so the supported dtypes are +# the ones it is built for rather than a silent widening. SUPPORTED_ROW_DTYPES = (torch.bfloat16, torch.float16) _MAX_BLOCK_HIDDEN = 512 _INVERT_INDEX_BLOCK = 256 -# The restore kernels hold a [slots, BLOCK_H] FP32 block live, so the hidden tile -# shrinks as top-k grows to keep that block in registers instead of spilling. +# The kernels hold a [slots, BLOCK_H] FP32 block live, so the hidden tile shrinks +# as top-k grows to keep that block in registers rather than spilling. _MAX_BLOCK_ELEMENTS = 2048 if _TRITON_AVAILABLE: - @triton.jit - def _gather_rows_kernel( - source_ptr, - index_ptr, - out_ptr, - hidden, - source_stride, - out_stride, - BLOCK_H: tl.constexpr, - ): - out_row = tl.program_id(0) - hidden_offsets = tl.program_id(1) * BLOCK_H + tl.arange(0, BLOCK_H) - hidden_mask = hidden_offsets < hidden - - source_row = tl.load(index_ptr + out_row).to(tl.int64) - # A negative index marks an alignment-padding slot. The eager path spelled - # that as an appended zero row; here it is a masked-off load. - row_valid = source_row >= 0 - safe_row = tl.where(row_valid, source_row, 0) - - values = tl.load( - source_ptr + safe_row * source_stride + hidden_offsets, - mask=hidden_mask & row_valid, - other=0.0, - ) - tl.store(out_ptr + out_row * out_stride + hidden_offsets, values, mask=hidden_mask) - @triton.jit def _invert_index_kernel( index_ptr, @@ -169,7 +134,9 @@ def _weighted_restore_backward_kernel( grad_rows_dtype = grad_rows_ptr.dtype.element_ty # One token per program, so the score gradient reduces over the hidden - # dimension in registers instead of through a cross-program atomic. + # dimension in registers. Splitting that dimension across programs and + # reducing the partials afterwards measured slower at this shape: the + # extra pass costs more than the added parallelism returns. score_partials = tl.zeros([K_PADDED, BLOCK_H], dtype=tl.float32) for hidden_start in range(0, hidden, BLOCK_H): @@ -204,42 +171,33 @@ def _weighted_restore_backward_kernel( ) -class FusedReorderContext(NamedTuple): - """Index metadata shared by the reorder forward and backward passes.""" - - permutation: torch.Tensor # [N_padded] int32; -1 marks an alignment-padding slot - inverse: torch.Tensor # [n_tokens] int32; -1 marks a row no slot claimed - aligned_counts: torch.Tensor # [E_local] int32 row counts for the grouped GEMM - n_tokens: int - - def is_available() -> bool: - """Whether this build can run the fused local token engine at all.""" + """Whether this build can run the fused weighted restore at all.""" return _TRITON_AVAILABLE def assert_supported(rows: torch.Tensor, *, score_apply: str) -> None: - """Reject configurations the fused path does not implement. + """Reject configurations the fused restore does not implement. Checked before any collective runs: a rank that raised while its peers proceeded would turn a clear error into a hang. """ if not _TRITON_AVAILABLE: - raise RuntimeError('expert_parallel.local_token_backend="fused" needs Triton, which is not installed in ' - 'this environment. Install Triton, or set local_token_backend to "eager".') + raise RuntimeError('combine_impl="fused_weighted_sum" needs Triton, which is not installed in this ' + "environment. Install Triton, or leave combine_impl unset.") if rows.device.type != "cuda": - raise RuntimeError('expert_parallel.local_token_backend="fused" runs CUDA kernels but this layer is on ' - f'device "{rows.device.type}". Set local_token_backend to "eager" to run here.') + raise RuntimeError('combine_impl="fused_weighted_sum" runs CUDA kernels but this layer is on device ' + f'"{rows.device.type}". Leave combine_impl unset to run here.') if rows.dtype not in SUPPORTED_ROW_DTYPES: - raise RuntimeError('expert_parallel.local_token_backend="fused" supports bfloat16 and float16 rows, got ' - f'{rows.dtype}. Set local_token_backend to "eager", or train in bf16/fp16.') + raise RuntimeError('combine_impl="fused_weighted_sum" supports bfloat16 and float16 rows, got ' + f"{rows.dtype}. Leave combine_impl unset, or train in bf16/fp16.") if score_apply != "post": - raise RuntimeError('expert_parallel.local_token_backend="fused" implements the post-expert weighted ' - f'restore, but this layer resolved score_apply="{score_apply}". Set local_token_backend ' - 'to "eager".') + raise RuntimeError('combine_impl="fused_weighted_sum" folds the routing weight into the top-k reduction, ' + f'which only exists for score_apply="post", but this layer resolved ' + f'score_apply="{score_apply}". Leave combine_impl unset.') -def _block_hidden(hidden: int, slots: int = 1) -> int: +def _block_hidden(hidden: int, slots: int) -> int: """Pick a power-of-two hidden tile that fits alongside ``slots`` rows of FP32. The floor keeps the budget honest for top-k values far wider than any real @@ -254,28 +212,6 @@ def _padded_top_k(top_k: int) -> int: return max(2, triton.next_power_of_2(top_k)) -def _gather_rows(source: torch.Tensor, index: torch.Tensor, num_out_rows: int) -> torch.Tensor: - """Gather the ``source`` rows named by ``index``, reading a negative index as zero.""" - hidden = source.shape[-1] - out = torch.empty((num_out_rows, hidden), dtype=source.dtype, device=source.device) - if num_out_rows == 0 or hidden == 0: - return out - - source = source.contiguous() - block_hidden = _block_hidden(hidden) - grid = (num_out_rows, triton.cdiv(hidden, block_hidden)) - _gather_rows_kernel[grid]( - source, - index, - out, - hidden, - source.stride(0), - out.stride(0), - BLOCK_H=block_hidden, - ) - return out - - def _invert_index(index: torch.Tensor, num_inverse_rows: int) -> torch.Tensor: """Invert a row permutation, leaving -1 wherever no slot claimed a row.""" inverse = torch.full((num_inverse_rows, ), -1, dtype=torch.int32, device=index.device) @@ -294,36 +230,6 @@ def _invert_index(index: torch.Tensor, num_inverse_rows: int) -> torch.Tensor: return inverse -class _FusedReorder(torch.autograd.Function): - """Gather routed rows into expert-major, alignment-padded order.""" - - @staticmethod - def forward(ctx, tokens, permutation, inverse, n_padded): - ctx.save_for_backward(inverse) - ctx.n_tokens = tokens.shape[0] - return _gather_rows(tokens, permutation, n_padded) - - @staticmethod - def backward(ctx, grad_out): - inverse, = ctx.saved_tensors - return _gather_rows(grad_out.contiguous(), inverse, ctx.n_tokens), None, None, None - - -class _FusedInverseReorder(torch.autograd.Function): - """Restore source-major order and drop the alignment padding.""" - - @staticmethod - def forward(ctx, expert_output, permutation, inverse): - ctx.save_for_backward(permutation) - ctx.n_padded = expert_output.shape[0] - return _gather_rows(expert_output, inverse, inverse.numel()) - - @staticmethod - def backward(ctx, grad_out): - permutation, = ctx.saved_tensors - return _gather_rows(grad_out.contiguous(), permutation, ctx.n_padded), None, None - - class _FusedWeightedRestore(torch.autograd.Function): """Weight rows by their routing score and reduce over top-k in one pass.""" @@ -370,7 +276,7 @@ def backward(ctx, grad_output): n_tokens, hidden = top_scores.shape[0], combined_rows.shape[-1] if n_tokens == 0 or hidden == 0: # Nothing is reduced, so the score gradient is zero rather than - # whatever the uninitialized buffer happened to hold. + # whatever an uninitialized buffer happened to hold. return grad_rows, torch.zeros_like(top_scores), None, None k_padded = _padded_top_k(ctx.top_k) @@ -394,42 +300,6 @@ def backward(ctx, grad_output): return grad_rows, grad_scores, None, None -def fused_permute_by_local_expert( - tokens: torch.Tensor, - local_counts: torch.Tensor, -) -> tuple[torch.Tensor, FusedReorderContext]: - """Reorder routed rows into expert-contiguous, alignment-padded order. - - Fused counterpart of ``permute_by_local_expert``. It shares that function's - index generation and produces the same rows, without materializing the padded - copy of ``tokens`` that the eager advanced index needs. - """ - from deepspeed.moe.ep_kernels import generate_local_expert_permute_indices - - permutation, aligned_counts = generate_local_expert_permute_indices( - n_tokens=tokens.shape[0], - local_counts=local_counts, - device=tokens.device, - ) - inverse = _invert_index(permutation, tokens.shape[0]) - context = FusedReorderContext( - permutation=permutation, - inverse=inverse, - aligned_counts=aligned_counts, - n_tokens=tokens.shape[0], - ) - permuted = _FusedReorder.apply(tokens, permutation, inverse, permutation.numel()) - return permuted, context - - -def fused_unpermute_by_local_expert( - expert_output: torch.Tensor, - context: FusedReorderContext, -) -> torch.Tensor: - """Reverse :func:`fused_permute_by_local_expert` and strip the padding.""" - return _FusedInverseReorder.apply(expert_output.contiguous(), context.permutation, context.inverse) - - def fused_weighted_restore( combined_rows: torch.Tensor, top_scores: torch.Tensor, diff --git a/deepspeed/moe/ep_kernels.py b/deepspeed/moe/ep_kernels.py index 6818d858c81a..6b3da7335e88 100644 --- a/deepspeed/moe/ep_kernels.py +++ b/deepspeed/moe/ep_kernels.py @@ -250,68 +250,6 @@ def generate_permute_indices( return permuted_indices, m_sizes, m_offsets.to(torch.int32) -# =================================================================== -# generate_local_expert_permute_indices -# =================================================================== - - -def generate_local_expert_permute_indices( - n_tokens: int, - local_counts: torch.Tensor, - device: torch.device, -) -> tuple: - """Build the expert-major row order for ``n_tokens`` received rows. - - Args: - n_tokens: Number of received rows, before alignment padding. - local_counts: ``(E_local,)`` when the counts are already aggregated over - sources, or ``(ep_degree, E_local)`` to keep the per-source layout - that correct regrouping depends on. - device: Device the returned tensors should live on. - - Returns: - Tuple of: - - permuted_indices: expert-major slot -> source row, ``-1`` where a - slot is alignment padding rather than a routed row. - - aligned_counts: per-expert row counts for the grouped GEMM. - """ - if local_counts.ndim == 1: - # [E_local]: already aggregated over sources (ep_degree=1) - ep_degree = 1 - num_local_experts = local_counts.shape[0] - local_counts_flat = local_counts - elif local_counts.ndim == 2: - # [ep_size, E_local]: preserve per-source layout for correct regrouping - ep_degree, num_local_experts = local_counts.shape - local_counts_flat = local_counts.reshape(-1) - else: - raise ValueError( - f"local_counts must have shape [E_local] or [ep_degree, E_local], got {tuple(local_counts.shape)}") - - alignment = TOKEN_GROUP_ALIGN_SIZE_M - x_padded_per_expert = n_tokens + num_local_experts * alignment - padded_max_len = _round_up(x_padded_per_expert, alignment) - - # Use the pure-PyTorch path for host tensors. The CPU accelerator reports - # CPU tensors as "on accelerator", but Triton still requires a GPU driver. - use_cpu = device.type == "cpu" - counts_for_permute = local_counts_flat.cpu() if use_cpu else local_counts_flat - with torch.no_grad(): - permuted_indices, m_sizes, _offsets = generate_permute_indices( - counts_for_permute, - num_local_experts, - ep_degree, - padded_max_len, - alignment, - use_cpu=use_cpu, - ) - if not use_cpu: - permuted_indices = permuted_indices.to(device) - m_sizes = m_sizes.to(device) - - return permuted_indices, m_sizes - - # =================================================================== # _permute / _unpermute / indices_padding_wrapper # =================================================================== diff --git a/docs/_pages/config-json.md b/docs/_pages/config-json.md index 6914986512f5..5d02220febc4 100644 --- a/docs/_pages/config-json.md +++ b/docs/_pages/config-json.md @@ -994,11 +994,11 @@ smoke coverage used for this AutoEP surface produced the following version gates | -------------------------------------------------------------------------------------------------------------- | -------- | | When to apply router scores: `"pre"` (before experts), `"post"` (during combine), or `"auto"` (from preset). | `"auto"` | -***local_token_backend***: [string] +***combine_impl***: [string] | Description | Default | | -------------------------------------------------------------------------------------------------------------- | -------- | -| GPU-local token movement backend. `"eager"` keeps the general-purpose reorder and combine. `"fused"` is experimental and replaces them with Triton kernels that group rows by expert, undo that grouping, and apply router scores while reducing over top-k, without materializing the padded token copy or the `[tokens, top_k, hidden]` FP32 intermediate. Collectives, routing and the grouped GEMM are unchanged. `"fused"` requires CUDA, Triton, bfloat16/float16 activations, `expert_tensor_parallel_size=1`, and a resolved `score_apply="post"`; it is rejected rather than silently ignored when any of those does not hold. | `"eager"` | +| How expert outputs are weighted by their router scores and reduced over top-k. `"auto"` resolves to `"weighted_sum"`. `"fused_weighted_sum"` is experimental and computes the same reduction in one Triton pass, without materializing the scattered assignment buffer or the `[tokens, top_k, hidden]` FP32 intermediate; it requires CUDA, Triton, bfloat16/float16 activations, `expert_tensor_parallel_size=1`, and a resolved `score_apply="post"`, and is rejected rather than silently ignored when any of those does not hold. `"legacy_bmm"` is a debug reduction retained for model-family verification. | `"auto"` | ***route_norm***: [boolean] diff --git a/docs/code-docs/source/autoep.rst b/docs/code-docs/source/autoep.rst index e960a7fc2c32..49ca41fed6af 100644 --- a/docs/code-docs/source/autoep.rst +++ b/docs/code-docs/source/autoep.rst @@ -84,11 +84,10 @@ Weights-only/module-only Universal Checkpoint loads use the converted 4. Expert parameters are marked for expert-data-parallel gradient reduction; router and shared-expert parameters use standard data-parallel reduction. -**Fused local token engine (experimental):** +**Fused weighted restore (experimental):** -``expert_parallel.local_token_backend`` selects how routed rows are moved around -the two expert all-to-alls. ``"eager"`` (default) keeps the existing -general-purpose tensor ops. ``"fused"`` replaces them with Triton kernels: +After the combine all-to-all, AutoEP holds one row per routed assignment and has +to turn it back into one row per token. ``combine_impl`` selects how: .. code-block:: json @@ -97,34 +96,35 @@ general-purpose tensor ops. ``"fused"`` replaces them with Triton kernels: "enabled": true, "autoep_size": 16, "preset_model": "qwen3_moe", - "local_token_backend": "fused" + "combine_impl": "fused_weighted_sum" } } -The fused engine groups rows by local expert, undoes that grouping, and applies -router scores while reducing over top-k. It reads and writes each row once, so -the padded copy of the token matrix, the zero-filled scatter buffer, and the -``[tokens, top_k, hidden]`` FP32 intermediate are never allocated. Routing -weights are still accumulated in FP32 and cast once, matching the eager result to -within the reduction order. +``"auto"`` (default) resolves to ``"weighted_sum"``, which scatters the rows into +a zero-filled ``[tokens * top_k, hidden]`` buffer, widens it to FP32 to apply the +routing weights, and reduces over top-k. ``"fused_weighted_sum"`` computes the +same result in a single pass: each program owns one token and one slice of the +hidden dimension, walks its top-k rows in registers and accumulates in FP32, so +neither the scattered buffer nor the FP32 intermediate is allocated. At the +canonical shape the FP32 intermediate alone is 64 MiB per layer. -The collectives, the router and the grouped GEMM are untouched, so a measured -difference between the two backends belongs to the local token engine alone. +Routing weights are still accumulated in FP32 and cast once, so the result +matches the eager reduction to within the order of the top-k summation. Only the +reduction changes: the collectives, the router, the grouped GEMM and the +expert-major reorder are untouched. -``"fused"`` is rejected, rather than quietly ignored, when it would have nothing -to replace or would change semantics: +``"fused_weighted_sum"`` is rejected, rather than quietly ignored, when it would +have nothing to replace or would change semantics: - ``expert_tensor_parallel_size`` greater than 1, which restores combined tokens - from assignment metadata instead of the weighted reduction; -- an explicit ``combine_impl="legacy_bmm"``, which selects the legacy reduction - kept for model-family verification; + from assignment metadata instead; - a resolved ``score_apply`` other than ``"post"``; - activations that are not bfloat16 or float16, a non-CUDA device, or a build without Triton. -Failing fast matters for measurement: a run that asked for ``"fused"`` and -silently got ``"eager"`` would report the difference between a backend and -itself. +Failing fast matters for measurement: a run that asked for the fused reduction +and silently got the eager one would report the difference between an +implementation and itself. **Constraints:** diff --git a/tests/unit/v1/moe/test_autoep_fused_parity.py b/tests/unit/v1/moe/test_autoep_fused_parity.py index 47ddcab5a60b..f480cdb77524 100644 --- a/tests/unit/v1/moe/test_autoep_fused_parity.py +++ b/tests/unit/v1/moe/test_autoep_fused_parity.py @@ -2,11 +2,11 @@ # SPDX-License-Identifier: Apache-2.0 # DeepSpeed Team -"""End-to-end parity between the eager and fused local token backends. +"""End-to-end parity between the eager and fused combine implementations. -The fused engine is only a layout change, so a step taken through it has to -produce the same loss, the same gradients on every trainable tensor, and the same -parameter update as the eager path it replaces. +The fused reduction only changes how the weighted sum is computed, so a step +taken through it has to produce the same loss, the same gradients on every +trainable tensor, and the same parameter update as the eager path. """ import functools @@ -44,10 +44,10 @@ def _fused_engine_available(): pytestmark = pytest.mark.skipif(not _fused_engine_available(), - reason="the fused local token engine needs CUDA and Triton") + reason="the fused weighted restore needs CUDA and Triton") -def _config(local_token_backend, ep_size): +def _config(combine_impl, ep_size): return { **mixed_precision_config(), "train_micro_batch_size_per_gpu": 1, @@ -65,19 +65,19 @@ def _config(local_token_backend, ep_size): "autoep_size": ep_size, "preset_model": "mixtral", "load_balance_coeff": None, - "local_token_backend": local_token_backend, + "combine_impl": combine_impl, }, } -def _build_engine(local_token_backend, ep_size, reference_state, seed): +def _build_engine(combine_impl, ep_size, reference_state, seed): seed_everything(seed) model = MockMoETransformer(num_layers=2, num_experts=NUM_EXPERTS, hidden_size=HIDDEN_SIZE, intermediate_size=2 * HIDDEN_SIZE) model.load_state_dict(reference_state) - engine, _, _, _ = deepspeed.initialize(model=model, config=_config(local_token_backend, ep_size)) + engine, _, _, _ = deepspeed.initialize(model=model, config=_config(combine_impl, ep_size)) return engine @@ -173,12 +173,12 @@ def test_fused_matches_eager_through_a_full_step(self, checkpoint_activations): hidden_size=HIDDEN_SIZE, intermediate_size=2 * HIDDEN_SIZE).state_dict() - eager_engine = _build_engine("eager", 2, reference_state, seed) + eager_engine = _build_engine("weighted_sum", 2, reference_state, seed) eager = _take_one_step(eager_engine, seed, checkpoint_activations=checkpoint_activations) - fused_engine = _build_engine("fused", 2, reference_state, seed) - assert all(module.local_token_backend == "fused" for module in fused_engine.module.modules() - if isinstance(module, AutoEPMoELayer)), "the fused backend was not actually selected" + fused_engine = _build_engine("fused_weighted_sum", 2, reference_state, seed) + assert all(module.combine_impl == "fused_weighted_sum" for module in fused_engine.module.modules() + if isinstance(module, AutoEPMoELayer)), "the fused reduction was not actually selected" fused = _take_one_step(fused_engine, seed, checkpoint_activations=checkpoint_activations) _assert_step_matches(fused, eager) @@ -195,7 +195,11 @@ def test_fused_matches_eager_without_expert_parallelism(self): hidden_size=HIDDEN_SIZE, intermediate_size=2 * HIDDEN_SIZE).state_dict() - eager = _take_one_step(_build_engine("eager", 1, reference_state, seed), seed, checkpoint_activations=False) - fused = _take_one_step(_build_engine("fused", 1, reference_state, seed), seed, checkpoint_activations=False) + eager = _take_one_step(_build_engine("weighted_sum", 1, reference_state, seed), + seed, + checkpoint_activations=False) + fused = _take_one_step(_build_engine("fused_weighted_sum", 1, reference_state, seed), + seed, + checkpoint_activations=False) _assert_step_matches(fused, eager) diff --git a/tests/unit/v1/moe/test_autoep_fused_token_ops.py b/tests/unit/v1/moe/test_autoep_fused_token_ops.py index 87df1c8b4d3c..a1d791b0a407 100644 --- a/tests/unit/v1/moe/test_autoep_fused_token_ops.py +++ b/tests/unit/v1/moe/test_autoep_fused_token_ops.py @@ -2,11 +2,11 @@ # SPDX-License-Identifier: Apache-2.0 # DeepSpeed Team -"""The fused local token engine against the eager reorder and weighted restore. +"""The fused weighted restore against the eager reduction it replaces. -The eager implementations are the reference: the fused engine is only worth -having if it is indistinguishable from them, so every assertion here compares the -two directly rather than against hand-written expectations. +``combine_from_routed`` is the reference: the fused reduction is only worth +having if it is indistinguishable from it, so every assertion compares the two +directly rather than against hand-written expectations. """ import pytest @@ -14,11 +14,7 @@ from deepspeed.accelerator import get_accelerator from deepspeed.moe import autoep_fused_token_ops as fused_ops -from deepspeed.module_inject.auto_ep_layer import ( - combine_from_routed, - permute_by_local_expert, - unpermute_by_local_expert, -) +from deepspeed.module_inject.auto_ep_layer import combine_from_routed def _fused_engine_available(): @@ -27,80 +23,13 @@ def _fused_engine_available(): pytestmark = pytest.mark.skipif(not _fused_engine_available(), - reason="the fused local token engine needs CUDA and Triton") - -# Row counts per local expert, or per [source rank, local expert] where nested. -# They are the only thing that decides the reorder, so they carry every shape the -# engine has to survive: idle experts, one expert taking everything, and sources -# that contribute nothing to a given expert. -REORDER_CASES = { - "balanced": [8, 8, 8, 8], - "empty_experts": [0, 12, 0, 4], - "extreme_skew": [40, 0, 0, 0], - "per_source": [[3, 5], [7, 1]], - "ragged_per_source": [[0, 9], [5, 0], [2, 2]], -} + reason="the fused weighted restore needs CUDA and Triton") def _device(): return get_accelerator().current_device_name() -def _counts(case): - return torch.tensor(REORDER_CASES[case], dtype=torch.int32, device=_device()) - - -@pytest.mark.parametrize("case", sorted(REORDER_CASES)) -@pytest.mark.parametrize("hidden", [64, 130]) -def test_fused_reorder_places_the_same_rows_as_eager(case, hidden): - counts = _counts(case) - tokens = torch.randn(int(counts.sum()), hidden, device=_device(), dtype=torch.bfloat16) - - eager_rows, _permutation, eager_counts, _n_tokens = permute_by_local_expert(tokens, counts) - fused_rows, context = fused_ops.fused_permute_by_local_expert(tokens, counts) - - # Pure data movement on both sides, so anything short of equality is a bug. - assert torch.equal(fused_rows, eager_rows) - assert torch.equal(context.aligned_counts, eager_counts) - - -def test_fused_reorder_handles_a_batch_no_expert_claimed(): - counts = torch.zeros(4, dtype=torch.int32, device=_device()) - tokens = torch.randn(0, 32, device=_device(), dtype=torch.bfloat16) - - eager_rows, _permutation, eager_counts, _n_tokens = permute_by_local_expert(tokens, counts) - fused_rows, context = fused_ops.fused_permute_by_local_expert(tokens, counts) - - assert torch.equal(fused_rows, eager_rows) - assert torch.equal(context.aligned_counts, eager_counts) - assert not fused_rows.any() - - -@pytest.mark.parametrize("case", sorted(REORDER_CASES)) -def test_fused_reorder_round_trip_matches_eager_including_gradients(case): - hidden = 96 - counts = _counts(case) - n_tokens = int(counts.sum()) - tokens = torch.randn(n_tokens, hidden, device=_device(), dtype=torch.bfloat16) - upstream = torch.randn(n_tokens, hidden, device=_device(), dtype=torch.bfloat16) - - eager_tokens = tokens.clone().requires_grad_(True) - eager_rows, permutation, _counts_out, n = permute_by_local_expert(eager_tokens, counts) - # Scaling by a power of two keeps the comparison exact while still putting a - # real op between the reorder and its inverse. - eager_output = unpermute_by_local_expert(eager_rows * 2.0, permutation, n) - - fused_tokens = tokens.clone().requires_grad_(True) - fused_rows, context = fused_ops.fused_permute_by_local_expert(fused_tokens, counts) - fused_output = fused_ops.fused_unpermute_by_local_expert(fused_rows * 2.0, context) - - assert torch.equal(fused_output, eager_output) - - eager_output.backward(upstream) - fused_output.backward(upstream) - assert torch.equal(fused_tokens.grad, eager_tokens.grad) - - @pytest.mark.parametrize("top_k", [2, 4, 6, 8]) @pytest.mark.parametrize("hidden", [128, 130]) @pytest.mark.parametrize("score_dtype", [torch.float32, torch.bfloat16]) diff --git a/tests/unit/v1/moe/test_autoep_unit.py b/tests/unit/v1/moe/test_autoep_unit.py index 642bd11b7476..1792556ecfc4 100644 --- a/tests/unit/v1/moe/test_autoep_unit.py +++ b/tests/unit/v1/moe/test_autoep_unit.py @@ -234,46 +234,33 @@ def test_validate_folding_routing_requires_boolean(self): tp_size=1, sp_size=1) - def test_local_token_backend_defaults_to_eager(self): - assert parse_autoep_config({}).local_token_backend == "eager" - assert parse_autoep_config({"enabled": True}).local_token_backend == "eager" - - def test_local_token_backend_rejects_unknown_value(self): - config = parse_autoep_config({"enabled": True, "local_token_backend": "triton"}) - with pytest.raises(ValueError, match="local_token_backend must be one of"): + def test_combine_impl_rejects_unknown_value(self): + config = parse_autoep_config({"enabled": True, "combine_impl": "triton"}) + with pytest.raises(ValueError, match="combine_impl must be one of"): validate_autoep_config(config, world_size=1, pp_size=1, tp_size=1, sp_size=1) - def test_fused_local_token_backend_rejects_folded_tensor_parallelism(self): + def test_fused_combine_rejects_folded_tensor_parallelism(self): config = parse_autoep_config({ "enabled": True, "autoep_size": 2, "expert_tensor_parallel_size": 2, - "local_token_backend": "fused", + "combine_impl": "fused_weighted_sum", }) with pytest.raises(ValueError, match="folded tensor parallelism"): validate_autoep_config(config, world_size=4, pp_size=1, tp_size=2, sp_size=1) - def test_fused_local_token_backend_rejects_explicit_legacy_bmm(self): - config = parse_autoep_config({ - "enabled": True, - "local_token_backend": "fused", - "combine_impl": "legacy_bmm", - }) - with pytest.raises(ValueError, match="legacy_bmm"): - validate_autoep_config(config, world_size=1, pp_size=1, tp_size=1, sp_size=1) - @pytest.mark.parametrize("score_apply, spec_score_apply", [("auto", "pre"), ("pre", "post")]) - def test_fused_local_token_backend_requires_post_score_apply(self, score_apply, spec_score_apply): + def test_fused_combine_requires_post_score_apply(self, score_apply, spec_score_apply): config = parse_autoep_config({ "enabled": True, - "local_token_backend": "fused", + "combine_impl": "fused_weighted_sum", "score_apply": score_apply, }) with pytest.raises(ValueError, match='requires score_apply="post"'): validate_autoep_post_detection(config, [_make_spec(score_apply=spec_score_apply)]) - def test_fused_local_token_backend_accepts_the_standard_path(self): - config = parse_autoep_config({"enabled": True, "autoep_size": 2, "local_token_backend": "fused"}) + def test_fused_combine_accepts_the_standard_path(self): + config = parse_autoep_config({"enabled": True, "autoep_size": 2, "combine_impl": "fused_weighted_sum"}) validate_autoep_config(config, world_size=2, pp_size=1, tp_size=1, sp_size=1) validate_autoep_post_detection(config, [_make_spec(num_experts=4, score_apply="post")]) From 5c14b012aaf1e93d2838f52b5998ee4f072d1c2f Mon Sep 17 00:00:00 2001 From: yh0903 Date: Wed, 26 Aug 2026 14:57:49 -0700 Subject: [PATCH 4/9] Reject fused restore for folded tensor parallelism early Validate AutoTP folding and expert tensor parallelism independently so neither configuration can mask the other before process-group setup. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: yh0903 --- deepspeed/module_inject/auto_ep_config.py | 15 ++++++++++----- deepspeed/module_inject/auto_ep_layer.py | 4 ++-- docs/_pages/config-json.md | 2 +- docs/code-docs/source/autoep.rst | 5 +++-- tests/unit/v1/moe/test_autoep_unit.py | 13 +++++++++++-- 5 files changed, 27 insertions(+), 12 deletions(-) diff --git a/deepspeed/module_inject/auto_ep_config.py b/deepspeed/module_inject/auto_ep_config.py index 42c8f370cdda..dfc3868a771a 100644 --- a/deepspeed/module_inject/auto_ep_config.py +++ b/deepspeed/module_inject/auto_ep_config.py @@ -122,11 +122,16 @@ def validate_autoep_config( # expert-parallel path runs. Where it has nothing to replace, say so instead # of running the eager reduction under a config that asked for the fused # one: a benchmark believing it measured the fused path would report noise. - if config.combine_impl == "fused_weighted_sum" and config.expert_tensor_parallel_size > 1: - raise ValueError('combine_impl="fused_weighted_sum" does not support folded tensor parallelism ' - f"(expert_tensor_parallel_size={config.expert_tensor_parallel_size}), which restores " - "combined tokens from assignment metadata instead of the weighted reduction it " - 'implements. Set expert_tensor_parallel_size to 1, or leave combine_impl unset.') + if config.combine_impl == "fused_weighted_sum": + if tp_size > 1: + raise ValueError('combine_impl="fused_weighted_sum" does not support folded tensor parallelism ' + f"(tensor_parallel.autotp_size={tp_size}), which restores combined tokens from " + "assignment metadata instead of the weighted reduction it implements. Set " + 'tensor_parallel.autotp_size to 1, or leave combine_impl unset.') + if config.expert_tensor_parallel_size > 1: + raise ValueError('combine_impl="fused_weighted_sum" requires expert_tensor_parallel_size=1, but got ' + f"{config.expert_tensor_parallel_size}. Set expert_tensor_parallel_size to 1, or leave " + "combine_impl unset.") folding_spec = build_folding_spec( world_size=world_size, diff --git a/deepspeed/module_inject/auto_ep_layer.py b/deepspeed/module_inject/auto_ep_layer.py index 9b1320965ca3..65f907816762 100644 --- a/deepspeed/module_inject/auto_ep_layer.py +++ b/deepspeed/module_inject/auto_ep_layer.py @@ -558,8 +558,8 @@ def set_deepspeed_parallelism( # never reaches the weighted reduction this replaces, so the # request would otherwise be silently ignored. raise ValueError('combine_impl="fused_weighted_sum" does not support folded tensor parallelism ' - f"(expert_tensor_parallel_size={folding_group_handles.spec.tp_size}). Set " - 'expert_tensor_parallel_size to 1, or leave combine_impl unset.') + f"(tensor_parallel.autotp_size={folding_group_handles.spec.tp_size}). Set " + 'tensor_parallel.autotp_size to 1, or leave combine_impl unset.') self.ep_group_name = folding_group_handles.ep_group_name self.ep_group = folding_group_handles.ep_group self.tp_group = folding_group_handles.tp_group diff --git a/docs/_pages/config-json.md b/docs/_pages/config-json.md index 5d02220febc4..7e9f4230f500 100644 --- a/docs/_pages/config-json.md +++ b/docs/_pages/config-json.md @@ -998,7 +998,7 @@ smoke coverage used for this AutoEP surface produced the following version gates | Description | Default | | -------------------------------------------------------------------------------------------------------------- | -------- | -| How expert outputs are weighted by their router scores and reduced over top-k. `"auto"` resolves to `"weighted_sum"`. `"fused_weighted_sum"` is experimental and computes the same reduction in one Triton pass, without materializing the scattered assignment buffer or the `[tokens, top_k, hidden]` FP32 intermediate; it requires CUDA, Triton, bfloat16/float16 activations, `expert_tensor_parallel_size=1`, and a resolved `score_apply="post"`, and is rejected rather than silently ignored when any of those does not hold. `"legacy_bmm"` is a debug reduction retained for model-family verification. | `"auto"` | +| How expert outputs are weighted by their router scores and reduced over top-k. `"auto"` resolves to `"weighted_sum"`. `"fused_weighted_sum"` is experimental and computes the same reduction in one Triton pass, without materializing the scattered assignment buffer or the `[tokens, top_k, hidden]` FP32 intermediate; it requires CUDA, Triton, bfloat16/float16 activations, `tensor_parallel.autotp_size=1`, `expert_tensor_parallel_size=1`, and a resolved `score_apply="post"`, and is rejected rather than silently ignored when any of those does not hold. `"legacy_bmm"` is a debug reduction retained for model-family verification. | `"auto"` | ***route_norm***: [boolean] diff --git a/docs/code-docs/source/autoep.rst b/docs/code-docs/source/autoep.rst index 49ca41fed6af..1864c967c8a1 100644 --- a/docs/code-docs/source/autoep.rst +++ b/docs/code-docs/source/autoep.rst @@ -116,8 +116,9 @@ expert-major reorder are untouched. ``"fused_weighted_sum"`` is rejected, rather than quietly ignored, when it would have nothing to replace or would change semantics: -- ``expert_tensor_parallel_size`` greater than 1, which restores combined tokens - from assignment metadata instead; +- ``tensor_parallel.autotp_size`` greater than 1, which uses folded tensor + parallelism and restores combined tokens from assignment metadata instead; +- ``expert_tensor_parallel_size`` greater than 1; - a resolved ``score_apply`` other than ``"post"``; - activations that are not bfloat16 or float16, a non-CUDA device, or a build without Triton. diff --git a/tests/unit/v1/moe/test_autoep_unit.py b/tests/unit/v1/moe/test_autoep_unit.py index 1792556ecfc4..30c71ab3676a 100644 --- a/tests/unit/v1/moe/test_autoep_unit.py +++ b/tests/unit/v1/moe/test_autoep_unit.py @@ -243,12 +243,21 @@ def test_fused_combine_rejects_folded_tensor_parallelism(self): config = parse_autoep_config({ "enabled": True, "autoep_size": 2, - "expert_tensor_parallel_size": 2, "combine_impl": "fused_weighted_sum", }) - with pytest.raises(ValueError, match="folded tensor parallelism"): + with pytest.raises(ValueError, match=r"tensor_parallel\.autotp_size=2"): validate_autoep_config(config, world_size=4, pp_size=1, tp_size=2, sp_size=1) + def test_fused_combine_rejects_expert_tensor_parallelism(self): + config = parse_autoep_config({ + "enabled": True, + "autoep_size": 2, + "expert_tensor_parallel_size": 2, + "combine_impl": "fused_weighted_sum", + }) + with pytest.raises(ValueError, match="requires expert_tensor_parallel_size=1"): + validate_autoep_config(config, world_size=4, pp_size=1, tp_size=1, sp_size=1) + @pytest.mark.parametrize("score_apply, spec_score_apply", [("auto", "pre"), ("pre", "post")]) def test_fused_combine_requires_post_score_apply(self, score_apply, spec_score_apply): config = parse_autoep_config({ From 9a4b79a7e970e174782ee2cfcefa2f681a3a8aa0 Mon Sep 17 00:00:00 2001 From: yh0903 Date: Wed, 26 Aug 2026 15:17:00 -0700 Subject: [PATCH 5/9] Validate fused restore inputs and preserve higher-order gradients Reject malformed tensor contracts before launching Triton, keep malformed permutations deterministic, and use a differentiable PyTorch backward when create_graph requests higher-order autograd. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: yh0903 --- deepspeed/moe/autoep_fused_token_ops.py | 54 +++++++++++++++++-- .../v1/moe/test_autoep_fused_token_ops.py | 45 ++++++++++++++++ 2 files changed, 95 insertions(+), 4 deletions(-) diff --git a/deepspeed/moe/autoep_fused_token_ops.py b/deepspeed/moe/autoep_fused_token_ops.py index b061589dd0d9..99ab1290fdbe 100644 --- a/deepspeed/moe/autoep_fused_token_ops.py +++ b/deepspeed/moe/autoep_fused_token_ops.py @@ -230,6 +230,22 @@ def _invert_index(index: torch.Tensor, num_inverse_rows: int) -> torch.Tensor: return inverse +def _differentiable_backward(grad_output, combined_rows, top_scores, inverse, top_k): + """Build the rare higher-order backward with regular PyTorch operations.""" + n_tokens, hidden = top_scores.shape[0], combined_rows.shape[-1] + valid = inverse >= 0 + safe_inverse = inverse.clamp_min(0).to(torch.int64) + gathered_rows = combined_rows.index_select(0, safe_inverse).reshape(n_tokens, top_k, hidden) + + grad_by_assignment = (grad_output[:, None, :] * top_scores[:, :, None]).to(combined_rows.dtype).reshape(-1, hidden) + grad_rows = torch.zeros_like(combined_rows) + grad_rows = grad_rows.index_copy(0, safe_inverse[valid], grad_by_assignment[valid]) + + grad_scores = (gathered_rows.float() * grad_output.float()[:, None, :]).sum(dim=-1) + grad_scores = torch.where(valid.reshape(n_tokens, top_k), grad_scores, 0.0).to(top_scores.dtype) + return grad_rows, grad_scores + + class _FusedWeightedRestore(torch.autograd.Function): """Weight rows by their routing score and reduce over top-k in one pass.""" @@ -267,10 +283,15 @@ def backward(ctx, grad_output): combined_rows, top_scores, inverse = ctx.saved_tensors grad_output = grad_output.contiguous() - # Every row is claimed by exactly one slot, which fused_weighted_restore - # checks by shape before building the inverse, so both gradients are - # written in full and neither buffer needs pre-zeroing. - grad_rows = torch.empty_like(combined_rows) + if torch.is_grad_enabled(): + grad_rows, grad_scores = _differentiable_backward(grad_output, combined_rows, top_scores, inverse, + ctx.top_k) + return grad_rows, grad_scores, None, None + + # AutoEP supplies an exact permutation, so every row is written once. + # Zero initialization also keeps malformed direct calls deterministic + # when an invalid or duplicate assignment leaves an inverse slot empty. + grad_rows = torch.zeros_like(combined_rows) grad_scores = torch.empty_like(top_scores) n_tokens, hidden = top_scores.shape[0], combined_rows.shape[-1] @@ -314,11 +335,36 @@ def fused_weighted_restore( assignment buffer nor the ``[T, K, H]`` FP32 intermediate is allocated. """ bsz, seqlen, hidden = shape + if top_k <= 0: + raise RuntimeError(f"fused weighted restore expects top_k > 0, got {top_k}.") + if bsz < 0 or seqlen < 0 or hidden < 0: + raise RuntimeError(f"fused weighted restore expects non-negative output dimensions, got {shape}.") + if combined_rows.ndim != 2: + raise RuntimeError(f"fused weighted restore expects combined_rows to be 2D, got shape " + f"{tuple(combined_rows.shape)}.") + if combined_rows.shape[1] != hidden: + raise RuntimeError(f"fused weighted restore output hidden size is {hidden}, but combined rows have hidden " + f"size {combined_rows.shape[1]}.") + n_tokens = bsz * seqlen expected_rows = n_tokens * top_k if combined_rows.shape[0] != expected_rows: raise RuntimeError(f"fused weighted restore expects one row per assignment: {expected_rows} rows for " f"{n_tokens} tokens at top_k={top_k}, got {combined_rows.shape[0]}.") + if tuple(top_scores.shape) != (n_tokens, top_k): + raise RuntimeError(f"fused weighted restore expects top_scores shape {(n_tokens, top_k)}, got " + f"{tuple(top_scores.shape)}.") + if token_indices_sorted.ndim != 1 or token_indices_sorted.numel() != expected_rows: + raise RuntimeError(f"fused weighted restore expects token_indices_sorted to contain {expected_rows} " + f"assignments, got shape {tuple(token_indices_sorted.shape)}.") + if token_indices_sorted.dtype not in (torch.int32, torch.int64): + raise RuntimeError("fused weighted restore expects token_indices_sorted to use int32 or int64 indices, got " + f"{token_indices_sorted.dtype}.") + if not torch.is_floating_point(top_scores): + raise RuntimeError(f"fused weighted restore expects floating-point top_scores, got {top_scores.dtype}.") + if combined_rows.device != top_scores.device or combined_rows.device != token_indices_sorted.device: + raise RuntimeError("fused weighted restore expects rows, scores, and indices on the same device, got " + f"{combined_rows.device}, {top_scores.device}, and {token_indices_sorted.device}.") inverse = _invert_index(token_indices_sorted, expected_rows) output = _FusedWeightedRestore.apply(combined_rows, top_scores.contiguous(), inverse, top_k) diff --git a/tests/unit/v1/moe/test_autoep_fused_token_ops.py b/tests/unit/v1/moe/test_autoep_fused_token_ops.py index a1d791b0a407..383f97615bb3 100644 --- a/tests/unit/v1/moe/test_autoep_fused_token_ops.py +++ b/tests/unit/v1/moe/test_autoep_fused_token_ops.py @@ -95,6 +95,51 @@ def test_fused_weighted_restore_requires_one_row_per_assignment(): ) +@pytest.mark.parametrize( + "rows_shape,scores_shape,index_count,index_dtype,error", + [ + ((8, 15), (4, 2), 8, torch.int64, "output hidden size"), + ((8, 16), (4, 1), 8, torch.int64, "top_scores shape"), + ((8, 16), (4, 2), 7, torch.int64, "token_indices_sorted"), + ((8, 16), (4, 2), 8, torch.float32, "int32 or int64"), + ], +) +def test_fused_weighted_restore_validates_input_contract(rows_shape, scores_shape, index_count, index_dtype, error): + device = _device() + with pytest.raises(RuntimeError, match=error): + fused_ops.fused_weighted_restore( + torch.randn(rows_shape, device=device, dtype=torch.bfloat16), + top_scores=torch.rand(scores_shape, device=device), + token_indices_sorted=torch.arange(index_count, device=device).to(index_dtype), + top_k=2, + shape=(1, 4, 16), + ) + + +def test_fused_weighted_restore_supports_double_backward(): + device = _device() + num_tokens, top_k, hidden = 4, 2, 16 + token_indices_sorted = torch.tensor([2, 0, 7, 1, 4, 6, 3, 5], device=device) + rows = torch.randn(num_tokens * top_k, hidden, device=device, dtype=torch.bfloat16, requires_grad=True) + scores = torch.rand(num_tokens, top_k, device=device, dtype=torch.float32, requires_grad=True) + + output = fused_ops.fused_weighted_restore( + rows, + top_scores=scores, + token_indices_sorted=token_indices_sorted, + top_k=top_k, + shape=(1, num_tokens, hidden), + ) + grad_rows, grad_scores = torch.autograd.grad(output.float().sum(), (rows, scores), create_graph=True) + score_cross_gradient = torch.autograd.grad(grad_rows.float().sum(), scores, retain_graph=True)[0] + row_cross_gradient = torch.autograd.grad(grad_scores.float().sum(), rows)[0] + + assert torch.isfinite(score_cross_gradient).all() + assert torch.isfinite(row_cross_gradient).all() + assert score_cross_gradient.abs().sum() > 0 + assert row_cross_gradient.abs().sum() > 0 + + def test_fused_engine_names_what_it_cannot_run(): device = _device() supported = torch.randn(8, 16, device=device, dtype=torch.bfloat16) From 248257b5ed7f7eff31575c0db0baf951449975e0 Mon Sep 17 00:00:00 2001 From: yh0903 Date: Wed, 26 Aug 2026 15:29:19 -0700 Subject: [PATCH 6/9] Avoid importing Triton on ROCm builds Keep the optional fused restore unavailable on HIP without importing pytorch-triton-rocm, matching DeepSpeed import-time device compatibility. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: yh0903 --- deepspeed/moe/autoep_fused_token_ops.py | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/deepspeed/moe/autoep_fused_token_ops.py b/deepspeed/moe/autoep_fused_token_ops.py index 99ab1290fdbe..d8d2b2e4e37e 100644 --- a/deepspeed/moe/autoep_fused_token_ops.py +++ b/deepspeed/moe/autoep_fused_token_ops.py @@ -26,13 +26,18 @@ import torch -try: - import triton - import triton.language as tl +_IS_ROCM_PYTORCH = getattr(torch.version, "hip", None) is not None - _TRITON_AVAILABLE = True -except ImportError: +if _IS_ROCM_PYTORCH: _TRITON_AVAILABLE = False +else: + try: + import triton + import triton.language as tl + + _TRITON_AVAILABLE = True + except ImportError: + _TRITON_AVAILABLE = False # The grouped GEMM produces the rows this consumes, so the supported dtypes are # the ones it is built for rather than a silent widening. From c02ea19b5fea6ecf98ca17560cd44d0a091363a5 Mon Sep 17 00:00:00 2001 From: yh0903 Date: Wed, 26 Aug 2026 16:04:49 -0700 Subject: [PATCH 7/9] Tighten fused restore comments Keep comments focused on numerical behavior, fail-fast ordering, and non-obvious kernel choices while removing repeated implementation narration. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: yh0903 --- deepspeed/module_inject/auto_ep_config.py | 5 +- deepspeed/module_inject/auto_ep_layer.py | 8 +-- deepspeed/moe/autoep_fused_token_ops.py | 57 ++++--------------- tests/unit/v1/moe/test_autoep_fused_parity.py | 15 +---- .../v1/moe/test_autoep_fused_token_ops.py | 13 +---- 5 files changed, 18 insertions(+), 80 deletions(-) diff --git a/deepspeed/module_inject/auto_ep_config.py b/deepspeed/module_inject/auto_ep_config.py index dfc3868a771a..547119549ec9 100644 --- a/deepspeed/module_inject/auto_ep_config.py +++ b/deepspeed/module_inject/auto_ep_config.py @@ -118,10 +118,7 @@ def validate_autoep_config( if not config.enabled: return - # The fused reduction only replaces the weighted sum the standard - # expert-parallel path runs. Where it has nothing to replace, say so instead - # of running the eager reduction under a config that asked for the fused - # one: a benchmark believing it measured the fused path would report noise. + # Reject configurations that would bypass the requested fused reduction. if config.combine_impl == "fused_weighted_sum": if tp_size > 1: raise ValueError('combine_impl="fused_weighted_sum" does not support folded tensor parallelism ' diff --git a/deepspeed/module_inject/auto_ep_layer.py b/deepspeed/module_inject/auto_ep_layer.py index 65f907816762..428bff8e34a2 100644 --- a/deepspeed/module_inject/auto_ep_layer.py +++ b/deepspeed/module_inject/auto_ep_layer.py @@ -554,9 +554,7 @@ def set_deepspeed_parallelism( if folding_group_handles is not None: self.folding_group_handles = folding_group_handles if self.combine_impl == "fused_weighted_sum" and folding_group_handles.spec.tp_size > 1: - # Folded TP restores combined tokens from assignment metadata and - # never reaches the weighted reduction this replaces, so the - # request would otherwise be silently ignored. + # Folded TP restores tokens through a different path. raise ValueError('combine_impl="fused_weighted_sum" does not support folded tensor parallelism ' f"(tensor_parallel.autotp_size={folding_group_handles.spec.tp_size}). Set " 'tensor_parallel.autotp_size to 1, or leave combine_impl unset.') @@ -598,9 +596,7 @@ def forward( bsz, seqlen, hdim = hidden_states.shape x = hidden_states.reshape(-1, hdim) # [T, H] - # Checked once, ahead of the router and of every collective, so an - # unsupported configuration fails on all ranks together rather than - # stalling the ones that carried on. + # Fail all ranks before any collective can stall. if self.combine_impl == "fused_weighted_sum" and not self._fused_combine_checked: fused_token_ops.assert_supported(x, score_apply=self.score_apply) self._fused_combine_checked = True diff --git a/deepspeed/moe/autoep_fused_token_ops.py b/deepspeed/moe/autoep_fused_token_ops.py index d8d2b2e4e37e..9b2b9512ab28 100644 --- a/deepspeed/moe/autoep_fused_token_ops.py +++ b/deepspeed/moe/autoep_fused_token_ops.py @@ -2,24 +2,10 @@ # SPDX-License-Identifier: Apache-2.0 # DeepSpeed Team -"""Fused weighted token restoration for AutoEP. - -After the combine all-to-all, the eager path returns one row per routed -assignment and turns it back into one row per token in general-purpose steps: it -scatters the rows into a zero-filled ``[tokens * top_k, hidden]`` buffer, views -that as ``[tokens, top_k, hidden]``, widens it to FP32 to apply routing weights, -and reduces over top-k. The FP32 intermediate alone is 64 MiB at the canonical -shape, and every step of that sequence costs a full pass over the routed -activations, in every MoE layer, on every step. - -This module does the same arithmetic in one pass. Each program owns one token -and one slice of the hidden dimension, walks its top-k rows in registers, -accumulates in FP32 and writes the token's output once, so neither the scattered -assignment buffer nor the FP32 intermediate is ever allocated. - -Only the reduction is replaced. The collectives, the router, the grouped GEMM -and the expert-major reorder are all untouched, so a measured difference belongs -to the reduction alone. +"""Fused AutoEP token restore without the eager scatter and FP32 intermediate. + +The kernel reduces each token's top-k rows in FP32. Communication, routing, +expert reorder, and grouped GEMM remain unchanged. """ from __future__ import annotations @@ -39,8 +25,6 @@ except ImportError: _TRITON_AVAILABLE = False -# The grouped GEMM produces the rows this consumes, so the supported dtypes are -# the ones it is built for rather than a silent widening. SUPPORTED_ROW_DTYPES = (torch.bfloat16, torch.float16) _MAX_BLOCK_HIDDEN = 512 @@ -100,8 +84,7 @@ def _weighted_restore_forward_kernel( other=0.0, ).to(tl.float32) - # FP32 product and reduction with a single cast on the way out, matching - # the dtype discipline of the eager weighted sum. + # Match the eager path's FP32 product and accumulation. weighted = tl.sum(values * scores[:, None], axis=0) tl.store( out_ptr + token * out_stride + hidden_offsets, @@ -138,10 +121,7 @@ def _weighted_restore_backward_kernel( scores = tl.load(scores_ptr + token * scores_stride + slots, mask=slot_mask, other=0.0).to(tl.float32) grad_rows_dtype = grad_rows_ptr.dtype.element_ty - # One token per program, so the score gradient reduces over the hidden - # dimension in registers. Splitting that dimension across programs and - # reducing the partials afterwards measured slower at this shape: the - # extra pass costs more than the added parallelism returns. + # Keeping one token per program avoids a second reduction pass for scores. score_partials = tl.zeros([K_PADDED, BLOCK_H], dtype=tl.float32) for hidden_start in range(0, hidden, BLOCK_H): @@ -182,11 +162,7 @@ def is_available() -> bool: def assert_supported(rows: torch.Tensor, *, score_apply: str) -> None: - """Reject configurations the fused restore does not implement. - - Checked before any collective runs: a rank that raised while its peers - proceeded would turn a clear error into a hang. - """ + """Reject unsupported configurations before collectives begin.""" if not _TRITON_AVAILABLE: raise RuntimeError('combine_impl="fused_weighted_sum" needs Triton, which is not installed in this ' "environment. Install Triton, or leave combine_impl unset.") @@ -203,11 +179,7 @@ def assert_supported(rows: torch.Tensor, *, score_apply: str) -> None: def _block_hidden(hidden: int, slots: int) -> int: - """Pick a power-of-two hidden tile that fits alongside ``slots`` rows of FP32. - - The floor keeps the budget honest for top-k values far wider than any real - router, so the tile shrinks rather than overrunning the element budget. - """ + """Choose a power-of-two tile within the FP32 register budget.""" budget = max(16, _MAX_BLOCK_ELEMENTS // slots) return min(_MAX_BLOCK_HIDDEN, budget, max(16, triton.next_power_of_2(hidden))) @@ -293,16 +265,12 @@ def backward(ctx, grad_output): ctx.top_k) return grad_rows, grad_scores, None, None - # AutoEP supplies an exact permutation, so every row is written once. - # Zero initialization also keeps malformed direct calls deterministic - # when an invalid or duplicate assignment leaves an inverse slot empty. + # Zero initialization keeps malformed direct calls deterministic. grad_rows = torch.zeros_like(combined_rows) grad_scores = torch.empty_like(top_scores) n_tokens, hidden = top_scores.shape[0], combined_rows.shape[-1] if n_tokens == 0 or hidden == 0: - # Nothing is reduced, so the score gradient is zero rather than - # whatever an uninitialized buffer happened to hold. return grad_rows, torch.zeros_like(top_scores), None, None k_padded = _padded_top_k(ctx.top_k) @@ -333,12 +301,7 @@ def fused_weighted_restore( top_k: int, shape: tuple[int, int, int], ) -> torch.Tensor: - """Weight combined rows by their routing scores and reduce over top-k. - - Fused counterpart of ``combine_from_routed`` for ``score_apply="post"``. It - goes straight from ``[T * K, H]`` to ``[B, S, H]``, so neither the scattered - assignment buffer nor the ``[T, K, H]`` FP32 intermediate is allocated. - """ + """Restore ``[T * K, H]`` rows directly to weighted ``[B, S, H]`` output.""" bsz, seqlen, hidden = shape if top_k <= 0: raise RuntimeError(f"fused weighted restore expects top_k > 0, got {top_k}.") diff --git a/tests/unit/v1/moe/test_autoep_fused_parity.py b/tests/unit/v1/moe/test_autoep_fused_parity.py index f480cdb77524..abc8fcee5eda 100644 --- a/tests/unit/v1/moe/test_autoep_fused_parity.py +++ b/tests/unit/v1/moe/test_autoep_fused_parity.py @@ -2,12 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 # DeepSpeed Team -"""End-to-end parity between the eager and fused combine implementations. - -The fused reduction only changes how the weighted sum is computed, so a step -taken through it has to produce the same loss, the same gradients on every -trainable tensor, and the same parameter update as the eager path. -""" +"""End-to-end parity for the eager and fused combine implementations.""" import functools @@ -31,10 +26,7 @@ SEQ_LEN = 16 NUM_EXPERTS = 4 -# Both backends run the same collectives, the same router and the same grouped -# GEMM; they differ only in the order the top-k reduction accumulates. That -# survives a full step as a last-few-bits difference, not a structural one, so -# the tolerance stays far tighter than a wrong permutation could hide behind. +# The top-k accumulation order differs, so parity allows last-bit noise only. PARITY_TOLERANCE = {"rtol": 1e-2, "atol": 1e-3} @@ -140,8 +132,7 @@ def _assert_step_matches(fused, eager): assert fused["gradients"], "no gradients were captured, so the comparison would be vacuous" assert set(fused["gradients"]) == set(eager["gradients"]) - # Router and expert gradients travel different routes through the fused - # restore, so they are named rather than left to a bulk comparison. + # Ensure both sides of the fused restore reached the comparison. assert any(".router." in name for name in fused["gradients"]), "no router gradient was captured" assert any(".experts.w" in name for name in fused["gradients"]), "no expert gradient was captured" diff --git a/tests/unit/v1/moe/test_autoep_fused_token_ops.py b/tests/unit/v1/moe/test_autoep_fused_token_ops.py index 383f97615bb3..c1bf7e6cb6f8 100644 --- a/tests/unit/v1/moe/test_autoep_fused_token_ops.py +++ b/tests/unit/v1/moe/test_autoep_fused_token_ops.py @@ -2,12 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 # DeepSpeed Team -"""The fused weighted restore against the eager reduction it replaces. - -``combine_from_routed`` is the reference: the fused reduction is only worth -having if it is indistinguishable from it, so every assertion compares the two -directly rather than against hand-written expectations. -""" +"""Compare the fused weighted restore with the eager reference.""" import pytest import torch @@ -40,7 +35,6 @@ def test_fused_weighted_restore_matches_eager_including_gradients(top_k, hidden, selected_experts = torch.randint(0, num_experts, (num_tokens, top_k), device=device, generator=generator) token_indices_sorted = torch.argsort(selected_experts.view(-1), stable=True) - # A restore that only had to undo the identity would not exercise anything. assert not torch.equal(token_indices_sorted, torch.arange(num_tokens * top_k, device=device)) rows = torch.randn(num_tokens * top_k, hidden, device=device, dtype=torch.bfloat16, generator=generator) @@ -75,10 +69,7 @@ def test_fused_weighted_restore_matches_eager_including_gradients(top_k, hidden, fused_output.backward(upstream) torch.testing.assert_close(fused_rows.grad, eager_rows.grad) - # The score gradient reduces over the hidden dimension, so the fused and eager - # summation orders differ even though both accumulate in FP32. That shows up - # in FP32 scores; a bfloat16 score rounds the difference away, and asking for - # FP32 precision there would fail on a rounding boundary rather than on a bug. + # Hidden reduction order only affects the last bits of FP32 score gradients. score_tolerance = {"rtol": 1e-4, "atol": 1e-5} if score_dtype == torch.float32 else {} torch.testing.assert_close(fused_scores.grad, eager_scores.grad, **score_tolerance) From f6e2dfd1e0535f5c8ee8a1328197381c5f3502a1 Mon Sep 17 00:00:00 2001 From: yh0903 Date: Sat, 29 Aug 2026 00:35:41 -0700 Subject: [PATCH 8/9] Address fused restore review feedback Signed-off-by: yh0903 --- deepspeed/module_inject/auto_ep_layer.py | 2 +- .../triton_ops}/autoep_fused_token_ops.py | 68 ++++--------------- tests/unit/v1/moe/test_autoep_fused_parity.py | 2 +- .../test_autoep_fused_token_ops.py | 58 ++++------------ 4 files changed, 27 insertions(+), 103 deletions(-) rename deepspeed/{moe => ops/triton_ops}/autoep_fused_token_ops.py (77%) rename tests/unit/v1/{moe => ops/triton_ops}/test_autoep_fused_token_ops.py (69%) diff --git a/deepspeed/module_inject/auto_ep_layer.py b/deepspeed/module_inject/auto_ep_layer.py index 428bff8e34a2..bd0080df3e1f 100644 --- a/deepspeed/module_inject/auto_ep_layer.py +++ b/deepspeed/module_inject/auto_ep_layer.py @@ -21,8 +21,8 @@ import deepspeed.comm as dist from deepspeed.module_inject.auto_ep_config import AutoEPConfig, MoELayerSpec, resolve_autoep_config_defaults from deepspeed.module_inject.auto_ep_folding import mark_autoep_folding_router_parameter +from deepspeed.ops.triton_ops import autoep_fused_token_ops as fused_token_ops from deepspeed.utils import logger -from deepspeed.moe import autoep_fused_token_ops as fused_token_ops from deepspeed.moe.ep_router import TokenChoiceTopKRouter from deepspeed.moe.ep_count import count_tokens_per_expert from deepspeed.moe.ep_experts import GroupedExperts diff --git a/deepspeed/moe/autoep_fused_token_ops.py b/deepspeed/ops/triton_ops/autoep_fused_token_ops.py similarity index 77% rename from deepspeed/moe/autoep_fused_token_ops.py rename to deepspeed/ops/triton_ops/autoep_fused_token_ops.py index 9b2b9512ab28..b0a222aca427 100644 --- a/deepspeed/moe/autoep_fused_token_ops.py +++ b/deepspeed/ops/triton_ops/autoep_fused_token_ops.py @@ -1,6 +1,4 @@ -# Copyright (c) DeepSpeed Team. # SPDX-License-Identifier: Apache-2.0 - # DeepSpeed Team """Fused AutoEP token restore without the eager scatter and FP32 intermediate. @@ -12,20 +10,11 @@ import torch -_IS_ROCM_PYTORCH = getattr(torch.version, "hip", None) is not None - -if _IS_ROCM_PYTORCH: - _TRITON_AVAILABLE = False -else: - try: - import triton - import triton.language as tl +from deepspeed.ops.triton_ops._triton import _TRITON_AVAILABLE, triton, tl - _TRITON_AVAILABLE = True - except ImportError: - _TRITON_AVAILABLE = False +_IS_ROCM_PYTORCH = getattr(torch.version, "hip", None) is not None -SUPPORTED_ROW_DTYPES = (torch.bfloat16, torch.float16) +SUPPORTED_ROW_DTYPES = (torch.bfloat16, torch.float16, torch.float32) _MAX_BLOCK_HIDDEN = 512 _INVERT_INDEX_BLOCK = 256 @@ -158,7 +147,7 @@ def _weighted_restore_backward_kernel( def is_available() -> bool: """Whether this build can run the fused weighted restore at all.""" - return _TRITON_AVAILABLE + return _TRITON_AVAILABLE and not _IS_ROCM_PYTORCH def assert_supported(rows: torch.Tensor, *, score_apply: str) -> None: @@ -166,12 +155,15 @@ def assert_supported(rows: torch.Tensor, *, score_apply: str) -> None: if not _TRITON_AVAILABLE: raise RuntimeError('combine_impl="fused_weighted_sum" needs Triton, which is not installed in this ' "environment. Install Triton, or leave combine_impl unset.") + if _IS_ROCM_PYTORCH: + raise RuntimeError('combine_impl="fused_weighted_sum" is not yet supported on ROCm. Leave combine_impl ' + "unset to run here.") if rows.device.type != "cuda": raise RuntimeError('combine_impl="fused_weighted_sum" runs CUDA kernels but this layer is on device ' f'"{rows.device.type}". Leave combine_impl unset to run here.') if rows.dtype not in SUPPORTED_ROW_DTYPES: - raise RuntimeError('combine_impl="fused_weighted_sum" supports bfloat16 and float16 rows, got ' - f"{rows.dtype}. Leave combine_impl unset, or train in bf16/fp16.") + raise RuntimeError('combine_impl="fused_weighted_sum" supports bfloat16, float16, and float32 rows, got ' + f"{rows.dtype}. Leave combine_impl unset, or use a supported floating-point dtype.") if score_apply != "post": raise RuntimeError('combine_impl="fused_weighted_sum" folds the routing weight into the top-k reduction, ' f'which only exists for score_apply="post", but this layer resolved ' @@ -190,11 +182,9 @@ def _padded_top_k(top_k: int) -> int: def _invert_index(index: torch.Tensor, num_inverse_rows: int) -> torch.Tensor: - """Invert a row permutation, leaving -1 wherever no slot claimed a row.""" - inverse = torch.full((num_inverse_rows, ), -1, dtype=torch.int32, device=index.device) + """Invert the row permutation produced by sorting routed assignments.""" + inverse = torch.empty((num_inverse_rows, ), dtype=torch.int32, device=index.device) num_indices = index.numel() - if num_indices == 0 or num_inverse_rows == 0: - return inverse grid = (triton.cdiv(num_indices, _INVERT_INDEX_BLOCK), ) _invert_index_kernel[grid]( @@ -234,8 +224,6 @@ def forward(ctx, combined_rows, top_scores, inverse, top_k): ctx.save_for_backward(combined_rows, top_scores, inverse) ctx.top_k = top_k - if n_tokens == 0 or hidden == 0: - return output k_padded = _padded_top_k(top_k) block_hidden = _block_hidden(hidden, slots=k_padded) @@ -265,13 +253,10 @@ def backward(ctx, grad_output): ctx.top_k) return grad_rows, grad_scores, None, None - # Zero initialization keeps malformed direct calls deterministic. - grad_rows = torch.zeros_like(combined_rows) + grad_rows = torch.empty_like(combined_rows) grad_scores = torch.empty_like(top_scores) n_tokens, hidden = top_scores.shape[0], combined_rows.shape[-1] - if n_tokens == 0 or hidden == 0: - return grad_rows, torch.zeros_like(top_scores), None, None k_padded = _padded_top_k(ctx.top_k) _weighted_restore_backward_kernel[(n_tokens, )]( @@ -303,37 +288,8 @@ def fused_weighted_restore( ) -> torch.Tensor: """Restore ``[T * K, H]`` rows directly to weighted ``[B, S, H]`` output.""" bsz, seqlen, hidden = shape - if top_k <= 0: - raise RuntimeError(f"fused weighted restore expects top_k > 0, got {top_k}.") - if bsz < 0 or seqlen < 0 or hidden < 0: - raise RuntimeError(f"fused weighted restore expects non-negative output dimensions, got {shape}.") - if combined_rows.ndim != 2: - raise RuntimeError(f"fused weighted restore expects combined_rows to be 2D, got shape " - f"{tuple(combined_rows.shape)}.") - if combined_rows.shape[1] != hidden: - raise RuntimeError(f"fused weighted restore output hidden size is {hidden}, but combined rows have hidden " - f"size {combined_rows.shape[1]}.") - n_tokens = bsz * seqlen expected_rows = n_tokens * top_k - if combined_rows.shape[0] != expected_rows: - raise RuntimeError(f"fused weighted restore expects one row per assignment: {expected_rows} rows for " - f"{n_tokens} tokens at top_k={top_k}, got {combined_rows.shape[0]}.") - if tuple(top_scores.shape) != (n_tokens, top_k): - raise RuntimeError(f"fused weighted restore expects top_scores shape {(n_tokens, top_k)}, got " - f"{tuple(top_scores.shape)}.") - if token_indices_sorted.ndim != 1 or token_indices_sorted.numel() != expected_rows: - raise RuntimeError(f"fused weighted restore expects token_indices_sorted to contain {expected_rows} " - f"assignments, got shape {tuple(token_indices_sorted.shape)}.") - if token_indices_sorted.dtype not in (torch.int32, torch.int64): - raise RuntimeError("fused weighted restore expects token_indices_sorted to use int32 or int64 indices, got " - f"{token_indices_sorted.dtype}.") - if not torch.is_floating_point(top_scores): - raise RuntimeError(f"fused weighted restore expects floating-point top_scores, got {top_scores.dtype}.") - if combined_rows.device != top_scores.device or combined_rows.device != token_indices_sorted.device: - raise RuntimeError("fused weighted restore expects rows, scores, and indices on the same device, got " - f"{combined_rows.device}, {top_scores.device}, and {token_indices_sorted.device}.") - inverse = _invert_index(token_indices_sorted, expected_rows) output = _FusedWeightedRestore.apply(combined_rows, top_scores.contiguous(), inverse, top_k) return output.reshape(bsz, seqlen, hidden) diff --git a/tests/unit/v1/moe/test_autoep_fused_parity.py b/tests/unit/v1/moe/test_autoep_fused_parity.py index abc8fcee5eda..00211912ae0e 100644 --- a/tests/unit/v1/moe/test_autoep_fused_parity.py +++ b/tests/unit/v1/moe/test_autoep_fused_parity.py @@ -11,8 +11,8 @@ import torch from deepspeed.accelerator import get_accelerator -from deepspeed.moe import autoep_fused_token_ops as fused_ops from deepspeed.module_inject.auto_ep_layer import AutoEPMoELayer +from deepspeed.ops.triton_ops import autoep_fused_token_ops as fused_ops from deepspeed.utils import safe_get_full_grad from unit.common import DistributedTest from unit.v1.moe.autoep_test_utils import ( diff --git a/tests/unit/v1/moe/test_autoep_fused_token_ops.py b/tests/unit/v1/ops/triton_ops/test_autoep_fused_token_ops.py similarity index 69% rename from tests/unit/v1/moe/test_autoep_fused_token_ops.py rename to tests/unit/v1/ops/triton_ops/test_autoep_fused_token_ops.py index c1bf7e6cb6f8..082c6bccf495 100644 --- a/tests/unit/v1/moe/test_autoep_fused_token_ops.py +++ b/tests/unit/v1/ops/triton_ops/test_autoep_fused_token_ops.py @@ -1,6 +1,4 @@ -# Copyright (c) DeepSpeed Team. # SPDX-License-Identifier: Apache-2.0 - # DeepSpeed Team """Compare the fused weighted restore with the eager reference.""" @@ -8,8 +6,8 @@ import torch from deepspeed.accelerator import get_accelerator -from deepspeed.moe import autoep_fused_token_ops as fused_ops from deepspeed.module_inject.auto_ep_layer import combine_from_routed +from deepspeed.ops.triton_ops import autoep_fused_token_ops as fused_ops def _fused_engine_available(): @@ -27,8 +25,9 @@ def _device(): @pytest.mark.parametrize("top_k", [2, 4, 6, 8]) @pytest.mark.parametrize("hidden", [128, 130]) +@pytest.mark.parametrize("row_dtype", [torch.float32, torch.float16, torch.bfloat16]) @pytest.mark.parametrize("score_dtype", [torch.float32, torch.bfloat16]) -def test_fused_weighted_restore_matches_eager_including_gradients(top_k, hidden, score_dtype): +def test_fused_weighted_restore_matches_eager_including_gradients(top_k, hidden, row_dtype, score_dtype): device = _device() num_tokens, num_experts = 24, 8 generator = torch.Generator(device=device).manual_seed(20260824) @@ -37,9 +36,9 @@ def test_fused_weighted_restore_matches_eager_including_gradients(top_k, hidden, token_indices_sorted = torch.argsort(selected_experts.view(-1), stable=True) assert not torch.equal(token_indices_sorted, torch.arange(num_tokens * top_k, device=device)) - rows = torch.randn(num_tokens * top_k, hidden, device=device, dtype=torch.bfloat16, generator=generator) + rows = torch.randn(num_tokens * top_k, hidden, device=device, dtype=row_dtype, generator=generator) scores = torch.rand(num_tokens, top_k, device=device, dtype=score_dtype, generator=generator) - upstream = torch.randn(1, num_tokens, hidden, device=device, dtype=torch.bfloat16, generator=generator) + upstream = torch.randn(1, num_tokens, hidden, device=device, dtype=row_dtype, generator=generator) eager_rows = rows.clone().requires_grad_(True) eager_scores = scores.clone().requires_grad_(True) @@ -63,50 +62,18 @@ def test_fused_weighted_restore_matches_eager_including_gradients(top_k, hidden, shape=(1, num_tokens, hidden), ) - torch.testing.assert_close(fused_output, eager_output) + output_tolerance = {"rtol": 1e-5, "atol": 1e-6} if row_dtype == torch.float32 else {} + torch.testing.assert_close(fused_output, eager_output, **output_tolerance) eager_output.backward(upstream) fused_output.backward(upstream) - torch.testing.assert_close(fused_rows.grad, eager_rows.grad) + torch.testing.assert_close(fused_rows.grad, eager_rows.grad, **output_tolerance) # Hidden reduction order only affects the last bits of FP32 score gradients. score_tolerance = {"rtol": 1e-4, "atol": 1e-5} if score_dtype == torch.float32 else {} torch.testing.assert_close(fused_scores.grad, eager_scores.grad, **score_tolerance) -def test_fused_weighted_restore_requires_one_row_per_assignment(): - device = _device() - with pytest.raises(RuntimeError, match="one row per assignment"): - fused_ops.fused_weighted_restore( - torch.randn(10, 16, device=device, dtype=torch.bfloat16), - top_scores=torch.rand(4, 2, device=device), - token_indices_sorted=torch.arange(8, device=device), - top_k=2, - shape=(1, 4, 16), - ) - - -@pytest.mark.parametrize( - "rows_shape,scores_shape,index_count,index_dtype,error", - [ - ((8, 15), (4, 2), 8, torch.int64, "output hidden size"), - ((8, 16), (4, 1), 8, torch.int64, "top_scores shape"), - ((8, 16), (4, 2), 7, torch.int64, "token_indices_sorted"), - ((8, 16), (4, 2), 8, torch.float32, "int32 or int64"), - ], -) -def test_fused_weighted_restore_validates_input_contract(rows_shape, scores_shape, index_count, index_dtype, error): - device = _device() - with pytest.raises(RuntimeError, match=error): - fused_ops.fused_weighted_restore( - torch.randn(rows_shape, device=device, dtype=torch.bfloat16), - top_scores=torch.rand(scores_shape, device=device), - token_indices_sorted=torch.arange(index_count, device=device).to(index_dtype), - top_k=2, - shape=(1, 4, 16), - ) - - def test_fused_weighted_restore_supports_double_backward(): device = _device() num_tokens, top_k, hidden = 4, 2, 16 @@ -133,12 +100,13 @@ def test_fused_weighted_restore_supports_double_backward(): def test_fused_engine_names_what_it_cannot_run(): device = _device() - supported = torch.randn(8, 16, device=device, dtype=torch.bfloat16) - fused_ops.assert_supported(supported, score_apply="post") + for dtype in fused_ops.SUPPORTED_ROW_DTYPES: + fused_ops.assert_supported(torch.randn(8, 16, device=device, dtype=dtype), score_apply="post") - with pytest.raises(RuntimeError, match="bfloat16 and float16"): - fused_ops.assert_supported(torch.randn(8, 16, device=device, dtype=torch.float32), score_apply="post") + with pytest.raises(RuntimeError, match="bfloat16, float16, and float32"): + fused_ops.assert_supported(torch.randn(8, 16, device=device, dtype=torch.float64), score_apply="post") + supported = torch.randn(8, 16, device=device, dtype=torch.bfloat16) with pytest.raises(RuntimeError, match='resolved score_apply="pre"'): fused_ops.assert_supported(supported, score_apply="pre") From 52a23734cfe9540b0e8bd331adcfc4a6878bd011 Mon Sep 17 00:00:00 2001 From: yh0903 Date: Sat, 29 Aug 2026 11:43:52 -0700 Subject: [PATCH 9/9] Remove unused fused restore double backward Signed-off-by: yh0903 --- .../ops/triton_ops/autoep_fused_token_ops.py | 21 ---------------- .../triton_ops/test_autoep_fused_token_ops.py | 24 ------------------- 2 files changed, 45 deletions(-) diff --git a/deepspeed/ops/triton_ops/autoep_fused_token_ops.py b/deepspeed/ops/triton_ops/autoep_fused_token_ops.py index b0a222aca427..1036eb500608 100644 --- a/deepspeed/ops/triton_ops/autoep_fused_token_ops.py +++ b/deepspeed/ops/triton_ops/autoep_fused_token_ops.py @@ -197,22 +197,6 @@ def _invert_index(index: torch.Tensor, num_inverse_rows: int) -> torch.Tensor: return inverse -def _differentiable_backward(grad_output, combined_rows, top_scores, inverse, top_k): - """Build the rare higher-order backward with regular PyTorch operations.""" - n_tokens, hidden = top_scores.shape[0], combined_rows.shape[-1] - valid = inverse >= 0 - safe_inverse = inverse.clamp_min(0).to(torch.int64) - gathered_rows = combined_rows.index_select(0, safe_inverse).reshape(n_tokens, top_k, hidden) - - grad_by_assignment = (grad_output[:, None, :] * top_scores[:, :, None]).to(combined_rows.dtype).reshape(-1, hidden) - grad_rows = torch.zeros_like(combined_rows) - grad_rows = grad_rows.index_copy(0, safe_inverse[valid], grad_by_assignment[valid]) - - grad_scores = (gathered_rows.float() * grad_output.float()[:, None, :]).sum(dim=-1) - grad_scores = torch.where(valid.reshape(n_tokens, top_k), grad_scores, 0.0).to(top_scores.dtype) - return grad_rows, grad_scores - - class _FusedWeightedRestore(torch.autograd.Function): """Weight rows by their routing score and reduce over top-k in one pass.""" @@ -248,11 +232,6 @@ def backward(ctx, grad_output): combined_rows, top_scores, inverse = ctx.saved_tensors grad_output = grad_output.contiguous() - if torch.is_grad_enabled(): - grad_rows, grad_scores = _differentiable_backward(grad_output, combined_rows, top_scores, inverse, - ctx.top_k) - return grad_rows, grad_scores, None, None - grad_rows = torch.empty_like(combined_rows) grad_scores = torch.empty_like(top_scores) diff --git a/tests/unit/v1/ops/triton_ops/test_autoep_fused_token_ops.py b/tests/unit/v1/ops/triton_ops/test_autoep_fused_token_ops.py index 082c6bccf495..2e065c4f7a08 100644 --- a/tests/unit/v1/ops/triton_ops/test_autoep_fused_token_ops.py +++ b/tests/unit/v1/ops/triton_ops/test_autoep_fused_token_ops.py @@ -74,30 +74,6 @@ def test_fused_weighted_restore_matches_eager_including_gradients(top_k, hidden, torch.testing.assert_close(fused_scores.grad, eager_scores.grad, **score_tolerance) -def test_fused_weighted_restore_supports_double_backward(): - device = _device() - num_tokens, top_k, hidden = 4, 2, 16 - token_indices_sorted = torch.tensor([2, 0, 7, 1, 4, 6, 3, 5], device=device) - rows = torch.randn(num_tokens * top_k, hidden, device=device, dtype=torch.bfloat16, requires_grad=True) - scores = torch.rand(num_tokens, top_k, device=device, dtype=torch.float32, requires_grad=True) - - output = fused_ops.fused_weighted_restore( - rows, - top_scores=scores, - token_indices_sorted=token_indices_sorted, - top_k=top_k, - shape=(1, num_tokens, hidden), - ) - grad_rows, grad_scores = torch.autograd.grad(output.float().sum(), (rows, scores), create_graph=True) - score_cross_gradient = torch.autograd.grad(grad_rows.float().sum(), scores, retain_graph=True)[0] - row_cross_gradient = torch.autograd.grad(grad_scores.float().sum(), rows)[0] - - assert torch.isfinite(score_cross_gradient).all() - assert torch.isfinite(row_cross_gradient).all() - assert score_cross_gradient.abs().sum() > 0 - assert row_cross_gradient.abs().sum() > 0 - - def test_fused_engine_names_what_it_cannot_run(): device = _device() for dtype in fused_ops.SUPPORTED_ROW_DTYPES: