diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/__init__.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/__init__.py deleted file mode 100644 index d7a782cd5..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/__init__.py +++ /dev/null @@ -1,52 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""MNNVL CuTe DSL AllReduce fusion backend internals.""" - -from importlib import import_module - -from .config import ( - KernelTarget, - MNNVLCuteDSLConfig, - MRangeDispatch, - ProtocolKind, - StaticProfile, -) - - -def __getattr__(name: str): - if name in { - "BT_ONLY_CONFIG", - "DEFAULT_CONFIG", - "HT_ONLY_CONFIG", - "LL_ONLY_CONFIG", - }: - presets = import_module(f"{__name__}.presets") - value = getattr(presets, name) - globals()[name] = value - return value - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") - - -__all__ = [ - "BT_ONLY_CONFIG", - "DEFAULT_CONFIG", - "HT_ONLY_CONFIG", - "LL_ONLY_CONFIG", - "KernelTarget", - "MNNVLCuteDSLConfig", - "MRangeDispatch", - "ProtocolKind", - "StaticProfile", -] diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/config.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/config.py deleted file mode 100644 index 60bb7a9ba..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/config.py +++ /dev/null @@ -1,231 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Configuration and token-count routing for the MNNVL CuTe DSL backend.""" - -from __future__ import annotations - -from bisect import bisect_left -from dataclasses import dataclass -from enum import Enum -from typing import Generic, TypeVar - -import torch - -__all__ = [ - "KernelTarget", - "MNNVLCuteDSLConfig", - "MRangeDispatch", - "ProtocolKind", - "StaticProfile", -] - - -class ProtocolKind(Enum): - LL = "ll" - BT = "bt" - HT = "ht" - - -PresetT = TypeVar("PresetT") - - -@dataclass(frozen=True, slots=True) -class KernelTarget(Generic[PresetT]): - protocol: ProtocolKind - preset: PresetT - - -TargetT = TypeVar("TargetT") - - -@dataclass(frozen=True, slots=True) -class MRangeDispatch(Generic[TargetT]): - """Map contiguous positive token-count ranges to kernel targets.""" - - upper_bounds: tuple[int | None, ...] - targets: tuple[TargetT, ...] - - def __post_init__(self) -> None: - if not self.upper_bounds: - raise ValueError("M range dispatch must contain at least one range") - if len(self.upper_bounds) != len(self.targets): - raise ValueError("M range upper bounds and targets must have equal length") - - previous = 0 - for index, upper_bound in enumerate(self.upper_bounds): - if upper_bound is None: - if index != len(self.upper_bounds) - 1: - raise ValueError("An unbounded M range must be the final range") - continue - if upper_bound <= previous: - raise ValueError("M range upper bounds must be strictly increasing") - previous = upper_bound - - @property - def is_unbounded(self) -> bool: - return self.upper_bounds[-1] is None - - @property - def finite_upper_bound(self) -> int | None: - return None if self.is_unbounded else self.upper_bounds[-1] - - def supports(self, m: int) -> bool: - if m <= 0: - return False - upper_bound = self.finite_upper_bound - return upper_bound is None or m <= upper_bound - - def select(self, m: int) -> TargetT: - if not self.supports(m): - raise ValueError(f"No kernel route supports M={m}") - - finite_bounds = tuple( - upper_bound for upper_bound in self.upper_bounds if upper_bound is not None - ) - index = bisect_left(finite_bounds, m) - return self.targets[index] - - def referenced_protocols(self) -> frozenset[ProtocolKind]: - protocols = { - target.protocol - for target in self.targets - if isinstance(target, KernelTarget) - } - return frozenset(protocols) - - def targets_for_capacity(self, capacity_m: int) -> tuple[TargetT, ...]: - if capacity_m <= 0: - return () - selected = [] - lower_bound = 1 - for upper_bound, target in zip(self.upper_bounds, self.targets, strict=True): - if lower_bound > capacity_m: - break - selected.append(target) - if upper_bound is None: - break - lower_bound = upper_bound + 1 - return tuple(selected) - - def max_m_for_protocol( - self, protocol: ProtocolKind, *, capacity_m: int - ) -> int | None: - lower_bound = 1 - maximum = None - for upper_bound, target in zip(self.upper_bounds, self.targets, strict=True): - effective_upper_bound = capacity_m if upper_bound is None else upper_bound - if ( - isinstance(target, KernelTarget) - and target.protocol is protocol - and lower_bound <= capacity_m - ): - maximum = min(effective_upper_bound, capacity_m) - lower_bound = effective_upper_bound + 1 - return maximum - - -@dataclass(frozen=True, slots=True) -class StaticProfile: - tp_size: int - hidden_size: int - top_k: int - dtype: torch.dtype - finalize_routes: MRangeDispatch[KernelTarget[object]] - all_reduce_routes: MRangeDispatch[KernelTarget[object]] - - def __post_init__(self) -> None: - if self.hidden_size <= 0 or self.hidden_size % 8: - raise ValueError("hidden_size must be a positive multiple of 8") - - def matches( - self, - *, - tp_size: int, - hidden_size: int, - top_k: int, - dtype: torch.dtype, - ) -> bool: - return ( - self.tp_size == tp_size - and self.hidden_size == hidden_size - and self.top_k == top_k - and self.dtype == dtype - ) - - def validate_capacity(self, capacity_m: int) -> None: - if capacity_m <= 0: - raise ValueError("capacity_m must be positive") - if not self.finalize_routes.supports(capacity_m): - raise ValueError( - "Finalize routes do not cover the requested workspace capacity" - ) - if not self.all_reduce_routes.supports(capacity_m): - raise ValueError( - "AllReduce routes do not cover the requested workspace capacity" - ) - - @property - def referenced_protocols(self) -> frozenset[ProtocolKind]: - return ( - self.finalize_routes.referenced_protocols() - | self.all_reduce_routes.referenced_protocols() - ) - - def protocol_capacity( - self, protocol: ProtocolKind, *, capacity_m: int - ) -> int | None: - maxima = ( - self.finalize_routes.max_m_for_protocol(protocol, capacity_m=capacity_m), - self.all_reduce_routes.max_m_for_protocol(protocol, capacity_m=capacity_m), - ) - present = tuple(value for value in maxima if value is not None) - return max(present) if present else None - - -@dataclass(frozen=True, slots=True) -class MNNVLCuteDSLConfig: - """Static profiles and routing policy for one backend configuration.""" - - profiles: tuple[StaticProfile, ...] - - def __post_init__(self) -> None: - keys = [ - (profile.tp_size, profile.hidden_size, profile.top_k, profile.dtype) - for profile in self.profiles - ] - if not keys: - raise ValueError("A backend config must contain at least one profile") - if len(keys) != len(set(keys)): - raise ValueError("Backend config profiles must have unique static shapes") - - def resolve( - self, - *, - tp_size: int, - hidden_size: int, - top_k: int, - dtype: torch.dtype, - capacity_m: int, - ) -> StaticProfile: - for profile in self.profiles: - if profile.matches( - tp_size=tp_size, - hidden_size=hidden_size, - top_k=top_k, - dtype=dtype, - ): - profile.validate_capacity(capacity_m) - return profile - raise ValueError("No MNNVL CuTe DSL profile supports this static shape") diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/cute_dsl_primitives.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/cute_dsl_primitives.py deleted file mode 100644 index 1b10ee6fb..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/cute_dsl_primitives.py +++ /dev/null @@ -1,876 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Small standalone CuTe DSL and PTX primitives shared by Kernel backends.""" - -import cutlass -import cutlass.cute as cute -from cutlass import BFloat16, Float32, Int32, Int64, Uint16, Uint32 -from cutlass._mlir import ir -from cutlass._mlir.dialects import llvm, vector -from cutlass.cutlass_dsl import T, dsl_user_op - -WARP_SIZE = 32 -VEC_BF16 = 8 -QUAD_BF16 = 4 -NEGATIVE_ZERO_BF16_BITS = 0x8000 -NEGATIVE_ZERO_BF16_PAIR = 0x80008000 -# CUTLASS cute::TMA::CacheHintSm100::EVICT_FIRST policy descriptor. -L2_EVICT_FIRST = 0x12F0000000000000 - - -@dsl_user_op -def load_global_u32x4( - pointer: cute.Pointer, - *, - volatile: cutlass.Constexpr[bool] = False, - loc=None, - ip=None, -): - address = pointer.toint(loc=loc, ip=ip) - if volatile: - opcode = "ld.volatile.global.v4.u32" - else: - opcode = "ld.global.v4.u32" - loaded = llvm.inline_asm( - llvm.StructType.get_literal([T.i32()] * 4), - [address.ir_value(loc=loc, ip=ip)], - f"{opcode} {{$0, $1, $2, $3}}, [$4];", - "=r,=r,=r,=r,l", - has_side_effects=volatile, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - packed = vector.from_elements( - ir.VectorType.get([4], T.i32(), loc=loc), - [ - llvm.extractvalue(T.i32(), loaded, [index], loc=loc, ip=ip) - for index in range(4) - ], - loc=loc, - ip=ip, - ) - return cute.TensorSSA(packed, 4, Uint32) - - -@dsl_user_op -def load_global_u32x4_predicated( - pointer: cute.Pointer, - predicate: Int32, - *, - loc=None, - ip=None, -): - address = pointer.toint(loc=loc, ip=ip) - loaded = llvm.inline_asm( - llvm.StructType.get_literal([T.i32()] * 4), - [ - address.ir_value(loc=loc, ip=ip), - Int32(predicate).ir_value(loc=loc, ip=ip), - ], - ( - "{\n\t" - ".reg .pred p;\n\t" - "setp.ne.s32 p, $5, 0;\n\t" - "@!p mov.u32 $0, 0;\n\t" - "@!p mov.u32 $1, 0;\n\t" - "@!p mov.u32 $2, 0;\n\t" - "@!p mov.u32 $3, 0;\n\t" - "@p ld.global.v4.u32 {$0, $1, $2, $3}, [$4];\n\t" - "}" - ), - "=r,=r,=r,=r,l,r", - has_side_effects=False, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - packed = vector.from_elements( - ir.VectorType.get([4], T.i32(), loc=loc), - [ - llvm.extractvalue(T.i32(), loaded, [index], loc=loc, ip=ip) - for index in range(4) - ], - loc=loc, - ip=ip, - ) - return cute.TensorSSA(packed, 4, Uint32) - - -@dsl_user_op -def load_global_u32(pointer: cute.Pointer, *, loc=None, ip=None) -> Uint32: - address = pointer.toint(loc=loc, ip=ip) - return Uint32( - llvm.inline_asm( - T.i32(), - [address.ir_value(loc=loc, ip=ip)], - "ld.global.u32 $0, [$1];", - "=r,l", - has_side_effects=False, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - ) - - -@dsl_user_op -def load_global_u32_predicated( - pointer: cute.Pointer, - predicate: Int32, - *, - loc=None, - ip=None, -) -> Uint32: - address = pointer.toint(loc=loc, ip=ip) - return Uint32( - llvm.inline_asm( - T.i32(), - [ - address.ir_value(loc=loc, ip=ip), - Int32(predicate).ir_value(loc=loc, ip=ip), - ], - ( - "{\n\t" - ".reg .pred p;\n\t" - "setp.ne.s32 p, $2, 0;\n\t" - "@!p mov.u32 $0, 0;\n\t" - "@p ld.global.u32 $0, [$1];\n\t" - "}" - ), - "=r,l,r", - has_side_effects=False, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - ) - - -@dsl_user_op -def load_global_u32x2(pointer: cute.Pointer, *, loc=None, ip=None): - address = pointer.toint(loc=loc, ip=ip) - loaded = llvm.inline_asm( - llvm.StructType.get_literal([T.i32()] * 2), - [address.ir_value(loc=loc, ip=ip)], - "ld.global.v2.u32 {$0, $1}, [$2];", - "=r,=r,l", - has_side_effects=False, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - packed = vector.from_elements( - ir.VectorType.get([2], T.i32(), loc=loc), - [ - llvm.extractvalue(T.i32(), loaded, [index], loc=loc, ip=ip) - for index in range(2) - ], - loc=loc, - ip=ip, - ) - return cute.TensorSSA(packed, 2, Uint32) - - -@dsl_user_op -def load_global_u32x2_predicated( - pointer: cute.Pointer, - predicate: Int32, - *, - loc=None, - ip=None, -): - address = pointer.toint(loc=loc, ip=ip) - loaded = llvm.inline_asm( - llvm.StructType.get_literal([T.i32()] * 2), - [ - address.ir_value(loc=loc, ip=ip), - Int32(predicate).ir_value(loc=loc, ip=ip), - ], - ( - "{\n\t" - ".reg .pred p;\n\t" - "setp.ne.s32 p, $3, 0;\n\t" - "@!p mov.u32 $0, 0;\n\t" - "@!p mov.u32 $1, 0;\n\t" - "@p ld.global.v2.u32 {$0, $1}, [$2];\n\t" - "}" - ), - "=r,=r,l,r", - has_side_effects=False, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - packed = vector.from_elements( - ir.VectorType.get([2], T.i32(), loc=loc), - [ - llvm.extractvalue(T.i32(), loaded, [index], loc=loc, ip=ip) - for index in range(2) - ], - loc=loc, - ip=ip, - ) - return cute.TensorSSA(packed, 2, Uint32) - - -@dsl_user_op -def store_global_u32x4(address: Int64, packed, *, loc=None, ip=None) -> None: - words = [packed[index].ir_value(loc=loc, ip=ip) for index in range(4)] - llvm.inline_asm( - None, - [address.ir_value(loc=loc, ip=ip), *words], - "st.global.v4.u32 [$0], {$1, $2, $3, $4};", - "l,r,r,r,r", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - - -@dsl_user_op -def store_global_u32_address( - address: Int64, - value: Uint32, - *, - loc=None, - ip=None, -) -> None: - llvm.inline_asm( - None, - [address.ir_value(loc=loc, ip=ip), value.ir_value(loc=loc, ip=ip)], - "st.global.u32 [$0], $1;", - "l,r", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - - -@dsl_user_op -def store_global_u32x2(address: Int64, packed, *, loc=None, ip=None) -> None: - words = [packed[index].ir_value(loc=loc, ip=ip) for index in range(2)] - llvm.inline_asm( - None, - [address.ir_value(loc=loc, ip=ip), *words], - "st.global.v2.u32 [$0], {$1, $2};", - "l,r,r", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - - -@dsl_user_op -def store_global_u16_bits( - address: Int64, - value: Uint32, - *, - loc=None, - ip=None, -) -> None: - llvm.inline_asm( - None, - [address.ir_value(loc=loc, ip=ip), value.ir_value(loc=loc, ip=ip)], - ( - "{\n\t" - ".reg .b16 bits;\n\t" - "cvt.u16.u32 bits, $1;\n\t" - "st.global.u16 [$0], bits;\n\t" - "}" - ), - "l,r", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - - -@dsl_user_op -def store_lamport_sentinel_u32x4( - address: Int64, - *, - loc=None, - ip=None, -) -> None: - sentinel = Uint32(NEGATIVE_ZERO_BF16_PAIR).ir_value(loc=loc, ip=ip) - llvm.inline_asm( - None, - [address.ir_value(loc=loc, ip=ip), sentinel, sentinel, sentinel, sentinel], - "st.global.v4.u32 [$0], {$1, $2, $3, $4};", - "l,r,r,r,r", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - - -@dsl_user_op -def load_global_bf16_as_f32( - address: Int64, - *, - loc=None, - ip=None, -) -> Float32: - return Float32( - llvm.inline_asm( - T.f32(), - [address.ir_value(loc=loc, ip=ip)], - ( - "{\n\t" - ".reg .b16 bits;\n\t" - "ld.global.b16 bits, [$1];\n\t" - "cvt.f32.bf16 $0, bits;\n\t" - "}" - ), - "=f,l", - has_side_effects=False, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - ) - - -@dsl_user_op -def load_global_bf16_as_f32_predicated( - address: Int64, - predicate: Int32, - *, - loc=None, - ip=None, -) -> Float32: - return Float32( - llvm.inline_asm( - T.f32(), - [ - address.ir_value(loc=loc, ip=ip), - Int32(predicate).ir_value(loc=loc, ip=ip), - ], - ( - "{\n\t" - ".reg .pred p;\n\t" - ".reg .b16 bits;\n\t" - "setp.ne.s32 p, $2, 0;\n\t" - "@!p mov.b16 bits, 0;\n\t" - "@p ld.global.b16 bits, [$1];\n\t" - "cvt.f32.bf16 $0, bits;\n\t" - "}" - ), - "=f,l,r", - has_side_effects=False, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - ) - - -@dsl_user_op -def f32_to_bf16_bits(value: Float32, *, loc=None, ip=None) -> Uint32: - return Uint32( - llvm.inline_asm( - T.i32(), - [value.ir_value(loc=loc, ip=ip)], - ( - "{\n\t" - ".reg .b16 bits;\n\t" - "cvt.rn.bf16.f32 bits, $1;\n\t" - "cvt.u32.u16 $0, bits;\n\t" - "}" - ), - "=r,f", - has_side_effects=False, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - ) - - -@dsl_user_op -def shuffle_sync_idx_u32( - value: Uint32, - source_lane: Int32, - *, - loc=None, - ip=None, -) -> Uint32: - return Uint32( - llvm.inline_asm( - T.i32(), - [ - value.ir_value(loc=loc, ip=ip), - source_lane.ir_value(loc=loc, ip=ip), - ], - "shfl.sync.idx.b32 $0, $1, $2, 0x1f, 0xffffffff;", - "=r,r,r", - # Preserve full-warp execution across later divergent consumers. - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - ) - - -@dsl_user_op -def load_volatile_u32(pointer: cute.Pointer, *, loc=None, ip=None) -> Uint32: - address = pointer.toint(loc=loc, ip=ip) - return Uint32( - llvm.inline_asm( - T.i32(), - [address.ir_value(loc=loc, ip=ip)], - "ld.volatile.global.u32 $0, [$1];", - "=r,l", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - ) - - -@dsl_user_op -def store_global_u32( - pointer: cute.Pointer, - value: Uint32, - *, - loc=None, - ip=None, -) -> None: - address = pointer.toint(loc=loc, ip=ip) - llvm.inline_asm( - None, - [address.ir_value(loc=loc, ip=ip), value.ir_value(loc=loc, ip=ip)], - "st.global.u32 [$0], $1;", - "l,r", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - - -@dsl_user_op -def packed_u32x4_to_bf16x8(packed, *, loc=None, ip=None): - values = llvm.bitcast( - ir.VectorType.get([VEC_BF16], BFloat16.mlir_type, loc=loc), - packed.ir_value(loc=loc, ip=ip), - loc=loc, - ip=ip, - ) - return cute.TensorSSA(values, VEC_BF16, BFloat16) - - -@dsl_user_op -def packed_u32_to_bf16x2(packed: Uint32, *, loc=None, ip=None): - values = llvm.bitcast( - ir.VectorType.get([2], BFloat16.mlir_type, loc=loc), - packed.ir_value(loc=loc, ip=ip), - loc=loc, - ip=ip, - ) - return cute.TensorSSA(values, 2, BFloat16) - - -@dsl_user_op -def packed_u32x2_to_bf16x4(packed, *, loc=None, ip=None): - values = llvm.bitcast( - ir.VectorType.get([QUAD_BF16], BFloat16.mlir_type, loc=loc), - packed.ir_value(loc=loc, ip=ip), - loc=loc, - ip=ip, - ) - return cute.TensorSSA(values, QUAD_BF16, BFloat16) - - -@dsl_user_op -def bf16x8_to_packed_u32x4(values, *, loc=None, ip=None): - packed = llvm.bitcast( - ir.VectorType.get([4], T.i32(), loc=loc), - values.ir_value(loc=loc, ip=ip), - loc=loc, - ip=ip, - ) - return cute.TensorSSA(packed, 4, Uint32) - - -@dsl_user_op -def bf16x2_to_packed_u32(values, *, loc=None, ip=None) -> Uint32: - return Uint32( - llvm.bitcast( - T.i32(), - values.ir_value(loc=loc, ip=ip), - loc=loc, - ip=ip, - ) - ) - - -@dsl_user_op -def bf16x4_to_packed_u32x2(values, *, loc=None, ip=None): - packed = llvm.bitcast( - ir.VectorType.get([2], T.i32(), loc=loc), - values.ir_value(loc=loc, ip=ip), - loc=loc, - ip=ip, - ) - return cute.TensorSSA(packed, 2, Uint32) - - -@cute.jit -def sanitize_negative_zero_u32x4(packed): - sanitized = cute.make_rmem_tensor(cute.make_layout((4,)), Uint32) - for index in cutlass.range_constexpr(4): - sanitized[index] = sanitize_negative_zero_u32(packed[index]) - return sanitized.load() - - -@cute.jit -def sanitize_negative_zero_u32(word: Uint32) -> Uint32: - low = Uint16(word & Uint32(0xFFFF)) - high = Uint16(word >> Uint32(16)) - if low == Uint16(NEGATIVE_ZERO_BF16_BITS): - word = word & Uint32(0xFFFF0000) - if high == Uint16(NEGATIVE_ZERO_BF16_BITS): - word = word & Uint32(0x0000FFFF) - return word - - -@cute.jit -def sanitize_negative_zero_u32x2(packed): - sanitized = cute.make_rmem_tensor(cute.make_layout((2,)), Uint32) - for index in cutlass.range_constexpr(2): - sanitized[index] = sanitize_negative_zero_u32(packed[index]) - return sanitized.load() - - -@cute.jit -def fragment_has_negative_zero(packed): - dirty = False - for index in cutlass.range_constexpr(4): - word = packed[index] - dirty = ( - dirty - | (Uint16(word & Uint32(0xFFFF)) == Uint16(NEGATIVE_ZERO_BF16_BITS)) - | (Uint16(word >> Uint32(16)) == Uint16(NEGATIVE_ZERO_BF16_BITS)) - ) - return dirty - - -@dsl_user_op -def map_shared_to_peer( - smem_pointer: cute.Pointer, - peer_rank: Int32, - *, - loc=None, - ip=None, -) -> Int32: - address = smem_pointer.toint(loc=loc, ip=ip).ir_value(loc=loc, ip=ip) - return Int32( - llvm.inline_asm( - T.i32(), - [address, peer_rank.ir_value(loc=loc, ip=ip)], - "mapa.shared::cluster.u32 $0, $1, $2;", - "=r,r,r", - has_side_effects=False, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - ) - - -@dsl_user_op -def store_shared_cluster_f32( - remote_address: Int32, - value: Float32, - *, - loc=None, - ip=None, -) -> None: - llvm.inline_asm( - None, - [ - remote_address.ir_value(loc=loc, ip=ip), - value.ir_value(loc=loc, ip=ip), - ], - "st.shared::cluster.f32 [$0], $1;", - "r,f", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - - -@dsl_user_op -def load_shared_u32x4(pointer: cute.Pointer, *, loc=None, ip=None): - address = pointer.toint(loc=loc, ip=ip) - # Prevent motion across the named-barrier pipeline protocol. - loaded = llvm.inline_asm( - llvm.StructType.get_literal([T.i32()] * 4), - [Int32(address).ir_value(loc=loc, ip=ip)], - "ld.shared.v4.u32 {$0, $1, $2, $3}, [$4];", - "=r,=r,=r,=r,r", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - packed = vector.from_elements( - ir.VectorType.get([4], T.i32(), loc=loc), - [ - llvm.extractvalue(T.i32(), loaded, [index], loc=loc, ip=ip) - for index in range(4) - ], - loc=loc, - ip=ip, - ) - return cute.TensorSSA(packed, 4, Uint32) - - -@dsl_user_op -def store_shared_u32x4( - pointer: cute.Pointer, - packed, - *, - loc=None, - ip=None, -) -> None: - address = pointer.toint(loc=loc, ip=ip) - words = [packed[index].ir_value(loc=loc, ip=ip) for index in range(4)] - llvm.inline_asm( - None, - [Int32(address).ir_value(loc=loc, ip=ip), *words], - "st.shared.v4.u32 [$0], {$1, $2, $3, $4};", - "r,r,r,r,r", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - - -@dsl_user_op -def load_global_u32x4_address( - address: Int64, - *, - volatile: cutlass.Constexpr[bool] = False, - loc=None, - ip=None, -): - opcode = "ld.volatile.global.v4.u32" if volatile else "ld.global.v4.u32" - loaded = llvm.inline_asm( - llvm.StructType.get_literal([T.i32()] * 4), - [address.ir_value(loc=loc, ip=ip)], - f"{opcode} {{$0, $1, $2, $3}}, [$4];", - "=r,=r,=r,=r,l", - has_side_effects=volatile, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - packed = vector.from_elements( - ir.VectorType.get([4], T.i32(), loc=loc), - [ - llvm.extractvalue(T.i32(), loaded, [index], loc=loc, ip=ip) - for index in range(4) - ], - loc=loc, - ip=ip, - ) - return cute.TensorSSA(packed, 4, Uint32) - - -@dsl_user_op -def packed_negative_zero_bf16x8(*, loc=None, ip=None): - word = Uint32(NEGATIVE_ZERO_BF16_PAIR).ir_value(loc=loc, ip=ip) - packed = vector.from_elements( - ir.VectorType.get([4], T.i32(), loc=loc), - [word, word, word, word], - loc=loc, - ip=ip, - ) - return cute.TensorSSA(packed, 4, Uint32) - - -@dsl_user_op -def cpasync_bulk_g2s( - gmem_ptr: cute.Pointer, - smem_ptr: cute.Pointer, - barrier_ptr: cute.Pointer, - size_bytes: Int32, - *, - loc=None, - ip=None, -) -> None: - operands = [ - gmem_ptr.toint(loc=loc, ip=ip).ir_value(), - smem_ptr.toint(loc=loc, ip=ip).ir_value(), - barrier_ptr.toint(loc=loc, ip=ip).ir_value(), - size_bytes.ir_value(loc=loc, ip=ip), - Int64(L2_EVICT_FIRST).ir_value(), - ] - llvm.inline_asm( - None, - operands, - ( - "cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes" - ".L2::cache_hint [$1], [$0], $3, [$2], $4;" - ), - "l,r,r,r,l", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - - -@dsl_user_op -def fence_proxy_async_shared_cta(*, loc=None, ip=None) -> None: - llvm.inline_asm( - None, - [], - "fence.proxy.async.shared::cta;", - "", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - - -@dsl_user_op -def remote_release_add1_u32(address: Int64, *, loc=None, ip=None) -> None: - llvm.inline_asm( - None, - [address.ir_value(loc=loc, ip=ip)], - "red.release.sys.global.add.u32 [$0], 1;", - "l", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - - -@dsl_user_op -def ldmc_bf16x8(address: Int64, *, loc=None, ip=None): - loaded = llvm.inline_asm( - llvm.StructType.get_literal([T.i32()] * 4), - [address.ir_value(loc=loc, ip=ip)], - "multimem.ld_reduce.relaxed.sys.global.add.acc::f32.v4.bf16x2 {$0, $1, $2, $3}, [$4];", - "=r,=r,=r,=r,l", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - packed = vector.from_elements( - ir.VectorType.get([4], T.i32(), loc=loc), - [ - llvm.extractvalue(T.i32(), loaded, [index], loc=loc, ip=ip) - for index in range(4) - ], - loc=loc, - ip=ip, - ) - return cute.TensorSSA(packed, 4, Uint32) - - -@dsl_user_op -def stmc_bf16x2( - address: Int64, - packed: Uint32, - *, - loc=None, - ip=None, -) -> None: - llvm.inline_asm( - None, - [address.ir_value(loc=loc, ip=ip), packed.ir_value(loc=loc, ip=ip)], - "multimem.st.relaxed.sys.global.bf16x2 [$0], $1;", - "l,r", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - - -@dsl_user_op -def stmc_bf16x4(address: Int64, values, *, loc=None, ip=None) -> None: - words = [values[index].ir_value(loc=loc, ip=ip) for index in range(2)] - llvm.inline_asm( - None, - [address.ir_value(loc=loc, ip=ip), *words], - "multimem.st.relaxed.sys.global.v2.bf16x2 [$0], {$1, $2};", - "l,r,r", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) - - -@dsl_user_op -def stmc_bf16x8(address: Int64, values, *, loc=None, ip=None) -> None: - words = [values[index].ir_value(loc=loc, ip=ip) for index in range(4)] - llvm.inline_asm( - None, - [address.ir_value(loc=loc, ip=ip), *words], - "multimem.st.relaxed.sys.global.v4.bf16x2 [$0], {$1, $2, $3, $4};", - "l,r,r,r,r", - has_side_effects=True, - is_align_stack=False, - asm_dialect=llvm.AsmDialect.AD_ATT, - loc=loc, - ip=ip, - ) diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_bt/__init__.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_bt/__init__.py deleted file mode 100644 index aaf171a7d..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_bt/__init__.py +++ /dev/null @@ -1,43 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Balanced MNNVL protocol.""" - -from .protocol import ( - BT_ALL_REDUCE_GB300_TP8_H8192_PRESET_0, - BT_ALL_REDUCE_GB300_TP8_H8192_PRESET_1, - BT_ALL_REDUCE_GB300_TP16_H8192_PRESET_0, - BT_ALL_REDUCE_GB300_TP16_H8192_PRESET_1, - BT_FINALIZE_GB300_TP8_H8192_K10_PRESET_0, - BT_FINALIZE_GB300_TP8_H8192_K10_PRESET_1, - BT_FINALIZE_GB300_TP16_H8192_K10_PRESET_0, - BT_FINALIZE_GB300_TP16_H8192_K10_PRESET_1, - BTAllReduceTuning, - BTCollectiveTuning, - BTFinalizeTuning, -) - -__all__ = [ - "BT_ALL_REDUCE_GB300_TP8_H8192_PRESET_0", - "BT_ALL_REDUCE_GB300_TP8_H8192_PRESET_1", - "BT_ALL_REDUCE_GB300_TP16_H8192_PRESET_0", - "BT_ALL_REDUCE_GB300_TP16_H8192_PRESET_1", - "BT_FINALIZE_GB300_TP8_H8192_K10_PRESET_0", - "BT_FINALIZE_GB300_TP8_H8192_K10_PRESET_1", - "BT_FINALIZE_GB300_TP16_H8192_K10_PRESET_0", - "BT_FINALIZE_GB300_TP16_H8192_K10_PRESET_1", - "BTAllReduceTuning", - "BTCollectiveTuning", - "BTFinalizeTuning", -] diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_bt/device_kernels.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_bt/device_kernels.py deleted file mode 100644 index 50af0ebb1..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_bt/device_kernels.py +++ /dev/null @@ -1,1285 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Three-stage BF16 MoE finalize, TP reduction, and RMSNorm for SM100.""" - -from __future__ import annotations - -import math - -import cuda.bindings.driver as cuda -import cutlass -import cutlass.cute as cute -from cutlass import BFloat16, Float32, Int32, Int64, Uint32 - -from ..cute_dsl_primitives import ( - NEGATIVE_ZERO_BF16_BITS, - VEC_BF16, - WARP_SIZE, - bf16x2_to_packed_u32, - bf16x4_to_packed_u32x2, - bf16x8_to_packed_u32x4, - f32_to_bf16_bits, - fragment_has_negative_zero, - load_global_bf16_as_f32, - load_global_bf16_as_f32_predicated, - load_global_u32, - load_global_u32_predicated, - load_global_u32x2, - load_global_u32x2_predicated, - load_global_u32x4, - load_volatile_u32, - packed_u32_to_bf16x2, - packed_u32x2_to_bf16x4, - packed_u32x4_to_bf16x8, - sanitize_negative_zero_u32, - sanitize_negative_zero_u32x2, - sanitize_negative_zero_u32x4, - stmc_bf16x8, - store_global_u16_bits, - store_global_u32, - store_global_u32_address, - store_global_u32x2, - store_global_u32x4, - store_lamport_sentinel_u32x4, -) - -LAMPORT_GENERATIONS = 3 -NEXT_STAGE = 0 -ACTIVE_STAGE = 1 - - -@cute.jit -def _block_sum( - value: Float32, - warp_sums: cute.Tensor, - warps: cutlass.Constexpr[int], -) -> Float32: - lane = cute.arch.lane_idx() - warp = cute.arch.warp_idx() - value = cute.arch.warp_reduction_sum(value) - if lane == 0: - cute.arch.store((warp_sums + warp).llvm_ptr, value) - cute.arch.barrier() - - result = Float32(0.0) - if warp == 0: - if lane < Int32(warps): - result = cute.arch.load((warp_sums + lane).llvm_ptr, Float32) - result = cute.arch.warp_reduction_sum(result) - if lane == 0: - cute.arch.store(warp_sums.llvm_ptr, result) - cute.arch.barrier() - return cute.arch.load(warp_sums.llvm_ptr, Float32) - - -class _ScalarFinalizeUnicastDeviceKernel: - def __init__( - self, - *, - hidden_size: int, - top_k: int, - tp_size: int, - rank: int, - local_capacity: int, - threads: int, - routed_scaling_factor: float, - include_shared_expert: bool, - load_shared_expert_before_pdl: bool, - enable_pdl: bool, - prefetch_group: int, - ) -> None: - self.hidden_size = hidden_size - self.top_k = top_k - self.tp_size = tp_size - self.rank = rank - self.local_capacity = local_capacity - self.threads = threads - self.routed_scaling_factor = routed_scaling_factor - self.include_shared_expert = include_shared_expert - self.load_shared_expert_before_pdl = load_shared_expert_before_pdl - self.enable_pdl = enable_pdl - self.prefetch_group = prefetch_group - self.prefetch_groups = (top_k + prefetch_group - 1) // prefetch_group - self.ctas_per_token = math.ceil(hidden_size / threads) - - @cute.jit - def __call__( - self, - routed_output: cute.Tensor, - expert_weights: cute.Tensor, - permuted_indices: cute.Tensor, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_peer_addresses: cute.Tensor, - m: Int32, - stream: cuda.CUstream, - ) -> None: - self.kernel( - routed_output, - expert_weights, - permuted_indices, - shared_output, - stage_state, - contribution_mailbox_peer_addresses, - ).launch( - grid=(m * self.ctas_per_token, 1, 1), - block=(self.threads, 1, 1), - smem=self.top_k * 8, - stream=stream, - use_pdl=self.enable_pdl, - ) - - @cute.kernel - def kernel( - self, - routed_output: cute.Tensor, - expert_weights: cute.Tensor, - permuted_indices: cute.Tensor, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_peer_addresses: cute.Tensor, - ) -> None: - block, _, _ = cute.arch.block_idx() - tidx, _, _ = cute.arch.thread_idx() - token = block // self.ctas_per_token - hidden_index = (block % self.ctas_per_token) * self.threads + tidx - - smem = cutlass.utils.SmemAllocator() - staged_indices = smem.allocate_array(Int32, self.top_k) - staged_weights = smem.allocate_array(Float32, self.top_k) - metadata_index = Int32(tidx) - while metadata_index < Int32(self.top_k): - metadata_offset = Int64(token) * self.top_k + Int64(metadata_index) - routed_index = cute.arch.load( - (permuted_indices.iterator + metadata_offset).llvm_ptr, - Int32, - ) - weight = load_global_bf16_as_f32( - Int64((expert_weights.iterator + metadata_offset).toint()) - ) - if cutlass.const_expr(self.routed_scaling_factor != 1.0): - weight = weight * Float32(self.routed_scaling_factor) - if routed_index == Int32(-1): - weight = Float32(0.0) - cute.arch.store( - (staged_indices + metadata_index).llvm_ptr, - routed_index, - ) - cute.arch.store( - (staged_weights + metadata_index).llvm_ptr, - weight, - ) - metadata_index = metadata_index + self.threads - cute.arch.barrier() - - shared_value = Float32(0.0) - if cutlass.const_expr( - self.include_shared_expert and self.load_shared_expert_before_pdl - ): - if hidden_index < Int32(self.hidden_size): - shared_offset = Int64(token) * self.hidden_size + Int64(hidden_index) - shared_value = load_global_bf16_as_f32( - Int64((shared_output.iterator + shared_offset).toint()) - ) - - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_wait() - - if cutlass.const_expr( - self.include_shared_expert and not self.load_shared_expert_before_pdl - ): - if hidden_index < Int32(self.hidden_size): - shared_offset = Int64(token) * self.hidden_size + Int64(hidden_index) - shared_value = load_global_bf16_as_f32( - Int64((shared_output.iterator + shared_offset).toint()) - ) - - stage = load_volatile_u32(stage_state.iterator + NEXT_STAGE) - if hidden_index < Int32(self.hidden_size): - accumulator = Float32(0.0) - if cutlass.const_expr(self.prefetch_group == 1): - for k in cutlass.range_constexpr(self.top_k): - routed_index = cute.arch.load( - (staged_indices + k).llvm_ptr, - Int32, - ) - weight = cute.arch.load( - (staged_weights + k).llvm_ptr, - Float32, - ) - source_offset = Int64(routed_index) * self.hidden_size + Int64( - hidden_index - ) - accumulator = ( - accumulator - + load_global_bf16_as_f32_predicated( - Int64((routed_output.iterator + source_offset).toint()), - Int32(routed_index != Int32(-1)), - ) - * weight - ) - else: - inputs = cute.make_rmem_tensor( - cute.make_layout((self.prefetch_group,)), - Float32, - ) - inputs.fill(Float32(0.0)) - for group in cutlass.range_constexpr(self.prefetch_groups): - for item in cutlass.range_constexpr(self.prefetch_group): - k = group * self.prefetch_group + item - if cutlass.const_expr(k < self.top_k): - routed_index = cute.arch.load( - (staged_indices + k).llvm_ptr, - Int32, - ) - source_offset = Int64( - routed_index - ) * self.hidden_size + Int64(hidden_index) - inputs[item] = load_global_bf16_as_f32_predicated( - Int64((routed_output.iterator + source_offset).toint()), - Int32(routed_index != Int32(-1)), - ) - for item in cutlass.range_constexpr(self.prefetch_group): - k = group * self.prefetch_group + item - if cutlass.const_expr(k < self.top_k): - accumulator = accumulator + inputs[item] * cute.arch.load( - (staged_weights + k).llvm_ptr, - Float32, - ) - - if cutlass.const_expr(self.include_shared_expert): - accumulator = accumulator + shared_value - bits = f32_to_bf16_bits(accumulator) - if bits == Uint32(NEGATIVE_ZERO_BF16_BITS): - bits = Uint32(0) - - destination_rank = token % self.tp_size - local_token = token // self.tp_size - destination_base = cute.arch.load( - ( - contribution_mailbox_peer_addresses.iterator + destination_rank - ).llvm_ptr, - Int64, - ) - destination_offset = ( - (Int64(stage) * self.tp_size + self.rank) * self.local_capacity - + Int64(local_token) - ) * self.hidden_size + Int64(hidden_index) - store_global_u16_bits( - destination_base + destination_offset * 2, - bits, - ) - - if block == 0 and tidx == 0: - store_global_u32(stage_state.iterator + ACTIVE_STAGE, stage) - cute.arch.barrier() - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_launch_dependents() - - -class _NarrowVectorFinalizeUnicastDeviceKernel: - def __init__( - self, - *, - hidden_size: int, - top_k: int, - tp_size: int, - rank: int, - local_capacity: int, - threads: int, - elements_per_thread: int, - routed_scaling_factor: float, - include_shared_expert: bool, - load_shared_expert_before_pdl: bool, - enable_pdl: bool, - prefetch_group: int, - ) -> None: - self.hidden_size = hidden_size - self.top_k = top_k - self.tp_size = tp_size - self.rank = rank - self.local_capacity = local_capacity - self.threads = threads - self.elements_per_thread = elements_per_thread - self.routed_scaling_factor = routed_scaling_factor - self.include_shared_expert = include_shared_expert - self.load_shared_expert_before_pdl = load_shared_expert_before_pdl - self.enable_pdl = enable_pdl - self.prefetch_group = prefetch_group - self.prefetch_groups = (top_k + prefetch_group - 1) // prefetch_group - self.words_per_fragment = elements_per_thread // 2 - self.fragments = hidden_size // elements_per_thread - self.ctas_per_token = math.ceil(self.fragments / threads) - - @cute.jit - def __call__( - self, - routed_output: cute.Tensor, - expert_weights: cute.Tensor, - permuted_indices: cute.Tensor, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_peer_addresses: cute.Tensor, - m: Int32, - stream: cuda.CUstream, - ) -> None: - self.kernel( - routed_output, - expert_weights, - permuted_indices, - shared_output, - stage_state, - contribution_mailbox_peer_addresses, - ).launch( - grid=(m * self.ctas_per_token, 1, 1), - block=(self.threads, 1, 1), - smem=self.top_k * 8, - stream=stream, - use_pdl=self.enable_pdl, - ) - - @cute.kernel - def kernel( - self, - routed_output: cute.Tensor, - expert_weights: cute.Tensor, - permuted_indices: cute.Tensor, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_peer_addresses: cute.Tensor, - ) -> None: - block, _, _ = cute.arch.block_idx() - tidx, _, _ = cute.arch.thread_idx() - token = block // self.ctas_per_token - fragment = (block % self.ctas_per_token) * self.threads + tidx - - smem = cutlass.utils.SmemAllocator() - staged_indices = smem.allocate_array(Int32, self.top_k) - staged_weights = smem.allocate_array(Float32, self.top_k) - metadata_index = Int32(tidx) - while metadata_index < Int32(self.top_k): - metadata_offset = Int64(token) * self.top_k + Int64(metadata_index) - routed_index = cute.arch.load( - (permuted_indices.iterator + metadata_offset).llvm_ptr, - Int32, - ) - weight = load_global_bf16_as_f32( - Int64((expert_weights.iterator + metadata_offset).toint()) - ) - if cutlass.const_expr(self.routed_scaling_factor != 1.0): - weight = weight * Float32(self.routed_scaling_factor) - if routed_index == Int32(-1): - weight = Float32(0.0) - cute.arch.store( - (staged_indices + metadata_index).llvm_ptr, - routed_index, - ) - cute.arch.store( - (staged_weights + metadata_index).llvm_ptr, - weight, - ) - metadata_index = metadata_index + self.threads - cute.arch.barrier() - - if cutlass.const_expr(self.include_shared_expert): - shared_values = cute.make_rmem_tensor( - cute.make_layout((self.elements_per_thread,)), - BFloat16, - ) - shared_values.fill(BFloat16(0.0)) - if cutlass.const_expr(self.load_shared_expert_before_pdl): - if fragment < self.fragments: - shared_offset = ( - Int64(token) * self.hidden_size - + Int64(fragment) * self.elements_per_thread - ) - shared_pointer = cute.make_ptr( - BFloat16, - (shared_output.iterator + shared_offset).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=self.elements_per_thread * 2, - ) - if cutlass.const_expr(self.elements_per_thread == 2): - shared_values.store( - packed_u32_to_bf16x2(load_global_u32(shared_pointer)) - ) - else: - shared_values.store( - packed_u32x2_to_bf16x4(load_global_u32x2(shared_pointer)) - ) - - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_wait() - - if cutlass.const_expr( - self.include_shared_expert and not self.load_shared_expert_before_pdl - ): - if fragment < self.fragments: - shared_offset = ( - Int64(token) * self.hidden_size - + Int64(fragment) * self.elements_per_thread - ) - shared_pointer = cute.make_ptr( - BFloat16, - (shared_output.iterator + shared_offset).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=self.elements_per_thread * 2, - ) - if cutlass.const_expr(self.elements_per_thread == 2): - shared_values.store( - packed_u32_to_bf16x2(load_global_u32(shared_pointer)) - ) - else: - shared_values.store( - packed_u32x2_to_bf16x4(load_global_u32x2(shared_pointer)) - ) - - stage = load_volatile_u32(stage_state.iterator + NEXT_STAGE) - if fragment < self.fragments: - accumulator = cute.make_rmem_tensor( - cute.make_layout((self.elements_per_thread,)), - Float32, - ) - accumulator.fill(Float32(0.0)) - if cutlass.const_expr(self.prefetch_group == 1): - for k in cutlass.range_constexpr(self.top_k): - routed_index = cute.arch.load( - (staged_indices + k).llvm_ptr, - Int32, - ) - weight = cute.arch.load( - (staged_weights + k).llvm_ptr, - Float32, - ) - source_offset = ( - Int64(routed_index) * self.hidden_size - + Int64(fragment) * self.elements_per_thread - ) - source_pointer = cute.make_ptr( - BFloat16, - (routed_output.iterator + source_offset).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=self.elements_per_thread * 2, - ) - if cutlass.const_expr(self.elements_per_thread == 2): - source = packed_u32_to_bf16x2( - load_global_u32_predicated( - source_pointer, - Int32(routed_index != Int32(-1)), - ) - ).to(Float32) - else: - source = packed_u32x2_to_bf16x4( - load_global_u32x2_predicated( - source_pointer, - Int32(routed_index != Int32(-1)), - ) - ).to(Float32) - accumulator.store(accumulator.load() + source * weight) - else: - inputs = cute.make_rmem_tensor( - cute.make_layout((self.prefetch_group, self.words_per_fragment)), - Uint32, - ) - inputs.fill(Uint32(0)) - for group in cutlass.range_constexpr(self.prefetch_groups): - for item in cutlass.range_constexpr(self.prefetch_group): - k = group * self.prefetch_group + item - if cutlass.const_expr(k < self.top_k): - routed_index = cute.arch.load( - (staged_indices + k).llvm_ptr, - Int32, - ) - source_offset = ( - Int64(routed_index) * self.hidden_size - + Int64(fragment) * self.elements_per_thread - ) - source_pointer = cute.make_ptr( - BFloat16, - (routed_output.iterator + source_offset).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=self.elements_per_thread * 2, - ) - if cutlass.const_expr(self.elements_per_thread == 2): - inputs[item, 0] = load_global_u32_predicated( - source_pointer, - Int32(routed_index != Int32(-1)), - ) - else: - source = load_global_u32x2_predicated( - source_pointer, - Int32(routed_index != Int32(-1)), - ) - for word in cutlass.range_constexpr( - self.words_per_fragment - ): - inputs[item, word] = source[word] - for item in cutlass.range_constexpr(self.prefetch_group): - k = group * self.prefetch_group + item - if cutlass.const_expr(k < self.top_k): - if cutlass.const_expr(self.elements_per_thread == 2): - source_values = packed_u32_to_bf16x2( - inputs[item, 0] - ).to(Float32) - else: - source = cute.make_rmem_tensor( - cute.make_layout((self.words_per_fragment,)), - Uint32, - ) - for word in cutlass.range_constexpr( - self.words_per_fragment - ): - source[word] = inputs[item, word] - source_values = packed_u32x2_to_bf16x4( - source.load() - ).to(Float32) - accumulator.store( - accumulator.load() - + source_values - * cute.arch.load( - (staged_weights + k).llvm_ptr, - Float32, - ) - ) - - result = accumulator.load() - if cutlass.const_expr(self.include_shared_expert): - result = result + shared_values.load().to(Float32) - result = result.to(BFloat16) - destination_rank = token % self.tp_size - local_token = token // self.tp_size - destination_base = cute.arch.load( - ( - contribution_mailbox_peer_addresses.iterator + destination_rank - ).llvm_ptr, - Int64, - ) - destination_offset = ( - (Int64(stage) * self.tp_size + self.rank) * self.local_capacity - + Int64(local_token) - ) * self.hidden_size + Int64(fragment) * self.elements_per_thread - if cutlass.const_expr(self.elements_per_thread == 2): - word = sanitize_negative_zero_u32(bf16x2_to_packed_u32(result)) - store_global_u32_address( - destination_base + destination_offset * 2, - word, - ) - else: - half = sanitize_negative_zero_u32x2(bf16x4_to_packed_u32x2(result)) - store_global_u32x2( - destination_base + destination_offset * 2, - half, - ) - - if block == 0 and tidx == 0: - store_global_u32(stage_state.iterator + ACTIVE_STAGE, stage) - cute.arch.barrier() - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_launch_dependents() - - -class _VectorFinalizeUnicastDeviceKernel: - def __init__( - self, - *, - hidden_size: int, - top_k: int, - tp_size: int, - rank: int, - local_capacity: int, - threads: int, - routed_scaling_factor: float, - include_shared_expert: bool, - load_shared_expert_before_pdl: bool, - enable_pdl: bool, - prefetch_group: int, - ) -> None: - self.hidden_size = hidden_size - self.top_k = top_k - self.tp_size = tp_size - self.rank = rank - self.local_capacity = local_capacity - self.threads = threads - self.routed_scaling_factor = routed_scaling_factor - self.include_shared_expert = include_shared_expert - self.load_shared_expert_before_pdl = load_shared_expert_before_pdl - self.enable_pdl = enable_pdl - self.prefetch_group = prefetch_group - self.prefetch_groups = (top_k + prefetch_group - 1) // prefetch_group - self.words_per_fragment = VEC_BF16 // 2 - self.fragments = hidden_size // VEC_BF16 - self.ctas_per_token = math.ceil(self.fragments / threads) - - @cute.jit - def __call__( - self, - routed_output: cute.Tensor, - expert_weights: cute.Tensor, - permuted_indices: cute.Tensor, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_peer_addresses: cute.Tensor, - m: Int32, - stream: cuda.CUstream, - ) -> None: - self.kernel( - routed_output, - expert_weights, - permuted_indices, - shared_output, - stage_state, - contribution_mailbox_peer_addresses, - ).launch( - grid=(m * self.ctas_per_token, 1, 1), - block=(self.threads, 1, 1), - smem=self.top_k * 8, - stream=stream, - use_pdl=self.enable_pdl, - ) - - @cute.kernel - def kernel( - self, - routed_output: cute.Tensor, - expert_weights: cute.Tensor, - permuted_indices: cute.Tensor, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_peer_addresses: cute.Tensor, - ) -> None: - block, _, _ = cute.arch.block_idx() - tidx, _, _ = cute.arch.thread_idx() - token = block // self.ctas_per_token - fragment = (block % self.ctas_per_token) * self.threads + tidx - - smem = cutlass.utils.SmemAllocator() - staged_indices = smem.allocate_array(Int32, self.top_k) - staged_weights = smem.allocate_array(Float32, self.top_k) - metadata_index = Int32(tidx) - while metadata_index < Int32(self.top_k): - metadata_offset = Int64(token) * self.top_k + Int64(metadata_index) - routed_index = cute.arch.load( - (permuted_indices.iterator + metadata_offset).llvm_ptr, - Int32, - ) - weight = load_global_bf16_as_f32( - Int64((expert_weights.iterator + metadata_offset).toint()) - ) - if cutlass.const_expr(self.routed_scaling_factor != 1.0): - weight = weight * Float32(self.routed_scaling_factor) - if routed_index == Int32(-1): - # Vector loads are unpredicated; row zero is safe and its weight is zero. - routed_index = Int32(0) - weight = Float32(0.0) - cute.arch.store( - (staged_indices + metadata_index).llvm_ptr, - routed_index, - ) - cute.arch.store( - (staged_weights + metadata_index).llvm_ptr, - weight, - ) - metadata_index = metadata_index + self.threads - cute.arch.barrier() - - if cutlass.const_expr(self.include_shared_expert): - shared_values = cute.make_rmem_tensor( - cute.make_layout((VEC_BF16,)), - BFloat16, - ) - shared_values.fill(BFloat16(0.0)) - if cutlass.const_expr(self.load_shared_expert_before_pdl): - if fragment < self.fragments: - shared_offset = ( - Int64(token) * self.hidden_size + Int64(fragment) * VEC_BF16 - ) - shared_pointer = cute.make_ptr( - BFloat16, - (shared_output.iterator + shared_offset).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=16, - ) - shared_values.store( - packed_u32x4_to_bf16x8(load_global_u32x4(shared_pointer)) - ) - - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_wait() - - if cutlass.const_expr( - self.include_shared_expert and not self.load_shared_expert_before_pdl - ): - if fragment < self.fragments: - shared_offset = ( - Int64(token) * self.hidden_size + Int64(fragment) * VEC_BF16 - ) - shared_pointer = cute.make_ptr( - BFloat16, - (shared_output.iterator + shared_offset).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=16, - ) - shared_values.store( - packed_u32x4_to_bf16x8(load_global_u32x4(shared_pointer)) - ) - - stage = load_volatile_u32(stage_state.iterator + NEXT_STAGE) - if fragment < self.fragments: - accumulator = cute.make_rmem_tensor( - cute.make_layout((VEC_BF16,)), - Float32, - ) - accumulator.fill(Float32(0.0)) - if cutlass.const_expr(self.prefetch_group == 1): - for k in cutlass.range_constexpr(self.top_k): - routed_index = cute.arch.load( - (staged_indices + k).llvm_ptr, - Int32, - ) - weight = cute.arch.load( - (staged_weights + k).llvm_ptr, - Float32, - ) - source_offset = ( - Int64(routed_index) * self.hidden_size - + Int64(fragment) * VEC_BF16 - ) - source_pointer = cute.make_ptr( - BFloat16, - (routed_output.iterator + source_offset).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=16, - ) - source_values = packed_u32x4_to_bf16x8( - load_global_u32x4(source_pointer) - ).to(Float32) - if weight != Float32(0.0): - accumulator.store(accumulator.load() + source_values * weight) - else: - inputs = cute.make_rmem_tensor( - cute.make_layout((self.prefetch_group, self.words_per_fragment)), - Uint32, - ) - inputs.fill(Uint32(0)) - for group in cutlass.range_constexpr(self.prefetch_groups): - for item in cutlass.range_constexpr(self.prefetch_group): - k = group * self.prefetch_group + item - if cutlass.const_expr(k < self.top_k): - routed_index = cute.arch.load( - (staged_indices + k).llvm_ptr, - Int32, - ) - source_offset = ( - Int64(routed_index) * self.hidden_size - + Int64(fragment) * VEC_BF16 - ) - source_pointer = cute.make_ptr( - BFloat16, - (routed_output.iterator + source_offset).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=16, - ) - source = load_global_u32x4(source_pointer) - for word in cutlass.range_constexpr( - self.words_per_fragment - ): - inputs[item, word] = source[word] - for item in cutlass.range_constexpr(self.prefetch_group): - k = group * self.prefetch_group + item - if cutlass.const_expr(k < self.top_k): - source = cute.make_rmem_tensor( - cute.make_layout((self.words_per_fragment,)), - Uint32, - ) - for word in cutlass.range_constexpr( - self.words_per_fragment - ): - source[word] = inputs[item, word] - weight = cute.arch.load( - (staged_weights + k).llvm_ptr, - Float32, - ) - if weight != Float32(0.0): - accumulator.store( - accumulator.load() - + packed_u32x4_to_bf16x8(source.load()).to(Float32) - * weight - ) - - result = accumulator.load() - if cutlass.const_expr(self.include_shared_expert): - result = result + shared_values.load().to(Float32) - result = result.to(BFloat16) - packed = sanitize_negative_zero_u32x4(bf16x8_to_packed_u32x4(result)) - destination_rank = token % self.tp_size - local_token = token // self.tp_size - destination_base = cute.arch.load( - ( - contribution_mailbox_peer_addresses.iterator + destination_rank - ).llvm_ptr, - Int64, - ) - destination_offset = ( - (Int64(stage) * self.tp_size + self.rank) * self.local_capacity - + Int64(local_token) - ) * self.hidden_size + Int64(fragment) * VEC_BF16 - store_global_u32x4( - destination_base + destination_offset * 2, - packed, - ) - - if block == 0 and tidx == 0: - store_global_u32(stage_state.iterator + ACTIVE_STAGE, stage) - cute.arch.barrier() - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_launch_dependents() - - -class _SharedOnlyPublishDeviceKernel: - def __init__( - self, - *, - hidden_size: int, - tp_size: int, - rank: int, - local_capacity: int, - threads: int, - vectors_per_thread: int, - enable_pdl: bool, - ) -> None: - self.hidden_size = hidden_size - self.tp_size = tp_size - self.rank = rank - self.local_capacity = local_capacity - self.threads = threads - self.vectors_per_thread = vectors_per_thread - self.enable_pdl = enable_pdl - self.fragments = hidden_size // VEC_BF16 - self.fragments_per_cta = threads * vectors_per_thread - self.ctas_per_token = math.ceil(self.fragments / self.fragments_per_cta) - - @cute.jit - def __call__( - self, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_peer_addresses: cute.Tensor, - m: Int32, - stream: cuda.CUstream, - ) -> None: - self.kernel( - shared_output, - stage_state, - contribution_mailbox_peer_addresses, - ).launch( - grid=(m * self.ctas_per_token, 1, 1), - block=(self.threads, 1, 1), - stream=stream, - use_pdl=self.enable_pdl, - ) - - @cute.kernel - def kernel( - self, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_peer_addresses: cute.Tensor, - ) -> None: - block, _, _ = cute.arch.block_idx() - tidx, _, _ = cute.arch.thread_idx() - token = block // self.ctas_per_token - fragment_base = (block % self.ctas_per_token) * self.fragments_per_cta + tidx - - inputs = cute.make_rmem_tensor( - cute.make_layout( - (self.vectors_per_thread, 4), - stride=(4, 1), - ), - Uint32, - ) - inputs.fill(Uint32(0)) - - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_wait() - - for trip in cutlass.range_constexpr(self.vectors_per_thread): - fragment = fragment_base + trip * self.threads - if fragment < self.fragments: - source_offset = ( - Int64(token) * self.hidden_size + Int64(fragment) * VEC_BF16 - ) - source_pointer = cute.make_ptr( - BFloat16, - (shared_output.iterator + source_offset).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=16, - ) - packed = load_global_u32x4(source_pointer) - for word in cutlass.range_constexpr(4): - inputs[trip, word] = packed[word] - - stage = load_volatile_u32(stage_state.iterator + NEXT_STAGE) - destination_rank = token % self.tp_size - local_token = token // self.tp_size - destination_base = cute.arch.load( - (contribution_mailbox_peer_addresses.iterator + destination_rank).llvm_ptr, - Int64, - ) - for trip in cutlass.range_constexpr(self.vectors_per_thread): - fragment = fragment_base + trip * self.threads - if fragment < self.fragments: - packed = cute.make_rmem_tensor(cute.make_layout((4,)), Uint32) - for word in cutlass.range_constexpr(4): - packed[word] = inputs[trip, word] - destination_offset = ( - (Int64(stage) * self.tp_size + self.rank) * self.local_capacity - + Int64(local_token) - ) * self.hidden_size + Int64(fragment) * VEC_BF16 - store_global_u32x4( - destination_base + destination_offset * 2, - sanitize_negative_zero_u32x4(packed.load()), - ) - - if block == 0 and tidx == 0: - store_global_u32(stage_state.iterator + ACTIVE_STAGE, stage) - cute.arch.barrier() - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_launch_dependents() - - -class _OwnerReduceMulticastDeviceKernel: - def __init__( - self, - *, - hidden_size: int, - tp_size: int, - rank: int, - capacity_m: int, - local_capacity: int, - threads: int, - add_residual: bool, - enable_pdl: bool, - ) -> None: - self.hidden_size = hidden_size - self.tp_size = tp_size - self.rank = rank - self.capacity_m = capacity_m - self.local_capacity = local_capacity - self.threads = threads - self.add_residual = add_residual - self.enable_pdl = enable_pdl - self.fragments = hidden_size // VEC_BF16 - self.ctas_per_token = math.ceil(self.fragments / threads) - - @cute.jit - def __call__( - self, - contribution_mailbox: cute.Tensor, - residual_source: cute.Tensor, - stage_state: cute.Tensor, - prenorm_mailbox_multicast_address: Int64, - m: Int32, - stream: cuda.CUstream, - ) -> None: - local_tokens = (m + Int32(self.tp_size - 1)) // Int32(self.tp_size) - self.kernel( - contribution_mailbox, - residual_source, - stage_state, - prenorm_mailbox_multicast_address, - m, - ).launch( - grid=(local_tokens * self.ctas_per_token, 1, 1), - block=(self.threads, 1, 1), - stream=stream, - use_pdl=self.enable_pdl, - ) - - @cute.kernel - def kernel( - self, - contribution_mailbox: cute.Tensor, - residual_source: cute.Tensor, - stage_state: cute.Tensor, - prenorm_mailbox_multicast_address: Int64, - m: Int32, - ) -> None: - block, _, _ = cute.arch.block_idx() - tidx, _, _ = cute.arch.thread_idx() - local_token = block // self.ctas_per_token - fragment = (block % self.ctas_per_token) * self.threads + tidx - token = local_token * self.tp_size + self.rank - active = token < m and fragment < self.fragments - - if cutlass.const_expr(self.add_residual): - residual = cute.make_rmem_tensor( - cute.make_layout((VEC_BF16,)), - BFloat16, - ) - residual.fill(BFloat16(0.0)) - if active: - residual_offset = ( - Int64(token) * self.hidden_size + Int64(fragment) * VEC_BF16 - ) - residual_pointer = cute.make_ptr( - BFloat16, - (residual_source.iterator + residual_offset).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=16, - ) - residual.store( - packed_u32x4_to_bf16x8(load_global_u32x4(residual_pointer)) - ) - - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_wait() - - stage = load_volatile_u32(stage_state.iterator + ACTIVE_STAGE) - rank_values = cute.make_rmem_tensor( - cute.make_layout((self.tp_size, 4)), - Uint32, - ) - rank_values.fill(Uint32(0)) - dirty = active - while dirty: - dirty = False - for source_rank in cutlass.range_constexpr(self.tp_size): - source_offset = ( - (Int64(stage) * self.tp_size + source_rank) * self.local_capacity - + Int64(local_token) - ) * self.hidden_size + Int64(fragment) * VEC_BF16 - source_pointer = cute.make_ptr( - BFloat16, - (contribution_mailbox.iterator + source_offset).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=16, - ) - packed = load_global_u32x4(source_pointer, volatile=True) - dirty = dirty | fragment_has_negative_zero(packed) - for word in cutlass.range_constexpr(4): - rank_values[source_rank, word] = packed[word] - if active: - reduced = cute.make_rmem_tensor( - cute.make_layout((VEC_BF16,)), - Float32, - ) - reduced.fill(Float32(0.0)) - for source_rank in cutlass.range_constexpr(self.tp_size): - packed = cute.make_rmem_tensor(cute.make_layout((4,)), Uint32) - for word in cutlass.range_constexpr(4): - packed[word] = rank_values[source_rank, word] - reduced.store( - reduced.load() + packed_u32x4_to_bf16x8(packed.load()).to(Float32) - ) - - prenorm = reduced.load() - if cutlass.const_expr(self.add_residual): - prenorm = prenorm + residual.load().to(Float32) - prenorm_bf16 = prenorm.to(BFloat16) - output_offset = ( - Int64(stage) * self.capacity_m + Int64(token) - ) * self.hidden_size + Int64(fragment) * VEC_BF16 - stmc_bf16x8( - prenorm_mailbox_multicast_address + output_offset * 2, - sanitize_negative_zero_u32x4(bf16x8_to_packed_u32x4(prenorm_bf16)), - ) - - if block == 0 and tidx == 0: - store_global_u32( - stage_state.iterator + NEXT_STAGE, - (stage + Uint32(1)) % Uint32(LAMPORT_GENERATIONS), - ) - cute.arch.barrier() - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_launch_dependents() - - if active: - for source_rank in cutlass.range_constexpr(self.tp_size): - source_offset = ( - (Int64(stage) * self.tp_size + source_rank) * self.local_capacity - + Int64(local_token) - ) * self.hidden_size + Int64(fragment) * VEC_BF16 - store_lamport_sentinel_u32x4( - Int64((contribution_mailbox.iterator + source_offset).toint()) - ) - - -class _MaterializeRMSNormDeviceKernel: - def __init__( - self, - *, - hidden_size: int, - capacity_m: int, - threads: int, - rms_epsilon: float, - weight_bias: float, - write_residual_output: bool, - enable_pdl: bool, - ) -> None: - self.hidden_size = hidden_size - self.capacity_m = capacity_m - self.threads = threads - self.rms_epsilon = rms_epsilon - self.weight_bias = weight_bias - self.write_residual_output = write_residual_output - self.enable_pdl = enable_pdl - self.fragments = hidden_size // VEC_BF16 - self.trips = math.ceil(self.fragments / threads) - self.warps = threads // WARP_SIZE - - @cute.jit - def __call__( - self, - prenorm_mailbox: cute.Tensor, - residual_output: cute.Tensor, - norm_output: cute.Tensor, - gamma: cute.Tensor, - stage_state: cute.Tensor, - m: Int32, - stream: cuda.CUstream, - ) -> None: - self.kernel( - prenorm_mailbox, - residual_output, - norm_output, - gamma, - stage_state, - ).launch( - grid=(m, 1, 1), - block=(self.threads, 1, 1), - smem=self.warps * 4, - stream=stream, - use_pdl=self.enable_pdl, - ) - - @cute.kernel - def kernel( - self, - prenorm_mailbox: cute.Tensor, - residual_output: cute.Tensor, - norm_output: cute.Tensor, - gamma: cute.Tensor, - stage_state: cute.Tensor, - ) -> None: - tidx, _, _ = cute.arch.thread_idx() - token, _, _ = cute.arch.block_idx() - - gamma_fragments = cute.make_rmem_tensor( - cute.make_layout( - (self.trips, VEC_BF16), - stride=(VEC_BF16, 1), - ), - BFloat16, - ) - prenorm_fragments = cute.make_rmem_tensor( - cute.make_layout( - (self.trips, VEC_BF16), - stride=(VEC_BF16, 1), - ), - BFloat16, - ) - gamma_fragments.fill(BFloat16(0.0)) - prenorm_fragments.fill(BFloat16(0.0)) - for trip in cutlass.range_constexpr(self.trips): - fragment = tidx + trip * self.threads - if fragment < self.fragments: - gamma_pointer = cute.make_ptr( - BFloat16, - (gamma.iterator + Int64(fragment) * VEC_BF16).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=16, - ) - gamma_fragments[trip, None].store( - packed_u32x4_to_bf16x8(load_global_u32x4(gamma_pointer)) - ) - - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_wait() - - stage = load_volatile_u32(stage_state.iterator + ACTIVE_STAGE) - thread_sum = Float32(0.0) - for trip in cutlass.range_constexpr(self.trips): - fragment = tidx + trip * self.threads - if fragment < self.fragments: - mailbox_offset = ( - Int64(stage) * self.capacity_m + Int64(token) - ) * self.hidden_size + Int64(fragment) * VEC_BF16 - mailbox_pointer = cute.make_ptr( - BFloat16, - (prenorm_mailbox.iterator + mailbox_offset).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=16, - ) - packed = load_global_u32x4(mailbox_pointer, volatile=True) - while fragment_has_negative_zero(packed): - packed = load_global_u32x4(mailbox_pointer, volatile=True) - prenorm = packed_u32x4_to_bf16x8(packed) - prenorm_fragments[trip, None].store(prenorm) - - if cutlass.const_expr(self.write_residual_output): - output_offset = ( - Int64(token) * self.hidden_size + Int64(fragment) * VEC_BF16 - ) - store_global_u32x4( - Int64((residual_output.iterator + output_offset).toint()), - packed, - ) - - values = prenorm.to(Float32) - thread_sum = thread_sum + (values * values).reduce( - cute.ReductionOp.ADD, - init_val=Float32(0.0), - reduction_profile=0, - ) - - smem = cutlass.utils.SmemAllocator() - warp_sums = smem.allocate_array(Float32, self.warps) - full_sum = _block_sum(thread_sum, warp_sums, self.warps) - inverse_rms = cute.math.rsqrt( - full_sum / Float32(self.hidden_size) + Float32(self.rms_epsilon), - fastmath=True, - ) - for trip in cutlass.range_constexpr(self.trips): - fragment = tidx + trip * self.threads - if fragment < self.fragments: - gamma_value = gamma_fragments[trip, None].load().to(Float32) - if cutlass.const_expr(self.weight_bias != 0.0): - gamma_value = gamma_value + Float32(self.weight_bias) - result = ( - prenorm_fragments[trip, None].load().to(Float32) - * inverse_rms - * gamma_value - ).to(BFloat16) - output_offset = ( - Int64(token) * self.hidden_size + Int64(fragment) * VEC_BF16 - ) - store_global_u32x4( - Int64((norm_output.iterator + output_offset).toint()), - bf16x8_to_packed_u32x4(result), - ) - - cute.arch.barrier() - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_launch_dependents() - - for trip in cutlass.range_constexpr(self.trips): - fragment = tidx + trip * self.threads - if fragment < self.fragments: - mailbox_offset = ( - Int64(stage) * self.capacity_m + Int64(token) - ) * self.hidden_size + Int64(fragment) * VEC_BF16 - store_lamport_sentinel_u32x4( - Int64((prenorm_mailbox.iterator + mailbox_offset).toint()) - ) diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_bt/protocol.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_bt/protocol.py deleted file mode 100644 index 373afd44f..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_bt/protocol.py +++ /dev/null @@ -1,511 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Balanced MNNVL protocol and its two operation paths.""" - -from __future__ import annotations - -import math -from dataclasses import dataclass -from typing import Any, TypedDict, cast - -import cutlass.cute as cute -import torch -import torch.distributed as dist -from cutlass import BFloat16, Int32, Int64 -from cutlass.cute.runtime import make_fake_compact_tensor - -from ..cute_dsl_primitives import VEC_BF16 -from ..runtime import ( - current_cu_stream, - make_fake_dynamic_compact_tensor, - to_cute, - to_cute_dynamic, -) -from ..symmetric_buffer import SymmetricBuffer -from .device_kernels import ( - LAMPORT_GENERATIONS, - _MaterializeRMSNormDeviceKernel, - _NarrowVectorFinalizeUnicastDeviceKernel, - _OwnerReduceMulticastDeviceKernel, - _ScalarFinalizeUnicastDeviceKernel, - _SharedOnlyPublishDeviceKernel, - _VectorFinalizeUnicastDeviceKernel, -) - - -@dataclass(frozen=True, slots=True) -class BTCollectiveTuning: - reduction_threads: int = 128 - rms_threads: int = 1024 - enable_pdl: bool = True - - -@dataclass(frozen=True, slots=True) -class BTFinalizeTuning: - elements_per_thread: int = VEC_BF16 - threads: int = 128 - prefetch_group: int = 1 - load_shared_expert_before_pdl: bool = False - collective: BTCollectiveTuning = BTCollectiveTuning() - - -@dataclass(frozen=True, slots=True) -class BTAllReduceTuning: - publish_threads: int = 128 - publish_vectors_per_thread: int = 1 - collective: BTCollectiveTuning = BTCollectiveTuning(reduction_threads=32) - - -BT_FINALIZE_GB300_TP8_H8192_K10_PRESET_0 = BTFinalizeTuning( - elements_per_thread=2, threads=256 -) -BT_FINALIZE_GB300_TP8_H8192_K10_PRESET_1 = BTFinalizeTuning() -BT_FINALIZE_GB300_TP16_H8192_K10_PRESET_0 = BTFinalizeTuning( - elements_per_thread=2, threads=256 -) -BT_FINALIZE_GB300_TP16_H8192_K10_PRESET_1 = BTFinalizeTuning() -BT_ALL_REDUCE_GB300_TP8_H8192_PRESET_0 = BTAllReduceTuning() -BT_ALL_REDUCE_GB300_TP8_H8192_PRESET_1 = BTAllReduceTuning( - collective=BTCollectiveTuning(reduction_threads=320) -) -BT_ALL_REDUCE_GB300_TP16_H8192_PRESET_0 = BTAllReduceTuning() -BT_ALL_REDUCE_GB300_TP16_H8192_PRESET_1 = BTAllReduceTuning( - collective=BTCollectiveTuning(reduction_threads=320) -) - - -@dataclass(slots=True) -class BTProtocolState: - contribution_mailbox: SymmetricBuffer - prenorm_mailbox: SymmetricBuffer - stage_state: torch.Tensor - - -@dataclass(frozen=True, slots=True) -class _CompiledTail: - reduce: Any - rms_norm: Any - - -@dataclass(frozen=True, slots=True) -class _CompiledFinalize: - publish: Any - tail: _CompiledTail - - -@dataclass(frozen=True, slots=True) -class _CompiledAllReduce: - publish: Any - tail: _CompiledTail - - -class _PathKwargs(TypedDict): - hidden_size: int - top_k: int - capacity_m: int - write_residual_output: bool - - -class _BTPath: - def __init__( - self, - *, - hidden_size: int, - top_k: int, - capacity_m: int, - write_residual_output: bool, - ) -> None: - self.hidden_size = hidden_size - self.top_k = top_k - self.capacity_m = capacity_m - self.write_residual_output = write_residual_output - - def _outputs( - self, - m: int, - norm_output: torch.Tensor | None, - residual_output: torch.Tensor | None, - ) -> tuple[torch.Tensor, torch.Tensor | None]: - shape = (m, self.hidden_size) - device = torch.device("cuda", torch.cuda.current_device()) - if norm_output is None: - norm_output = torch.empty(shape, dtype=torch.bfloat16, device=device) - if self.write_residual_output and residual_output is None: - residual_output = torch.empty(shape, dtype=torch.bfloat16, device=device) - return norm_output, residual_output - - def _validate_state(self, state: BTProtocolState, m: int) -> None: - if not 1 <= m <= self.capacity_m: - raise ValueError(f"m must be in [1, {self.capacity_m}]") - if state.contribution_mailbox.peer_addresses is None: - raise ValueError("BT contribution mailbox requires peer addresses") - address = state.prenorm_mailbox.multicast_address - if address is None or address % 16: - raise ValueError( - "BT prenorm mailbox requires a 16-byte-aligned multicast address" - ) - - def _launch_tail( - self, - tail: _CompiledTail, - residual_source: torch.Tensor | None, - gamma: torch.Tensor, - state: BTProtocolState, - norm_output: torch.Tensor, - residual_output: torch.Tensor | None, - m: int, - ) -> None: - residual_arg = residual_source if residual_source is not None else norm_output - residual_output_arg = ( - residual_output if residual_output is not None else norm_output - ) - stream = current_cu_stream() - tail.reduce( - to_cute(state.contribution_mailbox.tensor.flatten(), 16), - to_cute_dynamic(residual_arg.flatten(), 16, divisibility=self.hidden_size), - to_cute(state.stage_state, 4), - Int64(cast(int, state.prenorm_mailbox.multicast_address)), - Int32(m), - stream, - ) - tail.rms_norm( - to_cute(state.prenorm_mailbox.tensor.flatten(), 16), - to_cute_dynamic( - residual_output_arg.flatten(), - 16, - divisibility=self.hidden_size, - ), - to_cute_dynamic(norm_output.flatten(), 16, divisibility=self.hidden_size), - to_cute(gamma, 16), - to_cute(state.stage_state, 4), - Int32(m), - stream, - ) - - -class FinalizeAllReduceRMSNormBTKernel(_BTPath): - def __init__(self, *, compiled: _CompiledFinalize, **kwargs) -> None: - super().__init__(**kwargs) - self._compiled = compiled - - def __call__( - self, - routed_output: torch.Tensor, - expert_weights: torch.Tensor, - permuted_indices: torch.Tensor, - shared_output: torch.Tensor | None, - residual_source: torch.Tensor | None, - gamma: torch.Tensor, - m: int, - *, - state: BTProtocolState, - norm_output: torch.Tensor | None = None, - residual_output: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor | None]: - self._validate_state(state, m) - norm_output, residual_output = self._outputs(m, norm_output, residual_output) - shared_arg = shared_output if shared_output is not None else norm_output - peers = cast(torch.Tensor, state.contribution_mailbox.peer_addresses) - self._compiled.publish( - to_cute_dynamic(routed_output.flatten(), 16, divisibility=self.hidden_size), - to_cute_dynamic(expert_weights.flatten(), 2, divisibility=self.top_k), - to_cute_dynamic(permuted_indices.flatten(), 4, divisibility=self.top_k), - to_cute_dynamic(shared_arg.flatten(), 16, divisibility=self.hidden_size), - to_cute(state.stage_state, 4), - to_cute(peers, 8), - Int32(m), - current_cu_stream(), - ) - self._launch_tail( - self._compiled.tail, - residual_source, - gamma, - state, - norm_output, - residual_output, - m, - ) - return norm_output, residual_output - - -class AllReduceRMSNormBTKernel(_BTPath): - def __init__(self, *, compiled: _CompiledAllReduce, **kwargs) -> None: - super().__init__(**kwargs) - self._compiled = compiled - - def __call__( - self, - local_contribution: torch.Tensor, - residual_source: torch.Tensor | None, - gamma: torch.Tensor, - m: int, - *, - state: BTProtocolState, - norm_output: torch.Tensor | None = None, - residual_output: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor | None]: - self._validate_state(state, m) - norm_output, residual_output = self._outputs(m, norm_output, residual_output) - peers = cast(torch.Tensor, state.contribution_mailbox.peer_addresses) - self._compiled.publish( - to_cute_dynamic( - local_contribution.flatten(), - 16, - divisibility=self.hidden_size, - ), - to_cute(state.stage_state, 4), - to_cute(peers, 8), - Int32(m), - current_cu_stream(), - ) - self._launch_tail( - self._compiled.tail, - residual_source, - gamma, - state, - norm_output, - residual_output, - m, - ) - return norm_output, residual_output - - -class BTProtocol: - """Own BT State and protocol-local compiled variants for both paths.""" - - def __init__( - self, - hidden_size: int, - top_k: int, - tp_size: int, - rank: int, - capacity_m: int, - rms_epsilon: float, - routed_scaling_factor: float, - weight_bias: float, - *, - include_shared_expert: bool, - add_residual: bool, - write_residual_output: bool, - finalize_tunings: tuple[BTFinalizeTuning, ...], - all_reduce_tunings: tuple[BTAllReduceTuning, ...], - group: dist.ProcessGroup, - ) -> None: - self.hidden_size = hidden_size - self.top_k = top_k - self.tp_size = tp_size - self.rank = rank - self.capacity_m = capacity_m - self.local_capacity = math.ceil(capacity_m / tp_size) - self.rms_epsilon = rms_epsilon - self.routed_scaling_factor = routed_scaling_factor - self.weight_bias = weight_bias - self.include_shared_expert = include_shared_expert - self.add_residual = add_residual - self.write_residual_output = write_residual_output - - tail_cache = { - tuning: self._compile_tail(tuning) - for tuning in { - *(item.collective for item in finalize_tunings), - *(item.collective for item in all_reduce_tunings), - } - } - self.finalize_kernels = { - tuning: FinalizeAllReduceRMSNormBTKernel( - compiled=_CompiledFinalize( - publish=self._compile_finalize(tuning), - tail=tail_cache[tuning.collective], - ), - **self._path_kwargs(), - ) - for tuning in dict.fromkeys(finalize_tunings) - } - self.all_reduce_kernels = { - tuning: AllReduceRMSNormBTKernel( - compiled=_CompiledAllReduce( - publish=self._compile_all_reduce_publish(tuning), - tail=tail_cache[tuning.collective], - ), - **self._path_kwargs(), - ) - for tuning in dict.fromkeys(all_reduce_tunings) - } - self.state = self._create_state(group) - - def _path_kwargs(self) -> _PathKwargs: - return { - "hidden_size": self.hidden_size, - "top_k": self.top_k, - "capacity_m": self.capacity_m, - "write_residual_output": self.write_residual_output, - } - - def _compile_finalize(self, tuning: BTFinalizeTuning): - if tuning.elements_per_thread not in (1, 2, 4, VEC_BF16): - raise ValueError("BT finalize elements_per_thread must be 1, 2, 4, or 8") - kwargs: dict[str, Any] = { - "hidden_size": self.hidden_size, - "top_k": self.top_k, - "tp_size": self.tp_size, - "rank": self.rank, - "local_capacity": self.local_capacity, - "threads": tuning.threads, - "routed_scaling_factor": self.routed_scaling_factor, - "include_shared_expert": self.include_shared_expert, - "load_shared_expert_before_pdl": tuning.load_shared_expert_before_pdl, - "enable_pdl": tuning.collective.enable_pdl, - "prefetch_group": tuning.prefetch_group, - } - device_kernel: Any - if tuning.elements_per_thread == 1: - device_kernel = _ScalarFinalizeUnicastDeviceKernel(**kwargs) - elif tuning.elements_per_thread == VEC_BF16: - device_kernel = _VectorFinalizeUnicastDeviceKernel(**kwargs) - else: - device_kernel = _NarrowVectorFinalizeUnicastDeviceKernel( - **kwargs, elements_per_thread=tuning.elements_per_thread - ) - return cute.compile( - device_kernel, - *self._publish_compile_args(include_routed=True), - ) - - def _compile_all_reduce_publish(self, tuning: BTAllReduceTuning): - device_kernel = _SharedOnlyPublishDeviceKernel( - hidden_size=self.hidden_size, - tp_size=self.tp_size, - rank=self.rank, - local_capacity=self.local_capacity, - threads=tuning.publish_threads, - vectors_per_thread=tuning.publish_vectors_per_thread, - enable_pdl=tuning.collective.enable_pdl, - ) - return cute.compile( - device_kernel, - *self._publish_compile_args(include_routed=False), - ) - - def _publish_compile_args(self, *, include_routed: bool) -> tuple: - activation = make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ) - common = ( - make_fake_compact_tensor(Int32, (2,), assumed_align=4), - make_fake_compact_tensor(Int64, (self.tp_size,), assumed_align=8), - Int32(self.capacity_m), - current_cu_stream(), - ) - if not include_routed: - return (activation, *common) - return ( - activation, - make_fake_dynamic_compact_tensor( - BFloat16, alignment=2, divisibility=self.top_k - ), - make_fake_dynamic_compact_tensor( - Int32, alignment=4, divisibility=self.top_k - ), - activation, - *common, - ) - - def _compile_tail(self, tuning: BTCollectiveTuning) -> _CompiledTail: - reduce_kernel = _OwnerReduceMulticastDeviceKernel( - hidden_size=self.hidden_size, - tp_size=self.tp_size, - rank=self.rank, - capacity_m=self.capacity_m, - local_capacity=self.local_capacity, - threads=tuning.reduction_threads, - add_residual=self.add_residual, - enable_pdl=tuning.enable_pdl, - ) - reduce_elements = ( - LAMPORT_GENERATIONS * self.tp_size * self.local_capacity * self.hidden_size - ) - reduce = cute.compile( - reduce_kernel, - make_fake_compact_tensor(BFloat16, (reduce_elements,), assumed_align=16), - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_compact_tensor(Int32, (2,), assumed_align=4), - Int64(0), - Int32(self.capacity_m), - current_cu_stream(), - ) - rms_kernel = _MaterializeRMSNormDeviceKernel( - hidden_size=self.hidden_size, - capacity_m=self.capacity_m, - threads=tuning.rms_threads, - rms_epsilon=self.rms_epsilon, - weight_bias=self.weight_bias, - write_residual_output=self.write_residual_output, - enable_pdl=tuning.enable_pdl, - ) - prenorm_elements = LAMPORT_GENERATIONS * self.capacity_m * self.hidden_size - rms_norm = cute.compile( - rms_kernel, - make_fake_compact_tensor(BFloat16, (prenorm_elements,), assumed_align=16), - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_compact_tensor(BFloat16, (self.hidden_size,), assumed_align=16), - make_fake_compact_tensor(Int32, (2,), assumed_align=4), - Int32(self.capacity_m), - current_cu_stream(), - ) - return _CompiledTail(reduce=reduce, rms_norm=rms_norm) - - def _create_state(self, group: dist.ProcessGroup) -> BTProtocolState: - if dist.get_world_size(group) != self.tp_size: - raise ValueError("ProcessGroup size does not match tp_size") - if dist.get_rank(group) != self.rank: - raise ValueError("ProcessGroup rank does not match rank") - device = torch.device("cuda", torch.cuda.current_device()) - contribution = SymmetricBuffer.allocate( - ( - LAMPORT_GENERATIONS, - self.tp_size, - self.local_capacity, - self.hidden_size, - ), - torch.bfloat16, - device, - group, - materialize_peer_addresses=True, - ) - contribution.tensor.view(torch.int16).fill_(-32768) - prenorm = SymmetricBuffer.allocate( - ( - LAMPORT_GENERATIONS, - self.capacity_m, - self.hidden_size, - ), - torch.bfloat16, - device, - group, - require_multicast=True, - ) - prenorm.tensor.view(torch.int16).fill_(-32768) - return BTProtocolState( - contribution_mailbox=contribution, - prenorm_mailbox=prenorm, - stage_state=torch.zeros((2,), dtype=torch.int32, device=device), - ) diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ht/__init__.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ht/__init__.py deleted file mode 100644 index 136ec2c47..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ht/__init__.py +++ /dev/null @@ -1,37 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""High-throughput MNNVL protocol.""" - -from .protocol import ( - HT_ALL_REDUCE_GB300_TP8_H8192, - HT_ALL_REDUCE_GB300_TP16_H8192, - HT_FINALIZE_GB300_TP8_H8192_K10, - HT_FINALIZE_GB300_TP8_H8192_K10_M_GE_2049, - HT_FINALIZE_GB300_TP8_H8192_K10_M_LE_2048, - HT_FINALIZE_GB300_TP16_H8192_K10, - HTAllReduceTuning, - HTFinalizeTuning, -) - -__all__ = [ - "HT_ALL_REDUCE_GB300_TP8_H8192", - "HT_ALL_REDUCE_GB300_TP16_H8192", - "HT_FINALIZE_GB300_TP8_H8192_K10", - "HT_FINALIZE_GB300_TP8_H8192_K10_M_GE_2049", - "HT_FINALIZE_GB300_TP8_H8192_K10_M_LE_2048", - "HT_FINALIZE_GB300_TP16_H8192_K10", - "HTAllReduceTuning", - "HTFinalizeTuning", -] diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ht/device_kernel.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ht/device_kernel.py deleted file mode 100644 index 5c8421590..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ht/device_kernel.py +++ /dev/null @@ -1,978 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Persistent BF16 MoE finalize, TP reduction, and RMSNorm for SM100.""" - -from __future__ import annotations - -import cuda.bindings.driver as cuda -import cutlass -import cutlass.cute as cute -import cutlass.pipeline as pipeline -import cutlass.utils as utils -from cutlass import BFloat16, Float32, Int32, Int64, Uint32 - -from ..cute_dsl_primitives import ( - VEC_BF16, - WARP_SIZE, - bf16x8_to_packed_u32x4, - cpasync_bulk_g2s, - fence_proxy_async_shared_cta, - fragment_has_negative_zero, - ldmc_bf16x8, - load_global_bf16_as_f32, - load_global_u32x4_address, - load_shared_u32x4, - packed_negative_zero_bf16x8, - packed_u32x4_to_bf16x8, - remote_release_add1_u32, - sanitize_negative_zero_u32x4, - stmc_bf16x8, - store_global_u32x4, - store_shared_u32x4, -) - -SMEM_ALIGNMENT = 1024 - - -class _MoeFinalizeAllReduceRMSNormHTDeviceKernel: - def __init__( - self, - *, - hidden: int, - top_k: int, - tp: int, - rank: int, - active_ctas: int, - stages: int, - consumer_threads: int, - vectors_per_thread: int, - reduction_warps: int, - reduction_cta_groups: int | None, - rms_token_groups: int, - rms_pipeline_stages: int, - rms_shard_major: bool, - rms_epsilon: float, - routed_scaling_factor: float, - weight_bias: float, - include_shared_expert: bool, - add_residual: bool, - write_residual_output: bool, - enable_pdl: bool, - ) -> None: - if tp not in (2, 4, 8, 16): - raise ValueError("tp must be 2, 4, 8, or 16") - if rank < 0 or rank >= tp: - raise ValueError("rank must be in [0, tp)") - if hidden <= 0 or hidden % VEC_BF16: - raise ValueError("hidden must be a positive multiple of 8") - if top_k < 0: - raise ValueError("top_k must be nonnegative") - if active_ctas <= 0 or active_ctas % tp: - raise ValueError("active_ctas must be positive and divisible by tp") - if stages < 2: - raise ValueError("stages must be at least 2") - if consumer_threads <= 0 or consumer_threads % WARP_SIZE: - raise ValueError("consumer_threads must be a positive warp multiple") - if vectors_per_thread <= 0: - raise ValueError("vectors_per_thread must be positive") - if reduction_warps not in (1, 2, 4, 8): - raise ValueError("reduction_warps must be 1, 2, 4, or 8") - if rms_token_groups not in (1, 2, 4): - raise ValueError("rms_token_groups must be 1, 2, or 4") - if consumer_threads % rms_token_groups: - raise ValueError("consumer threads must divide across RMS token groups") - if rms_pipeline_stages not in (1, 2, 3): - raise ValueError("rms_pipeline_stages must be 1, 2, or 3") - block_threads = consumer_threads + (2 + reduction_warps) * WARP_SIZE - if block_threads > 1024: - raise ValueError("warp roles exceed the CUDA block limit") - shard_elements = consumer_threads * VEC_BF16 * vectors_per_thread - if hidden <= 0 or hidden % shard_elements: - raise ValueError(f"hidden must be divisible by {shard_elements}") - cta_groups = active_ctas // tp - if reduction_cta_groups is None: - reduction_cta_groups = active_ctas // tp - if reduction_cta_groups <= 0 or reduction_cta_groups * tp > active_ctas: - raise ValueError("reduction CTA groups and shards must fit the grid") - contributions = top_k + int(include_shared_expert) - if contributions <= 0: - raise ValueError("at least one local contribution is required") - self.hidden = hidden - self.top_k = top_k - self.tp = tp - self.rank = rank - self.active_ctas = active_ctas - self.stages = stages - self.vectors_per_thread = vectors_per_thread - self.consumer_threads = consumer_threads - self.reduction_warps = reduction_warps - self.reduction_cta_groups = reduction_cta_groups - self.reduction_ctas = reduction_cta_groups * tp - self.rms_token_groups = rms_token_groups - self.rms_pipeline_stages = rms_pipeline_stages - self.rms_shard_major = rms_shard_major - self.rms_epsilon = rms_epsilon - self.routed_scaling_factor = routed_scaling_factor - self.weight_bias = weight_bias - self.include_shared_expert = include_shared_expert - self.add_residual = add_residual - self.write_residual_output = write_residual_output - self.enable_pdl = enable_pdl - self.metadata_chunks = (top_k + WARP_SIZE - 1) // WARP_SIZE - self.metadata_slots = max(top_k, 1) - self.consumer_warps = consumer_threads // WARP_SIZE - self.rms_threads_per_token = consumer_threads // rms_token_groups - self.rms_warps_per_token = self.rms_threads_per_token // WARP_SIZE - self.rms_stage_slots = rms_token_groups * rms_pipeline_stages - self.rms_warp_sum_slots = self.rms_stage_slots * self.rms_warps_per_token - self.publisher_warp = 1 + self.consumer_warps - self.reduction_warp_begin = self.publisher_warp + 1 - self.reduction_threads = reduction_warps * WARP_SIZE - self.block_threads = block_threads - self.shard_elements = shard_elements - self.shard_bytes = shard_elements * 2 - self.hidden_shards = hidden // shard_elements - self.contributions = contributions - if ( - rms_pipeline_stages > 1 - and self.rms_stage_slots * hidden > self.shard_elements * stages - ): - raise ValueError("finalize stage storage cannot hold the RMS pipeline") - self.cta_groups = cta_groups - self.packs_per_token = hidden // VEC_BF16 - if self.packs_per_token % tp: - raise ValueError("hidden vector count must be divisible by tp") - if self.packs_per_token % consumer_threads: - raise ValueError("token vectors must divide evenly across consumers") - self.clear_vectors_per_thread = self.packs_per_token // consumer_threads - self.copy_threads = self.rms_threads_per_token - if self.packs_per_token % self.copy_threads: - raise ValueError("token vectors must divide evenly across finalize threads") - self.rms_vectors_per_thread = self.packs_per_token // self.copy_threads - self.packs_per_reduction_shard = self.packs_per_token // tp - if rms_shard_major: - if tp < self.rms_warps_per_token or tp % self.rms_warps_per_token: - raise ValueError( - "shard-major RMS requires an integer number of reduction " - "shards per RMS warp" - ) - self.reduction_shards_per_rms_warp = tp // self.rms_warps_per_token - if ( - self.rms_vectors_per_thread * WARP_SIZE - != self.packs_per_reduction_shard * self.reduction_shards_per_rms_warp - ): - raise ValueError( - "shard-major RMS warp coverage must match its reduction shards" - ) - else: - self.reduction_shards_per_rms_warp = 0 - if self.packs_per_reduction_shard % self.reduction_threads: - raise ValueError("the reduction shard must divide evenly across threads") - self.reduction_vectors_per_thread = ( - self.packs_per_reduction_shard // self.reduction_threads - ) - - @cute.jit - def _rms_arrive_and_wait(self, rms_group: Int32) -> None: - barrier_0 = pipeline.NamedBarrier( - barrier_id=2, num_threads=self.rms_threads_per_token - ) - if cutlass.const_expr(self.rms_token_groups > 1): - barrier_1 = pipeline.NamedBarrier( - barrier_id=3, num_threads=self.rms_threads_per_token - ) - if cutlass.const_expr(self.rms_token_groups == 4): - barrier_2 = pipeline.NamedBarrier( - barrier_id=4, num_threads=self.rms_threads_per_token - ) - barrier_3 = pipeline.NamedBarrier( - barrier_id=5, num_threads=self.rms_threads_per_token - ) - if rms_group == 0: - barrier_0.arrive_and_wait() - elif cutlass.const_expr(self.rms_token_groups == 2): # noqa: SIM114 - barrier_1.arrive_and_wait() - elif rms_group == 1: - barrier_1.arrive_and_wait() - elif rms_group == 2: - barrier_2.arrive_and_wait() - else: - barrier_3.arrive_and_wait() - else: - barrier_0.arrive_and_wait() - - @cute.jit - def __call__( - self, - routed_output: cute.Tensor, - expert_weights: cute.Tensor, - permuted_indices: cute.Tensor, - shared_output: cute.Tensor, - residual_source: cute.Tensor, - gamma: cute.Tensor, - local_contributions: cute.Tensor, - prenorm_mailbox: cute.Tensor, - residual_output: cute.Tensor, - norm_output: cute.Tensor, - ready_counter_peer_addresses: cute.Tensor, - ready_counters: cute.Tensor, - processed_counters: cute.Tensor, - local_contributions_multicast_address: Int64, - prenorm_mailbox_multicast_address: Int64, - m: Int32, - stream: cuda.CUstream, - ) -> None: - smem_layout = cute.make_layout( - (self.shard_elements * self.stages,), stride=(1,) - ) - - @cute.struct - class SharedStorage: - barriers: cute.struct.MemRange[Int64, 2 * self.stages] - stage_probs: cute.struct.MemRange[Float32, self.stages] - cached_rows: cute.struct.MemRange[Int32, self.metadata_slots] - cached_probs: cute.struct.MemRange[Float32, self.metadata_slots] - consumer_progress: cute.struct.MemRange[Int32, self.consumer_warps] - norm_warp_sums: cute.struct.MemRange[Float32, self.rms_warp_sum_slots] - norm_inv_rms: cute.struct.MemRange[Float32, self.rms_stage_slots] - rows: cute.struct.Align[ - cute.struct.MemRange[BFloat16, cute.cosize(smem_layout)], - SMEM_ALIGNMENT, - ] - - self.shared_storage: type[cute.struct.Struct] = SharedStorage - self.kernel( - routed_output, - shared_output, - residual_source, - gamma, - expert_weights, - permuted_indices, - local_contributions, - prenorm_mailbox, - residual_output, - norm_output, - ready_counter_peer_addresses, - ready_counters, - processed_counters, - local_contributions_multicast_address, - prenorm_mailbox_multicast_address, - m, - smem_layout, - ).launch( - grid=(self.active_ctas, 1, 1), - block=(self.block_threads, 1, 1), - min_blocks_per_mp=1, - use_pdl=self.enable_pdl, - stream=stream, - ) - - @cute.kernel - def kernel( - self, - routed_source: cute.Tensor, - shared_source: cute.Tensor, - residual_source: cute.Tensor, - gamma: cute.Tensor, - expert_weights: cute.Tensor, - permuted_indices: cute.Tensor, - local_contributions: cute.Tensor, - prenorm_mailbox: cute.Tensor, - residual_output: cute.Tensor, - norm_output: cute.Tensor, - ready_counter_peer_addresses: cute.Tensor, - ready_counters: cute.Tensor, - processed_counters: cute.Tensor, - local_contributions_multicast_address: Int64, - prenorm_mailbox_multicast_address: Int64, - m: Int32, - smem_layout: cute.Layout, - ) -> None: - block = cute.arch.block_idx()[0] - tidx = cute.arch.thread_idx()[0] - warp = cute.arch.make_warp_uniform(cute.arch.warp_idx()) - lane = cute.arch.lane_idx() - cta_group = block // self.tp - cta_slot = block % self.tp - wave = Int64(cta_group) - token = wave * self.tp + cta_slot - smem = utils.SmemAllocator() - storage = smem.allocate(self.shared_storage) - rows = storage.rows.get_tensor(smem_layout) - barrier_storage = storage.barriers.data_ptr() - stage_probs = storage.stage_probs.data_ptr() - cached_rows = storage.cached_rows.data_ptr() - cached_probs = storage.cached_probs.data_ptr() - consumer_progress = storage.consumer_progress.data_ptr() - norm_warp_sums = storage.norm_warp_sums.data_ptr() - norm_inv_rms = storage.norm_inv_rms.data_ptr() - if tidx < self.consumer_warps: - cute.arch.store((consumer_progress + tidx).llvm_ptr, Int32(0)) - cute.arch.sync_threads() - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_wait() - load_pipeline = pipeline.PipelineTmaAsync.create( - barrier_storage=barrier_storage, - num_stages=self.stages, - producer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, 1), - consumer_group=pipeline.CooperativeGroup( - pipeline.Agent.Thread, self.consumer_warps - ), - tx_count=self.shard_bytes, - ) - if warp == 0: - producer_state = pipeline.make_pipeline_state( - pipeline.PipelineUserType.Producer, self.stages - ) - peek_empty = cutlass.Boolean(1) - if token < Int64(m): - peek_empty = load_pipeline.producer_try_acquire(producer_state) - while token < Int64(m): - for metadata_chunk in cutlass.range_constexpr(self.metadata_chunks): - metadata_slot = metadata_chunk * WARP_SIZE + lane - if metadata_slot < self.top_k: - item = Int64(token) * self.top_k + metadata_slot - row = cute.arch.load( - (permuted_indices.iterator + item).llvm_ptr, Int32 - ) - prob = load_global_bf16_as_f32( - Int64((expert_weights.iterator + item).toint()) - ) - if cutlass.const_expr(self.routed_scaling_factor != 1.0): - prob = prob * Float32(self.routed_scaling_factor) - if row == Int32(-1): - row = Int32(0) - prob = Float32(0.0) - cute.arch.store((cached_rows + metadata_slot).llvm_ptr, row) - cute.arch.store((cached_probs + metadata_slot).llvm_ptr, prob) - cute.arch.sync_warp() - for shard in cutlass.range_constexpr(self.hidden_shards): - for contribution in cutlass.range_constexpr(self.contributions): - load_pipeline.producer_acquire(producer_state, peek_empty) - if lane == 0: - if cutlass.const_expr(contribution < self.top_k): - prob = cute.arch.load( - (cached_probs + contribution).llvm_ptr, - Float32, - ) - row = cute.arch.load( - (cached_rows + contribution).llvm_ptr, - Int32, - ) - source_element = ( - Int64(row) * self.hidden - + shard * self.shard_elements - ) - source = routed_source.iterator + source_element - else: - prob = Float32(1.0) - source_element = ( - Int64(token) * self.hidden - + shard * self.shard_elements - ) - source = shared_source.iterator + source_element - cute.arch.store( - (stage_probs + producer_state.index).llvm_ptr, - prob, - ) - fence_proxy_async_shared_cta() - cpasync_bulk_g2s( - source, - rows.iterator - + producer_state.index * self.shard_elements, - load_pipeline.producer_get_barrier(producer_state), - Int32(self.shard_bytes), - ) - producer_state.advance() - peek_empty = load_pipeline.producer_try_acquire(producer_state) - wave += self.cta_groups - token = wave * self.tp + cta_slot - load_pipeline.producer_tail(producer_state) - elif warp > 0 and warp <= self.consumer_warps: - finalize_join = pipeline.NamedBarrier( - barrier_id=6, num_threads=self.consumer_threads - ) - consumer_state = pipeline.make_pipeline_state( - pipeline.PipelineUserType.Consumer, self.stages - ) - consumer_tid = tidx - WARP_SIZE - consumer_wave = Int64(cta_group) - token = consumer_wave * self.tp + cta_slot - consumer_token_progress = Int32(0) - while token < Int64(m): - for shard in cutlass.range_constexpr(self.hidden_shards): - if cutlass.const_expr(self.top_k == 0): - load_pipeline.consumer_wait(consumer_state) - for trip in cutlass.range_constexpr(self.vectors_per_thread): - output_element = ( - Int64(token) * self.hidden - + shard * self.shard_elements - + trip * self.consumer_threads * VEC_BF16 - + consumer_tid * VEC_BF16 - ) - store_global_u32x4( - Int64( - ( - local_contributions.iterator + output_element - ).toint() - ), - load_shared_u32x4( - rows.iterator - + consumer_state.index * self.shard_elements - + trip * self.consumer_threads * VEC_BF16 - + consumer_tid * VEC_BF16 - ), - ) - load_pipeline.consumer_release(consumer_state) - consumer_state.advance() - accum = cute.make_rmem_tensor( - cute.make_layout( - (self.vectors_per_thread, VEC_BF16), - stride=(VEC_BF16, 1), - ), - Float32, - ) - accum.fill(Float32(0.0)) - for _ in cutlass.range_constexpr( - self.contributions if self.top_k > 0 else 0 - ): - load_pipeline.consumer_wait(consumer_state) - prob = cute.arch.load( - (stage_probs + consumer_state.index).llvm_ptr, - Float32, - ) - if prob != Float32(0.0): - for trip in cutlass.range_constexpr( - self.vectors_per_thread - ): - stage_ptr = ( - rows.iterator - + consumer_state.index * self.shard_elements - + trip * self.consumer_threads * VEC_BF16 - + consumer_tid * VEC_BF16 - ) - values = packed_u32x4_to_bf16x8( - load_shared_u32x4(stage_ptr) - ).to(Float32) - accum[trip, None].store( - accum[trip, None].load() + values * prob - ) - load_pipeline.consumer_release(consumer_state) - consumer_state.advance() - for trip in cutlass.range_constexpr( - self.vectors_per_thread if self.top_k > 0 else 0 - ): - output_element = ( - Int64(token) * self.hidden - + shard * self.shard_elements - + trip * self.consumer_threads * VEC_BF16 - + consumer_tid * VEC_BF16 - ) - store_global_u32x4( - Int64( - (local_contributions.iterator + output_element).toint() - ), - bf16x8_to_packed_u32x4( - accum[trip, None].load().to(BFloat16) - ), - ) - clear_value = packed_negative_zero_bf16x8() - token_pack = token * self.packs_per_token - for clear_item in cutlass.range_constexpr( - self.clear_vectors_per_thread - ): - clear_pack = consumer_tid + clear_item * self.consumer_threads - clear_element = (token_pack + clear_pack) * VEC_BF16 - store_global_u32x4( - Int64((prenorm_mailbox.iterator + clear_element).toint()), - clear_value, - ) - cute.arch.sync_warp() - consumer_token_progress += 1 - if lane == 0: - cute.arch.store( - (consumer_progress + warp - 1).llvm_ptr, - consumer_token_progress, - sem="release", - scope="cta", - ) - consumer_wave += self.cta_groups - token = consumer_wave * self.tp + cta_slot - finalize_join.arrive_and_wait() - if cutlass.const_expr(self.rms_token_groups > 1): - rms_group = (warp - 1) // self.rms_warps_per_token - rms_group_warp = warp - 1 - rms_group * self.rms_warps_per_token - copy_tid = rms_group_warp * WARP_SIZE + lane - else: - rms_group = Int32(0) - rms_group_warp = warp - 1 - copy_tid = consumer_tid - rms_pack_base = copy_tid - rms_pack_stride = self.copy_threads - if cutlass.const_expr(self.rms_shard_major): - rms_pack_base = ( - rms_group_warp - * self.reduction_shards_per_rms_warp - * self.packs_per_reduction_shard - + lane - ) - rms_pack_stride = WARP_SIZE - copy_wave = Int64(cta_group) + Int64(rms_group) * self.cta_groups - copy_token = copy_wave * self.tp + cta_slot - if cutlass.const_expr(self.rms_pipeline_stages > 1): - rms_wave_stride = self.cta_groups * self.rms_token_groups - while copy_token < Int64(m): - for rms_stage in cutlass.range_constexpr(self.rms_pipeline_stages): - stage_wave = copy_wave + rms_stage * rms_wave_stride - stage_token = stage_wave * self.tp + cta_slot - if stage_token < Int64(m): - stage_slot = ( - rms_group * self.rms_pipeline_stages + rms_stage - ) - token_pack = stage_token * self.packs_per_token - copy_fragments = cute.make_rmem_tensor( - cute.make_layout( - (self.rms_vectors_per_thread, 4), - stride=(4, 1), - ), - Uint32, - ) - all_ready = cutlass.Boolean(0) - while not all_ready: - all_ready = cutlass.Boolean(1) - for item in cutlass.range_constexpr( - self.rms_vectors_per_thread - ): - pack = rms_pack_base + item * rms_pack_stride - linear_pack = token_pack + pack - packed = load_global_u32x4_address( - Int64( - ( - prenorm_mailbox.iterator - + linear_pack * VEC_BF16 - ).toint() - ), - volatile=True, - ) - copy_fragments[item, None].store(packed) - all_ready = all_ready and ( - not fragment_has_negative_zero(packed) - ) - thread_sum = Float32(0.0) - if cutlass.const_expr(self.write_residual_output): - prenorm_packed = [] - for item in cutlass.range_constexpr( - self.rms_vectors_per_thread - ): - pack = rms_pack_base + item * rms_pack_stride - prenorm = packed_u32x4_to_bf16x8( - copy_fragments[item, None].load() - ) - packed_prenorm = bf16x8_to_packed_u32x4(prenorm) - if cutlass.const_expr(self.write_residual_output): - prenorm_packed.append(packed_prenorm) - # Finalize has drained `rows`; __init__ verifies the RMS layout fits. - store_shared_u32x4( - rows.iterator - + stage_slot * self.hidden - + pack * VEC_BF16, - packed_prenorm, - ) - prenorm_f32 = prenorm.to(Float32) - thread_sum = thread_sum + ( - prenorm_f32 * prenorm_f32 - ).reduce( - cute.ReductionOp.ADD, - init_val=Float32(0.0), - reduction_profile=0, - ) - if cutlass.const_expr(self.write_residual_output): - for item in cutlass.range_constexpr( - self.rms_vectors_per_thread - ): - pack = rms_pack_base + item * rms_pack_stride - linear_pack = token_pack + pack - residual_address = Int64( - ( - residual_output.iterator - + linear_pack * VEC_BF16 - ).toint() - ) - store_global_u32x4( - residual_address, prenorm_packed[item] - ) - warp_sum = cute.arch.warp_reduction_sum(thread_sum) - if lane == 0: - cute.arch.store( - ( - norm_warp_sums - + stage_slot * self.rms_warps_per_token - + rms_group_warp - ).llvm_ptr, - warp_sum, - ) - self._rms_arrive_and_wait(rms_group) - if warp == 1 + rms_group * self.rms_warps_per_token: - for rms_stage in cutlass.range_constexpr( - self.rms_pipeline_stages - ): - stage_wave = copy_wave + rms_stage * rms_wave_stride - stage_token = stage_wave * self.tp + cta_slot - if stage_token < Int64(m): - stage_slot = ( - rms_group * self.rms_pipeline_stages + rms_stage - ) - cta_sum = Float32(0.0) - if lane < self.rms_warps_per_token: - cta_sum = cute.arch.load( - ( - norm_warp_sums - + stage_slot * self.rms_warps_per_token - + lane - ).llvm_ptr, - Float32, - ) - cta_sum = cute.arch.warp_reduction_sum(cta_sum) - if lane == 0: - inv_rms = cute.math.rsqrt( - cta_sum / Float32(self.hidden) - + Float32(self.rms_epsilon), - fastmath=True, - ) - cute.arch.store( - (norm_inv_rms + stage_slot).llvm_ptr, - inv_rms, - ) - self._rms_arrive_and_wait(rms_group) - gamma_values = [] - for item in cutlass.range_constexpr(self.rms_vectors_per_thread): - pack = rms_pack_base + item * rms_pack_stride - gamma_value = packed_u32x4_to_bf16x8( - load_global_u32x4_address( - Int64((gamma.iterator + pack * VEC_BF16).toint()) - ) - ).to(Float32) - if cutlass.const_expr(self.weight_bias != 0.0): - gamma_value = gamma_value + Float32(self.weight_bias) - gamma_values.append(gamma_value) - for rms_stage in cutlass.range_constexpr(self.rms_pipeline_stages): - stage_wave = copy_wave + rms_stage * rms_wave_stride - stage_token = stage_wave * self.tp + cta_slot - if stage_token < Int64(m): - stage_slot = ( - rms_group * self.rms_pipeline_stages + rms_stage - ) - token_pack = stage_token * self.packs_per_token - inv_rms = cute.arch.load( - (norm_inv_rms + stage_slot).llvm_ptr, Float32 - ) - norm_packed = [] - for item in cutlass.range_constexpr( - self.rms_vectors_per_thread - ): - pack = rms_pack_base + item * rms_pack_stride - prenorm = packed_u32x4_to_bf16x8( - load_shared_u32x4( - rows.iterator - + stage_slot * self.hidden - + pack * VEC_BF16 - ) - ).to(Float32) - result = (prenorm * inv_rms * gamma_values[item]).to( - BFloat16 - ) - norm_packed.append(bf16x8_to_packed_u32x4(result)) - for item in cutlass.range_constexpr( - self.rms_vectors_per_thread - ): - pack = rms_pack_base + item * rms_pack_stride - linear_pack = token_pack + pack - store_global_u32x4( - Int64( - ( - norm_output.iterator - + linear_pack * VEC_BF16 - ).toint() - ), - norm_packed[item], - ) - copy_wave += rms_wave_stride * self.rms_pipeline_stages - copy_token = copy_wave * self.tp + cta_slot - if cutlass.const_expr(self.rms_pipeline_stages == 1): - while copy_token < Int64(m): - token_pack = copy_token * self.packs_per_token - copy_values = [] - copy_sources = [] - if cutlass.const_expr(self.write_residual_output): - copy_destinations = [] - for item in cutlass.range_constexpr(self.rms_vectors_per_thread): - pack = rms_pack_base + item * rms_pack_stride - linear_pack = token_pack + pack - source_address = Int64( - (prenorm_mailbox.iterator + linear_pack * VEC_BF16).toint() - ) - copy_sources.append(source_address) - if cutlass.const_expr(self.write_residual_output): - copy_destinations.append( - Int64( - ( - residual_output.iterator - + linear_pack * VEC_BF16 - ).toint() - ) - ) - copy_fragments = cute.make_rmem_tensor( - cute.make_layout( - (self.rms_vectors_per_thread, 4), - stride=(4, 1), - ), - Uint32, - ) - all_ready = cutlass.Boolean(0) - while not all_ready: - all_ready = cutlass.Boolean(1) - for item in cutlass.range_constexpr( - self.rms_vectors_per_thread - ): - packed = load_global_u32x4_address( - copy_sources[item], - volatile=True, - ) - copy_fragments[item, None].store(packed) - all_ready = all_ready and ( - not fragment_has_negative_zero(packed) - ) - for item in cutlass.range_constexpr(self.rms_vectors_per_thread): - copy_values.append(copy_fragments[item, None].load()) - prenorm_fragments = cute.make_rmem_tensor( - cute.make_layout( - (self.rms_vectors_per_thread, VEC_BF16), - stride=(VEC_BF16, 1), - ), - BFloat16, - ) - thread_sum = Float32(0.0) - if cutlass.const_expr(self.write_residual_output): - prenorm_packed = [] - for item in cutlass.range_constexpr(self.rms_vectors_per_thread): - pack = rms_pack_base + item * rms_pack_stride - prenorm = packed_u32x4_to_bf16x8(copy_values[item]) - prenorm_fragments[item, None].store(prenorm) - packed_prenorm = bf16x8_to_packed_u32x4(prenorm) - if cutlass.const_expr(self.write_residual_output): - prenorm_packed.append(packed_prenorm) - prenorm_f32 = prenorm.to(Float32) - thread_sum = thread_sum + (prenorm_f32 * prenorm_f32).reduce( - cute.ReductionOp.ADD, - init_val=Float32(0.0), - reduction_profile=0, - ) - if cutlass.const_expr(self.write_residual_output): - for item in cutlass.range_constexpr( - self.rms_vectors_per_thread - ): - store_global_u32x4( - copy_destinations[item], prenorm_packed[item] - ) - warp_sum = cute.arch.warp_reduction_sum(thread_sum) - if lane == 0: - cute.arch.store((norm_warp_sums + warp - 1).llvm_ptr, warp_sum) - self._rms_arrive_and_wait(rms_group) - if warp == 1 + rms_group * self.rms_warps_per_token: - cta_sum = Float32(0.0) - if lane < self.rms_warps_per_token: - cta_sum = cute.arch.load( - ( - norm_warp_sums - + rms_group * self.rms_warps_per_token - + lane - ).llvm_ptr, - Float32, - ) - cta_sum = cute.arch.warp_reduction_sum(cta_sum) - if lane == 0: - inv_rms = cute.math.rsqrt( - cta_sum / Float32(self.hidden) - + Float32(self.rms_epsilon), - fastmath=True, - ) - cute.arch.store( - (norm_inv_rms + rms_group).llvm_ptr, inv_rms - ) - self._rms_arrive_and_wait(rms_group) - inv_rms = cute.arch.load( - (norm_inv_rms + rms_group).llvm_ptr, Float32 - ) - gamma_values = [] - for item in cutlass.range_constexpr(self.rms_vectors_per_thread): - pack = rms_pack_base + item * rms_pack_stride - gamma_value = packed_u32x4_to_bf16x8( - load_global_u32x4_address( - Int64((gamma.iterator + pack * VEC_BF16).toint()) - ) - ).to(Float32) - if cutlass.const_expr(self.weight_bias != 0.0): - gamma_value = gamma_value + Float32(self.weight_bias) - gamma_values.append(gamma_value) - norm_packed = [] - for item in cutlass.range_constexpr(self.rms_vectors_per_thread): - pack = rms_pack_base + item * rms_pack_stride - linear_pack = token_pack + pack - prenorm_for_norm = ( - prenorm_fragments[item, None].load().to(Float32) - ) - result = (prenorm_for_norm * inv_rms * gamma_values[item]).to( - BFloat16 - ) - norm_packed.append(bf16x8_to_packed_u32x4(result)) - for item in cutlass.range_constexpr(self.rms_vectors_per_thread): - pack = rms_pack_base + item * rms_pack_stride - linear_pack = token_pack + pack - store_global_u32x4( - Int64( - (norm_output.iterator + linear_pack * VEC_BF16).toint() - ), - norm_packed[item], - ) - copy_wave += self.cta_groups * self.rms_token_groups - copy_token = copy_wave * self.tp + cta_slot - elif warp == self.publisher_warp: - owner_ready_address = cute.arch.load( - (ready_counter_peer_addresses.iterator + cta_slot).llvm_ptr, - Int64, - ) - first_token = Int64(cta_group) * self.tp + cta_slot - token_count = Int32(0) - if first_token < Int64(m): - token_count = Int32( - (Int64(m) + self.active_ctas - 1 - first_token) // self.active_ctas - ) - published = Int32(0) - while published < token_count: - observed = token_count - if lane < self.consumer_warps: - observed = cute.arch.load( - (consumer_progress + lane).llvm_ptr, - Int32, - sem="relaxed", - scope="cta", - ) - frontier = cute.arch.warp_reduction( - observed, lambda x, y: cutlass.min(x, y) - ) - if frontier > published: - acquired = token_count - if lane < self.consumer_warps: - acquired = cute.arch.load( - (consumer_progress + lane).llvm_ptr, - Int32, - sem="acquire", - scope="cta", - ) - frontier = cute.arch.warp_reduction( - acquired, lambda x, y: cutlass.min(x, y) - ) - cute.arch.sync_warp() - batch = cutlass.min(frontier - published, Int32(WARP_SIZE)) - if lane < batch: - sequence = Int64(published + lane) - publish_token = ( - Int64(cta_group) + sequence * self.cta_groups - ) * self.tp + cta_slot - owner_token = publish_token // self.tp - remote_release_add1_u32(owner_ready_address + owner_token * 4) - published += batch - elif ( - block < self.reduction_ctas - and warp >= self.reduction_warp_begin - and (warp < self.reduction_warp_begin + self.reduction_warps) - ): - reduction_warp = warp - self.reduction_warp_begin - reduction_tid = reduction_warp * WARP_SIZE + lane - reduction_barrier = pipeline.NamedBarrier( - barrier_id=1, num_threads=self.reduction_threads - ) - reduction_shard = block % self.tp - local_token = Int64(block // self.tp) - token = local_token * self.tp + self.rank - while token < Int64(m): - processed_index = local_token * self.tp + reduction_shard - target = Uint32(0) - if reduction_tid == 0: - ready_counter_address = ( - ready_counters.iterator + local_token - ).llvm_ptr - processed_counter_address = ( - processed_counters.iterator + processed_index - ).llvm_ptr - target = cute.arch.load(processed_counter_address, Uint32) + Uint32( - self.tp - ) - observed = Uint32(0) - while observed != target: - observed = cute.arch.load( - ready_counter_address, - Uint32, - sem="relaxed", - scope="sys", - ) - cute.arch.load( - ready_counter_address, - Uint32, - sem="acquire", - scope="sys", - ) - reduction_barrier.arrive_and_wait() - values = [] - addresses = [] - token_pack = token * self.packs_per_token - shard_pack = reduction_shard * self.packs_per_reduction_shard - for item in cutlass.range_constexpr(self.reduction_vectors_per_thread): - pack = shard_pack + reduction_tid + item * self.reduction_threads - input_address = ( - local_contributions_multicast_address + (token_pack + pack) * 16 - ) - output_address = ( - prenorm_mailbox_multicast_address + (token_pack + pack) * 16 - ) - reduced_packed = ldmc_bf16x8(input_address) - reduced_values = packed_u32x4_to_bf16x8(reduced_packed).to(Float32) - if cutlass.const_expr(self.add_residual): - residual_values = packed_u32x4_to_bf16x8( - load_global_u32x4_address( - Int64( - ( - residual_source.iterator - + (token_pack + pack) * VEC_BF16 - ).toint() - ) - ) - ).to(Float32) - reduced_values = reduced_values + residual_values - reduced_packed = bf16x8_to_packed_u32x4(reduced_values.to(BFloat16)) - values.append(sanitize_negative_zero_u32x4(reduced_packed)) - addresses.append(output_address) - for item in cutlass.range_constexpr(self.reduction_vectors_per_thread): - stmc_bf16x8(addresses[item], values[item]) - if reduction_tid == 0: - cute.arch.store( - (processed_counters.iterator + processed_index).llvm_ptr, - target, - ) - local_token += self.reduction_cta_groups - token = local_token * self.tp + self.rank - cute.arch.sync_threads() - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_launch_dependents() diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ht/protocol.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ht/protocol.py deleted file mode 100644 index b024387b9..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ht/protocol.py +++ /dev/null @@ -1,469 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""High-throughput MNNVL protocol and its two operation paths.""" - -from __future__ import annotations - -from dataclasses import dataclass -from typing import Any, TypedDict, cast - -import cutlass.cute as cute -import torch -import torch.distributed as dist -from cutlass import BFloat16, Int32, Int64, Uint32 -from cutlass.cute.runtime import make_fake_compact_tensor - -from ..runtime import ( - current_cu_stream, - make_fake_dynamic_compact_tensor, - to_cute, - to_cute_dynamic, -) -from ..symmetric_buffer import SymmetricBuffer -from .device_kernel import _MoeFinalizeAllReduceRMSNormHTDeviceKernel - - -@dataclass(frozen=True, slots=True) -class HTFinalizeTuning: - persistent_ctas: int | None = None - consumer_threads: int = 512 - vectors_per_thread: int = 2 - stages: int = 6 - reduction_warps: int = 1 - reduction_cta_groups: int | None = None - rms_token_groups: int = 2 - rms_pipeline_stages: int = 2 - rms_shard_major: bool = False - enable_pdl: bool = True - - -@dataclass(frozen=True, slots=True) -class HTAllReduceTuning: - persistent_ctas: int | None = None - consumer_threads: int = 512 - vectors_per_thread: int = 2 - stages: int = 2 - reduction_warps: int = 2 - reduction_cta_groups: int | None = None - rms_token_groups: int = 2 - rms_pipeline_stages: int = 1 - rms_shard_major: bool = False - enable_pdl: bool = True - - -HT_FINALIZE_GB300_TP8_H8192_K10_M_LE_2048 = HTFinalizeTuning( - stages=7, - reduction_warps=2, - rms_pipeline_stages=3, - rms_shard_major=True, -) -HT_FINALIZE_GB300_TP8_H8192_K10_M_GE_2049 = HTFinalizeTuning() -HT_FINALIZE_GB300_TP8_H8192_K10 = HT_FINALIZE_GB300_TP8_H8192_K10_M_GE_2049 -HT_FINALIZE_GB300_TP16_H8192_K10 = HTFinalizeTuning( - stages=7, - reduction_warps=2, - rms_pipeline_stages=3, - rms_shard_major=True, -) -HT_ALL_REDUCE_GB300_TP8_H8192 = HTAllReduceTuning() -HT_ALL_REDUCE_GB300_TP16_H8192 = HTAllReduceTuning() - - -@dataclass(slots=True) -class HTProtocolState: - local_contributions: SymmetricBuffer - prenorm_mailbox: SymmetricBuffer - routed_ready_counters: SymmetricBuffer - routed_processed_counters: torch.Tensor - all_reduce_ready_counters: SymmetricBuffer - all_reduce_processed_counters: torch.Tensor - - -class _PathKwargs(TypedDict): - hidden_size: int - top_k: int - capacity_m: int - write_residual_output: bool - - -class _HTPath: - def __init__( - self, - *, - compiled: Any, - hidden_size: int, - top_k: int, - capacity_m: int, - write_residual_output: bool, - ) -> None: - self._compiled = compiled - self.hidden_size = hidden_size - self.top_k = top_k - self.capacity_m = capacity_m - self.write_residual_output = write_residual_output - - def _outputs( - self, - m: int, - norm_output: torch.Tensor | None, - residual_output: torch.Tensor | None, - ) -> tuple[torch.Tensor, torch.Tensor | None]: - shape = (m, self.hidden_size) - device = torch.device("cuda", torch.cuda.current_device()) - if norm_output is None: - norm_output = torch.empty(shape, dtype=torch.bfloat16, device=device) - if self.write_residual_output and residual_output is None: - residual_output = torch.empty(shape, dtype=torch.bfloat16, device=device) - return norm_output, residual_output - - def _state_buffers( - self, state: HTProtocolState - ) -> tuple[SymmetricBuffer, SymmetricBuffer]: - return state.local_contributions, state.prenorm_mailbox - - def _validate_m(self, m: int) -> None: - if not 1 <= m <= self.capacity_m: - raise ValueError(f"m must be in [1, {self.capacity_m}]") - - -class FinalizeAllReduceRMSNormHTKernel(_HTPath): - def __call__( - self, - routed_output: torch.Tensor, - expert_weights: torch.Tensor, - permuted_indices: torch.Tensor, - shared_output: torch.Tensor | None, - residual_source: torch.Tensor | None, - gamma: torch.Tensor, - m: int, - *, - state: HTProtocolState, - norm_output: torch.Tensor | None = None, - residual_output: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor | None]: - self._validate_m(m) - norm_output, residual_output = self._outputs(m, norm_output, residual_output) - local, prenorm = self._state_buffers(state) - peers = cast(torch.Tensor, state.routed_ready_counters.peer_addresses) - shared_arg = shared_output if shared_output is not None else norm_output - residual_arg = residual_source if residual_source is not None else norm_output - residual_output_arg = ( - residual_output if residual_output is not None else norm_output - ) - self._compiled( - to_cute_dynamic(routed_output.flatten(), 16, divisibility=self.hidden_size), - to_cute_dynamic(expert_weights.flatten(), 2, divisibility=self.top_k), - to_cute_dynamic(permuted_indices.flatten(), 4, divisibility=self.top_k), - to_cute_dynamic(shared_arg.flatten(), 16, divisibility=self.hidden_size), - to_cute_dynamic(residual_arg.flatten(), 16, divisibility=self.hidden_size), - to_cute(gamma, 16), - to_cute(local.tensor.flatten(), 16), - to_cute(prenorm.tensor.flatten(), 16), - to_cute_dynamic( - residual_output_arg.flatten(), - 16, - divisibility=self.hidden_size, - ), - to_cute_dynamic(norm_output.flatten(), 16, divisibility=self.hidden_size), - to_cute(peers, 8), - to_cute(state.routed_ready_counters.tensor, 4), - to_cute(state.routed_processed_counters.flatten(), 4), - Int64(cast(int, local.multicast_address)), - Int64(cast(int, prenorm.multicast_address)), - Int32(m), - current_cu_stream(), - ) - return norm_output, residual_output - - -class AllReduceRMSNormHTKernel(_HTPath): - def __call__( - self, - local_contribution: torch.Tensor, - residual_source: torch.Tensor | None, - gamma: torch.Tensor, - m: int, - *, - state: HTProtocolState, - norm_output: torch.Tensor | None = None, - residual_output: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor | None]: - self._validate_m(m) - norm_output, residual_output = self._outputs(m, norm_output, residual_output) - local, prenorm = self._state_buffers(state) - peers = cast(torch.Tensor, state.all_reduce_ready_counters.peer_addresses) - residual_arg = residual_source if residual_source is not None else norm_output - residual_output_arg = ( - residual_output if residual_output is not None else norm_output - ) - index_arg = state.all_reduce_processed_counters.view(torch.int32) - # top_k=0 disables metadata reads, so the aliased placeholders stay unused. - self._compiled( - to_cute_dynamic( - local_contribution.flatten(), - 16, - divisibility=self.hidden_size, - ), - to_cute_dynamic(local_contribution.flatten(), 2, divisibility=1), - to_cute_dynamic(index_arg.flatten(), 4, divisibility=1), - to_cute_dynamic( - local_contribution.flatten(), - 16, - divisibility=self.hidden_size, - ), - to_cute_dynamic(residual_arg.flatten(), 16, divisibility=self.hidden_size), - to_cute(gamma, 16), - to_cute(local.tensor.flatten(), 16), - to_cute(prenorm.tensor.flatten(), 16), - to_cute_dynamic( - residual_output_arg.flatten(), - 16, - divisibility=self.hidden_size, - ), - to_cute_dynamic(norm_output.flatten(), 16, divisibility=self.hidden_size), - to_cute(peers, 8), - to_cute(state.all_reduce_ready_counters.tensor, 4), - to_cute(state.all_reduce_processed_counters.flatten(), 4), - Int64(cast(int, local.multicast_address)), - Int64(cast(int, prenorm.multicast_address)), - Int32(m), - current_cu_stream(), - ) - return norm_output, residual_output - - -class HTProtocol: - """Own tuning-independent HT State and both persistent path variants.""" - - def __init__( - self, - hidden_size: int, - top_k: int, - tp_size: int, - rank: int, - capacity_m: int, - rms_epsilon: float, - routed_scaling_factor: float, - weight_bias: float, - *, - include_shared_expert: bool, - add_residual: bool, - write_residual_output: bool, - finalize_tunings: tuple[HTFinalizeTuning, ...], - all_reduce_tunings: tuple[HTAllReduceTuning, ...], - group: dist.ProcessGroup, - ) -> None: - self.hidden_size = hidden_size - self.top_k = top_k - self.tp_size = tp_size - self.rank = rank - self.capacity_m = capacity_m - self.rms_epsilon = rms_epsilon - self.routed_scaling_factor = routed_scaling_factor - self.weight_bias = weight_bias - self.include_shared_expert = include_shared_expert - self.add_residual = add_residual - self.write_residual_output = write_residual_output - - self.finalize_kernels = { - tuning: FinalizeAllReduceRMSNormHTKernel( - compiled=self._compile_finalize(tuning), - **self._path_kwargs(), - ) - for tuning in dict.fromkeys(finalize_tunings) - } - self.all_reduce_kernels = { - tuning: AllReduceRMSNormHTKernel( - compiled=self._compile_all_reduce(tuning), - **self._path_kwargs(), - ) - for tuning in dict.fromkeys(all_reduce_tunings) - } - self.state = self._create_state(group) - - def _path_kwargs(self) -> _PathKwargs: - return { - "hidden_size": self.hidden_size, - "top_k": self.top_k, - "capacity_m": self.capacity_m, - "write_residual_output": self.write_residual_output, - } - - def _resolve_ctas(self, persistent_ctas: int | None) -> int: - sm_count = torch.cuda.get_device_properties( - torch.cuda.current_device() - ).multi_processor_count - # min_blocks_per_mp=1 guarantees one resident CTA per SM for this kernel. - resident_ctas = (sm_count // self.tp_size) * self.tp_size - if resident_ctas == 0: - raise ValueError("tp_size exceeds the available SM count") - if persistent_ctas is None: - return resident_ctas - if persistent_ctas <= 0 or persistent_ctas % self.tp_size: - raise ValueError( - "persistent_ctas must be positive and divisible by tp_size" - ) - return min(persistent_ctas, resident_ctas) - - def _compile_finalize(self, tuning: HTFinalizeTuning): - active_ctas = self._resolve_ctas(tuning.persistent_ctas) - groups = tuning.reduction_cta_groups or active_ctas // self.tp_size - kernel = _MoeFinalizeAllReduceRMSNormHTDeviceKernel( - hidden=self.hidden_size, - top_k=self.top_k, - tp=self.tp_size, - rank=self.rank, - active_ctas=active_ctas, - stages=tuning.stages, - consumer_threads=tuning.consumer_threads, - vectors_per_thread=tuning.vectors_per_thread, - reduction_warps=tuning.reduction_warps, - reduction_cta_groups=groups, - rms_token_groups=tuning.rms_token_groups, - rms_pipeline_stages=tuning.rms_pipeline_stages, - rms_shard_major=tuning.rms_shard_major, - rms_epsilon=self.rms_epsilon, - routed_scaling_factor=self.routed_scaling_factor, - weight_bias=self.weight_bias, - include_shared_expert=self.include_shared_expert, - add_residual=self.add_residual, - write_residual_output=self.write_residual_output, - enable_pdl=tuning.enable_pdl, - ) - return self._compile(kernel, top_k=self.top_k) - - def _compile_all_reduce(self, tuning: HTAllReduceTuning): - active_ctas = self._resolve_ctas(tuning.persistent_ctas) - groups = tuning.reduction_cta_groups or active_ctas // self.tp_size - kernel = _MoeFinalizeAllReduceRMSNormHTDeviceKernel( - hidden=self.hidden_size, - top_k=0, - tp=self.tp_size, - rank=self.rank, - active_ctas=active_ctas, - stages=tuning.stages, - consumer_threads=tuning.consumer_threads, - vectors_per_thread=tuning.vectors_per_thread, - reduction_warps=tuning.reduction_warps, - reduction_cta_groups=groups, - rms_token_groups=tuning.rms_token_groups, - rms_pipeline_stages=tuning.rms_pipeline_stages, - rms_shard_major=tuning.rms_shard_major, - rms_epsilon=self.rms_epsilon, - routed_scaling_factor=1.0, - weight_bias=self.weight_bias, - include_shared_expert=True, - add_residual=self.add_residual, - write_residual_output=self.write_residual_output, - enable_pdl=tuning.enable_pdl, - ) - return self._compile(kernel, top_k=0) - - def _compile( - self, - kernel: _MoeFinalizeAllReduceRMSNormHTDeviceKernel, - *, - top_k: int, - ): - activation = self.capacity_m * self.hidden_size - token_slots = (self.capacity_m + self.tp_size - 1) // self.tp_size - args = ( - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_dynamic_compact_tensor( - BFloat16, alignment=2, divisibility=max(top_k, 1) - ), - make_fake_dynamic_compact_tensor( - Int32, alignment=4, divisibility=max(top_k, 1) - ), - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_compact_tensor(BFloat16, (self.hidden_size,), assumed_align=16), - make_fake_compact_tensor(BFloat16, (activation,), assumed_align=16), - make_fake_compact_tensor(BFloat16, (activation,), assumed_align=16), - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_compact_tensor(Int64, (self.tp_size,), assumed_align=8), - make_fake_compact_tensor(Uint32, (token_slots,), assumed_align=4), - make_fake_compact_tensor( - Uint32, (token_slots * self.tp_size,), assumed_align=4 - ), - Int64(0), - Int64(0), - Int32(self.capacity_m), - current_cu_stream(), - ) - return cute.compile(kernel, *args) - - def _allocate_large_buffers( - self, group: dist.ProcessGroup - ) -> tuple[SymmetricBuffer, SymmetricBuffer]: - shape = (self.capacity_m, self.hidden_size) - device = torch.device("cuda", torch.cuda.current_device()) - return ( - SymmetricBuffer.allocate( - shape, - torch.bfloat16, - device, - group, - require_multicast=True, - ), - SymmetricBuffer.allocate( - shape, - torch.bfloat16, - device, - group, - require_multicast=True, - ), - ) - - def _create_state(self, group: dist.ProcessGroup) -> HTProtocolState: - device = torch.device("cuda", torch.cuda.current_device()) - token_slots = (self.capacity_m + self.tp_size - 1) // self.tp_size - - def counters() -> tuple[SymmetricBuffer, torch.Tensor]: - ready = SymmetricBuffer.allocate( - (token_slots,), - torch.uint32, - device, - group, - materialize_peer_addresses=True, - ) - ready.tensor.zero_() - processed = torch.zeros( - (token_slots, self.tp_size), dtype=torch.uint32, device=device - ) - return ready, processed - - routed_ready, routed_processed = counters() - all_reduce_ready, all_reduce_processed = counters() - local, prenorm = self._allocate_large_buffers(group) - return HTProtocolState( - local_contributions=local, - prenorm_mailbox=prenorm, - routed_ready_counters=routed_ready, - routed_processed_counters=routed_processed, - all_reduce_ready_counters=all_reduce_ready, - all_reduce_processed_counters=all_reduce_processed, - ) diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ll/__init__.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ll/__init__.py deleted file mode 100644 index 11c888a6e..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ll/__init__.py +++ /dev/null @@ -1,47 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Low-latency MNNVL protocol.""" - -from .protocol import ( - LL_ALL_REDUCE_GB300_TP8_H8192, - LL_ALL_REDUCE_GB300_TP8_H8192_M_GE_5, - LL_ALL_REDUCE_GB300_TP8_H8192_M_LE_4, - LL_ALL_REDUCE_GB300_TP16_H8192, - LL_ALL_REDUCE_GB300_TP16_H8192_M_11_TO_17, - LL_ALL_REDUCE_GB300_TP16_H8192_M_GE_18, - LL_ALL_REDUCE_GB300_TP16_H8192_M_LE_10, - LL_FINALIZE_GB300_TP8_H8192_K10, - LL_FINALIZE_GB300_TP8_H8192_K10_M_GE_20, - LL_FINALIZE_GB300_TP16_H8192_K10, - LLAllReduceTuning, - LLCollectiveTuning, - LLFinalizeTuning, -) - -__all__ = [ - "LL_ALL_REDUCE_GB300_TP8_H8192", - "LL_ALL_REDUCE_GB300_TP8_H8192_M_GE_5", - "LL_ALL_REDUCE_GB300_TP8_H8192_M_LE_4", - "LL_ALL_REDUCE_GB300_TP16_H8192", - "LL_ALL_REDUCE_GB300_TP16_H8192_M_11_TO_17", - "LL_ALL_REDUCE_GB300_TP16_H8192_M_GE_18", - "LL_ALL_REDUCE_GB300_TP16_H8192_M_LE_10", - "LL_FINALIZE_GB300_TP8_H8192_K10", - "LL_FINALIZE_GB300_TP8_H8192_K10_M_GE_20", - "LL_FINALIZE_GB300_TP16_H8192_K10", - "LLAllReduceTuning", - "LLCollectiveTuning", - "LLFinalizeTuning", -] diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ll/device_kernels.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ll/device_kernels.py deleted file mode 100644 index 92b6b1755..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ll/device_kernels.py +++ /dev/null @@ -1,1003 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Split low-latency BF16 MoE finalize, TP reduction, and RMSNorm.""" - -from __future__ import annotations - -import math - -import cuda.bindings.driver as cuda -import cutlass -import cutlass.cute as cute -from cutlass import BFloat16, Float32, Int32, Int64, Uint32 - -from ..cute_dsl_primitives import ( - NEGATIVE_ZERO_BF16_BITS, - QUAD_BF16, - VEC_BF16, - WARP_SIZE, - bf16x4_to_packed_u32x2, - bf16x8_to_packed_u32x4, - f32_to_bf16_bits, - fragment_has_negative_zero, - load_global_bf16_as_f32, - load_global_bf16_as_f32_predicated, - load_global_u32x2, - load_global_u32x4, - load_volatile_u32, - map_shared_to_peer, - packed_u32x2_to_bf16x4, - packed_u32x4_to_bf16x8, - sanitize_negative_zero_u32x2, - sanitize_negative_zero_u32x4, - shuffle_sync_idx_u32, - stmc_bf16x2, - stmc_bf16x4, - stmc_bf16x8, - store_global_u32, - store_global_u32x4, - store_lamport_sentinel_u32x4, - store_shared_cluster_f32, -) - -LAMPORT_GENERATIONS = 3 -NEXT_STAGE = 0 -ACTIVE_STAGE = 1 - - -@cute.jit -def _group_leader_block_sum( - value: Float32, - warp_sums: cute.Tensor, - warps: cutlass.Constexpr[int], - leader_stride: cutlass.Constexpr[int], -) -> Float32: - lane = cute.arch.lane_idx() - warp = cute.arch.warp_idx() - for offset in cutlass.range_constexpr(1, WARP_SIZE): - if cutlass.const_expr(offset >= leader_stride and (offset & (offset - 1)) == 0): - value = value + cute.arch.shuffle_sync_bfly( - value, - offset=offset, - mask=-1, - mask_and_clamp=31, - ) - if lane == 0: - cute.arch.store((warp_sums + warp).llvm_ptr, value) - cute.arch.barrier() - - result = Float32(0.0) - if warp == 0: - if lane < Int32(warps): - result = cute.arch.load( - (warp_sums + lane).llvm_ptr, - Float32, - ) - result = cute.arch.warp_reduction_sum(result) - if lane == 0: - cute.arch.store(warp_sums.llvm_ptr, result) - cute.arch.barrier() - return cute.arch.load(warp_sums.llvm_ptr, Float32) - - -class _ScalarFinalizePublishDeviceKernel: - def __init__( - self, - *, - hidden: int, - top_k: int, - tp: int, - rank: int, - capacity_m: int, - threads: int, - routed_scaling_factor: float, - include_shared_expert: bool, - load_shared_expert_before_pdl: bool, - enable_pdl: bool, - prefetch_group: int, - ) -> None: - if hidden <= 0 or hidden % 2: - raise ValueError("hidden must be a positive multiple of 2") - self.hidden = hidden - self.top_k = top_k - self.tp = tp - self.rank = rank - self.capacity_m = capacity_m - self.threads = threads - self.routed_scaling_factor = routed_scaling_factor - self.ctas_per_token = math.ceil(hidden / threads) - self.include_shared_expert = include_shared_expert - self.load_shared_expert_before_pdl = load_shared_expert_before_pdl - self.enable_pdl = enable_pdl - self.prefetch_group = prefetch_group - self.prefetch_groups = (top_k + prefetch_group - 1) // prefetch_group - - def smem_size_in_bytes(self) -> int: - return self.top_k * 8 - - @cute.jit - def __call__( - self, - routed_output: cute.Tensor, - expert_weights: cute.Tensor, - permuted_indices: cute.Tensor, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_multicast_address: Int64, - m: Int32, - stream: cuda.CUstream, - ) -> None: - self.kernel( - routed_output, - expert_weights, - permuted_indices, - shared_output, - stage_state, - contribution_mailbox_multicast_address, - ).launch( - grid=(m * self.ctas_per_token, 1, 1), - block=(self.threads, 1, 1), - smem=self.smem_size_in_bytes(), - stream=stream, - use_pdl=self.enable_pdl, - ) - - @cute.kernel - def kernel( - self, - routed_output: cute.Tensor, - expert_weights: cute.Tensor, - permuted_indices: cute.Tensor, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_multicast_address: Int64, - ) -> None: - block, _, _ = cute.arch.block_idx() - tidx, _, _ = cute.arch.thread_idx() - token = block // self.ctas_per_token - cta_in_token = block % self.ctas_per_token - hidden_index = cta_in_token * self.threads + tidx - - smem = cutlass.utils.SmemAllocator() - staged_indices = smem.allocate_array(Int32, self.top_k) - staged_weights = smem.allocate_array(Float32, self.top_k) - metadata_index = Int32(tidx) - while metadata_index < Int32(self.top_k): - element = Int64(token) * self.top_k + Int64(metadata_index) - row = cute.arch.load( - (permuted_indices.iterator + element).llvm_ptr, - Int32, - ) - weight = load_global_bf16_as_f32( - Int64((expert_weights.iterator + element).toint()) - ) - if cutlass.const_expr(self.routed_scaling_factor != 1.0): - weight = weight * Float32(self.routed_scaling_factor) - if row == Int32(-1): - weight = Float32(0.0) - cute.arch.store( - (staged_indices + metadata_index).llvm_ptr, - row, - ) - cute.arch.store( - (staged_weights + metadata_index).llvm_ptr, - weight, - ) - metadata_index = metadata_index + self.threads - cute.arch.barrier() - - shared_value = Float32(0.0) - if cutlass.const_expr( - self.include_shared_expert and self.load_shared_expert_before_pdl - ): - if hidden_index < Int32(self.hidden): - shared_element = Int64(token) * self.hidden + Int64(hidden_index) - shared_value = load_global_bf16_as_f32( - Int64((shared_output.iterator + shared_element).toint()) - ) - - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_wait() - - if cutlass.const_expr( - self.include_shared_expert and not self.load_shared_expert_before_pdl - ): - if hidden_index < Int32(self.hidden): - shared_element = Int64(token) * self.hidden + Int64(hidden_index) - shared_value = load_global_bf16_as_f32( - Int64((shared_output.iterator + shared_element).toint()) - ) - - stage = load_volatile_u32(stage_state.iterator + NEXT_STAGE) - bits = Uint32(0) - if hidden_index < Int32(self.hidden): - accumulator = Float32(0.0) - if cutlass.const_expr(self.prefetch_group == 1): - for k in cutlass.range_constexpr(self.top_k): - row = cute.arch.load((staged_indices + k).llvm_ptr, Int32) - weight = cute.arch.load( - (staged_weights + k).llvm_ptr, - Float32, - ) - source_element = Int64(row) * self.hidden + Int64(hidden_index) - accumulator = ( - accumulator - + load_global_bf16_as_f32_predicated( - Int64((routed_output.iterator + source_element).toint()), - Int32(row != Int32(-1)), - ) - * weight - ) - else: - inputs = cute.make_rmem_tensor( - cute.make_layout((self.prefetch_group,)), - Float32, - ) - inputs.fill(Float32(0.0)) - for group in cutlass.range_constexpr(self.prefetch_groups): - for item in cutlass.range_constexpr(self.prefetch_group): - k = group * self.prefetch_group + item - if cutlass.const_expr(k < self.top_k): - row = cute.arch.load( - (staged_indices + k).llvm_ptr, - Int32, - ) - source_element = Int64(row) * self.hidden + Int64( - hidden_index - ) - inputs[item] = load_global_bf16_as_f32_predicated( - Int64( - (routed_output.iterator + source_element).toint() - ), - Int32(row != Int32(-1)), - ) - for item in cutlass.range_constexpr(self.prefetch_group): - k = group * self.prefetch_group + item - if cutlass.const_expr(k < self.top_k): - accumulator = accumulator + inputs[item] * cute.arch.load( - (staged_weights + k).llvm_ptr, - Float32, - ) - if cutlass.const_expr(self.include_shared_expert): - accumulator = accumulator + shared_value - bits = f32_to_bf16_bits(accumulator) - if bits == Uint32(NEGATIVE_ZERO_BF16_BITS): - bits = Uint32(0) - - # Unlike the quad path, this publishes first; Lamport sentinels gate consumers. - if block == 0 and tidx == 0: - store_global_u32( - stage_state.iterator + ACTIVE_STAGE, - stage, - ) - cute.arch.barrier() - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_launch_dependents() - - lane = cute.arch.lane_idx() - partner_bits = shuffle_sync_idx_u32( - bits, - Int32(lane | Int32(1)), - ) - if (lane & Int32(1)) == Int32(0) and hidden_index < Int32(self.hidden): - packed = bits | (partner_bits << Uint32(16)) - mailbox_element = ( - (Int64(stage) * self.tp + self.rank) * self.capacity_m + Int64(token) - ) * self.hidden + Int64(hidden_index) - stmc_bf16x2( - contribution_mailbox_multicast_address + mailbox_element * 2, - packed, - ) - - -class _QuadFinalizePublishDeviceKernel: - def __init__( - self, - *, - hidden: int, - top_k: int, - tp: int, - rank: int, - capacity_m: int, - threads: int, - routed_scaling_factor: float, - include_shared_expert: bool, - load_shared_expert_before_pdl: bool, - enable_pdl: bool, - prefetch_group: int, - ) -> None: - if hidden <= 0 or hidden % QUAD_BF16: - raise ValueError("hidden must be a positive multiple of 4") - self.hidden = hidden - self.top_k = top_k - self.tp = tp - self.rank = rank - self.capacity_m = capacity_m - self.threads = threads - self.routed_scaling_factor = routed_scaling_factor - self.fragments = hidden // QUAD_BF16 - self.ctas_per_token = math.ceil(self.fragments / threads) - self.include_shared_expert = include_shared_expert - self.load_shared_expert_before_pdl = load_shared_expert_before_pdl - self.enable_pdl = enable_pdl - self.prefetch_group = prefetch_group - self.prefetch_groups = (top_k + prefetch_group - 1) // prefetch_group - - def smem_size_in_bytes(self) -> int: - return self.top_k * 8 - - @cute.jit - def __call__( - self, - routed_output: cute.Tensor, - expert_weights: cute.Tensor, - permuted_indices: cute.Tensor, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_multicast_address: Int64, - m: Int32, - stream: cuda.CUstream, - ) -> None: - self.kernel( - routed_output, - expert_weights, - permuted_indices, - shared_output, - stage_state, - contribution_mailbox_multicast_address, - ).launch( - grid=(m * self.ctas_per_token, 1, 1), - block=(self.threads, 1, 1), - smem=self.smem_size_in_bytes(), - stream=stream, - use_pdl=self.enable_pdl, - ) - - @cute.kernel - def kernel( - self, - routed_output: cute.Tensor, - expert_weights: cute.Tensor, - permuted_indices: cute.Tensor, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_multicast_address: Int64, - ) -> None: - block, _, _ = cute.arch.block_idx() - tidx, _, _ = cute.arch.thread_idx() - token = block // self.ctas_per_token - cta_in_token = block % self.ctas_per_token - fragment = cta_in_token * self.threads + tidx - - smem = cutlass.utils.SmemAllocator() - staged_indices = smem.allocate_array(Int32, self.top_k) - staged_weights = smem.allocate_array(Float32, self.top_k) - metadata_index = Int32(tidx) - while metadata_index < Int32(self.top_k): - element = Int64(token) * self.top_k + Int64(metadata_index) - row = cute.arch.load( - (permuted_indices.iterator + element).llvm_ptr, - Int32, - ) - weight = load_global_bf16_as_f32( - Int64((expert_weights.iterator + element).toint()) - ) - if cutlass.const_expr(self.routed_scaling_factor != 1.0): - weight = weight * Float32(self.routed_scaling_factor) - if row == Int32(-1): - weight = Float32(0.0) - cute.arch.store( - (staged_indices + metadata_index).llvm_ptr, - row, - ) - cute.arch.store( - (staged_weights + metadata_index).llvm_ptr, - weight, - ) - metadata_index = metadata_index + self.threads - cute.arch.barrier() - - if cutlass.const_expr(self.include_shared_expert): - shared_values = cute.make_rmem_tensor( - cute.make_layout((QUAD_BF16,)), BFloat16 - ) - shared_values.fill(BFloat16(0.0)) - if cutlass.const_expr(self.load_shared_expert_before_pdl): - if fragment < self.fragments: - shared_element = ( - Int64(token) * self.hidden + Int64(fragment) * QUAD_BF16 - ) - shared_pointer = cute.make_ptr( - BFloat16, - (shared_output.iterator + shared_element).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=8, - ) - shared_values.store( - packed_u32x2_to_bf16x4(load_global_u32x2(shared_pointer)) - ) - - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_wait() - - if cutlass.const_expr( - self.include_shared_expert and not self.load_shared_expert_before_pdl - ): - if fragment < self.fragments: - shared_element = ( - Int64(token) * self.hidden + Int64(fragment) * QUAD_BF16 - ) - shared_pointer = cute.make_ptr( - BFloat16, - (shared_output.iterator + shared_element).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=8, - ) - shared_values.store( - packed_u32x2_to_bf16x4(load_global_u32x2(shared_pointer)) - ) - - stage = load_volatile_u32(stage_state.iterator + NEXT_STAGE) - if fragment < self.fragments: - accumulator = cute.make_rmem_tensor(cute.make_layout((QUAD_BF16,)), Float32) - accumulator.fill(Float32(0.0)) - if cutlass.const_expr(self.prefetch_group == 1): - for k in cutlass.range_constexpr(self.top_k): - row = cute.arch.load((staged_indices + k).llvm_ptr, Int32) - weight = cute.arch.load( - (staged_weights + k).llvm_ptr, - Float32, - ) - if row != Int32(-1): - source_element = ( - Int64(row) * self.hidden + Int64(fragment) * QUAD_BF16 - ) - source_pointer = cute.make_ptr( - BFloat16, - (routed_output.iterator + source_element).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=8, - ) - accumulator.store( - accumulator.load() - + packed_u32x2_to_bf16x4( - load_global_u32x2(source_pointer) - ).to(Float32) - * weight - ) - else: - inputs = cute.make_rmem_tensor( - cute.make_layout((self.prefetch_group, 2)), - Uint32, - ) - inputs.fill(Uint32(0)) - for group in cutlass.range_constexpr(self.prefetch_groups): - for item in cutlass.range_constexpr(self.prefetch_group): - k = group * self.prefetch_group + item - if cutlass.const_expr(k < self.top_k): - row = cute.arch.load( - (staged_indices + k).llvm_ptr, - Int32, - ) - for word in cutlass.range_constexpr(2): - inputs[item, word] = Uint32(0) - if row != Int32(-1): - source_element = ( - Int64(row) * self.hidden - + Int64(fragment) * QUAD_BF16 - ) - source_pointer = cute.make_ptr( - BFloat16, - (routed_output.iterator + source_element).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=8, - ) - source = load_global_u32x2(source_pointer) - for word in cutlass.range_constexpr(2): - inputs[item, word] = source[word] - for item in cutlass.range_constexpr(self.prefetch_group): - k = group * self.prefetch_group + item - if cutlass.const_expr(k < self.top_k): - source = cute.make_rmem_tensor( - cute.make_layout((2,)), - Uint32, - ) - for word in cutlass.range_constexpr(2): - source[word] = inputs[item, word] - accumulator.store( - accumulator.load() - + packed_u32x2_to_bf16x4(source.load()).to(Float32) - * cute.arch.load( - (staged_weights + k).llvm_ptr, - Float32, - ) - ) - result = accumulator.load() - if cutlass.const_expr(self.include_shared_expert): - result = result + shared_values.load().to(Float32) - result_packed = bf16x4_to_packed_u32x2(result.to(BFloat16)) - packed = sanitize_negative_zero_u32x2(result_packed) - mailbox_element = ( - (Int64(stage) * self.tp + self.rank) * self.capacity_m + Int64(token) - ) * self.hidden + Int64(fragment) * QUAD_BF16 - stmc_bf16x4( - contribution_mailbox_multicast_address + mailbox_element * 2, - packed, - ) - - if block == 0 and tidx == 0: - store_global_u32( - stage_state.iterator + ACTIVE_STAGE, - stage, - ) - cute.arch.barrier() - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_launch_dependents() - - -class _SharedOnlyPublishDeviceKernel: - def __init__( - self, - *, - hidden: int, - tp: int, - rank: int, - capacity_m: int, - elements_per_thread: int, - threads: int, - release_before_store: bool, - enable_pdl: bool, - ) -> None: - if elements_per_thread not in (1, QUAD_BF16, VEC_BF16): - raise ValueError("elements_per_thread must be 1, 4, or 8") - if hidden <= 0 or hidden % elements_per_thread: - raise ValueError("hidden must divide evenly across thread fragments") - self.hidden = hidden - self.tp = tp - self.rank = rank - self.capacity_m = capacity_m - self.elements_per_thread = elements_per_thread - self.threads = threads - self.fragments = hidden // elements_per_thread - self.ctas_per_token = math.ceil(self.fragments / threads) - self.release_before_store = release_before_store - self.enable_pdl = enable_pdl - - @cute.jit - def __call__( - self, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_multicast_address: Int64, - m: Int32, - stream: cuda.CUstream, - ) -> None: - self.kernel( - shared_output, - stage_state, - contribution_mailbox_multicast_address, - ).launch( - grid=(m * self.ctas_per_token, 1, 1), - block=(self.threads, 1, 1), - stream=stream, - use_pdl=self.enable_pdl, - ) - - @cute.kernel - def kernel( - self, - shared_output: cute.Tensor, - stage_state: cute.Tensor, - contribution_mailbox_multicast_address: Int64, - ) -> None: - block, _, _ = cute.arch.block_idx() - tidx, _, _ = cute.arch.thread_idx() - token = block // self.ctas_per_token - cta_in_token = block % self.ctas_per_token - fragment = cta_in_token * self.threads + tidx - - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_wait() - - if cutlass.const_expr(self.elements_per_thread == 1): - bits = Uint32(0) - if fragment < self.fragments: - element = Int64(token) * self.hidden + Int64(fragment) - bits = f32_to_bf16_bits( - load_global_bf16_as_f32( - Int64((shared_output.iterator + element).toint()) - ) - ) - if bits == Uint32(NEGATIVE_ZERO_BF16_BITS): - bits = Uint32(0) - partner_bits = shuffle_sync_idx_u32( - bits, - Int32(cute.arch.lane_idx() | Int32(1)), - ) - packed_word = bits | (partner_bits << Uint32(16)) - elif cutlass.const_expr(self.elements_per_thread == QUAD_BF16): - packed_words = cute.make_rmem_tensor(cute.make_layout((2,)), Uint32) - packed_words.fill(Uint32(0)) - if fragment < self.fragments: - element = Int64(token) * self.hidden + Int64(fragment) * QUAD_BF16 - pointer = cute.make_ptr( - BFloat16, - (shared_output.iterator + element).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=8, - ) - packed_words.store( - sanitize_negative_zero_u32x2(load_global_u32x2(pointer)) - ) - else: - packed_words = cute.make_rmem_tensor(cute.make_layout((4,)), Uint32) - packed_words.fill(Uint32(0)) - if fragment < self.fragments: - element = Int64(token) * self.hidden + Int64(fragment) * VEC_BF16 - pointer = cute.make_ptr( - BFloat16, - (shared_output.iterator + element).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=16, - ) - packed_words.store( - sanitize_negative_zero_u32x4(load_global_u32x4(pointer)) - ) - - stage = load_volatile_u32(stage_state.iterator + NEXT_STAGE) - if cutlass.const_expr(self.release_before_store): - if block == 0 and tidx == 0: - store_global_u32( - stage_state.iterator + ACTIVE_STAGE, - stage, - ) - cute.arch.barrier() - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_launch_dependents() - - if cutlass.const_expr(self.elements_per_thread == 1): - if (cute.arch.lane_idx() & Int32(1)) == Int32( - 0 - ) and fragment < self.fragments: - mailbox_element = ( - (Int64(stage) * self.tp + self.rank) * self.capacity_m - + Int64(token) - ) * self.hidden + Int64(fragment) - stmc_bf16x2( - contribution_mailbox_multicast_address + mailbox_element * 2, - packed_word, - ) - elif cutlass.const_expr(self.elements_per_thread == QUAD_BF16): - if fragment < self.fragments: - mailbox_element = ( - (Int64(stage) * self.tp + self.rank) * self.capacity_m - + Int64(token) - ) * self.hidden + Int64(fragment) * QUAD_BF16 - stmc_bf16x4( - contribution_mailbox_multicast_address + mailbox_element * 2, - packed_words, - ) - else: - if fragment < self.fragments: - mailbox_element = ( - (Int64(stage) * self.tp + self.rank) * self.capacity_m - + Int64(token) - ) * self.hidden + Int64(fragment) * VEC_BF16 - stmc_bf16x8( - contribution_mailbox_multicast_address + mailbox_element * 2, - packed_words, - ) - - if cutlass.const_expr(not self.release_before_store): - if block == 0 and tidx == 0: - store_global_u32( - stage_state.iterator + ACTIVE_STAGE, - stage, - ) - cute.arch.barrier() - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_launch_dependents() - - -class _LamportResidualRMSNormDeviceKernel: - def __init__( - self, - *, - hidden: int, - tp: int, - capacity_m: int, - cluster_size: int, - rank_lanes: int, - threads: int, - rms_epsilon: float, - weight_bias: float, - add_residual: bool, - write_residual_output: bool, - enable_pdl: bool, - ) -> None: - if rank_lanes not in (1, 2, 4, 8): - raise ValueError("rank_lanes must be 1, 2, 4, or 8") - if tp % rank_lanes: - raise ValueError("tp must be divisible by rank_lanes") - if threads <= 0 or threads % WARP_SIZE or threads % rank_lanes: - raise ValueError("threads must be a positive warp and rank-lane multiple") - if hidden <= 0 or hidden % VEC_BF16: - raise ValueError("hidden must be a positive multiple of 8") - self.hidden = hidden - self.tp = tp - self.capacity_m = capacity_m - self.cluster_size = cluster_size - self.rank_lanes = rank_lanes - self.threads = threads - self.rms_epsilon = rms_epsilon - self.weight_bias = weight_bias - self.add_residual = add_residual - self.write_residual_output = write_residual_output - self.enable_pdl = enable_pdl - self.fragments = hidden // VEC_BF16 - self.groups_per_cta = threads // rank_lanes - self.fragment_stride = cluster_size * self.groups_per_cta - self.trips = math.ceil(self.fragments / self.fragment_stride) - self.warps = threads // WARP_SIZE - self.rank_waves = tp // rank_lanes - - def smem_size_in_bytes(self) -> int: - return (self.warps + self.cluster_size) * 4 - - @cute.jit - def __call__( - self, - contribution_mailbox: cute.Tensor, - residual_source: cute.Tensor, - gamma: cute.Tensor, - residual_output: cute.Tensor, - norm_output: cute.Tensor, - stage_state: cute.Tensor, - m: Int32, - stream: cuda.CUstream, - ) -> None: - self.kernel( - contribution_mailbox, - residual_source, - gamma, - residual_output, - norm_output, - stage_state, - ).launch( - grid=(m, self.cluster_size, 1), - block=(self.threads, 1, 1), - cluster=(1, self.cluster_size, 1), - smem=self.smem_size_in_bytes(), - stream=stream, - use_pdl=self.enable_pdl, - ) - - @cute.kernel - def kernel( - self, - contribution_mailbox: cute.Tensor, - residual_source: cute.Tensor, - gamma: cute.Tensor, - residual_output: cute.Tensor, - norm_output: cute.Tensor, - stage_state: cute.Tensor, - ) -> None: - tidx, _, _ = cute.arch.thread_idx() - token, _, _ = cute.arch.block_idx() - cluster_rank = cute.arch.make_warp_uniform(cute.arch.block_idx_in_cluster()) - rank_lane = tidx % self.rank_lanes - group = tidx // self.rank_lanes - base_fragment = cluster_rank * self.groups_per_cta + group - - prenorm_fragments = cute.make_rmem_tensor( - cute.make_layout( - (self.trips, VEC_BF16), - stride=(VEC_BF16, 1), - ), - BFloat16, - ) - prenorm_fragments.fill(BFloat16(0.0)) - gamma_fragments = cute.make_rmem_tensor( - cute.make_layout( - (self.trips, VEC_BF16), - stride=(VEC_BF16, 1), - ), - BFloat16, - ) - gamma_fragments.fill(BFloat16(0.0)) - if cutlass.const_expr(self.add_residual): - residual_fragments = cute.make_rmem_tensor( - cute.make_layout( - (self.trips, VEC_BF16), - stride=(VEC_BF16, 1), - ), - BFloat16, - ) - residual_fragments.fill(BFloat16(0.0)) - - for trip in cutlass.range_constexpr(self.trips): - fragment = base_fragment + trip * self.fragment_stride - if fragment < self.fragments and rank_lane == 0: - gamma_element = Int64(fragment) * VEC_BF16 - gamma_pointer = cute.make_ptr( - BFloat16, - (gamma.iterator + gamma_element).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=16, - ) - gamma_fragments[trip, None].store( - packed_u32x4_to_bf16x8(load_global_u32x4(gamma_pointer)) - ) - if cutlass.const_expr(self.add_residual): - residual_element = Int64(token) * self.hidden + gamma_element - residual_pointer = cute.make_ptr( - BFloat16, - (residual_source.iterator + residual_element).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=16, - ) - residual_fragments[trip, None].store( - packed_u32x4_to_bf16x8(load_global_u32x4(residual_pointer)) - ) - - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_wait() - - active_stage = load_volatile_u32(stage_state.iterator + ACTIVE_STAGE) - for trip in cutlass.range_constexpr(self.trips): - fragment = base_fragment + trip * self.fragment_stride - lane_packed = cute.make_rmem_tensor( - cute.make_layout((self.rank_waves, 4)), - Uint32, - ) - lane_packed.fill(Uint32(0)) - dirty = fragment < self.fragments - while dirty: - dirty = False - for wave in cutlass.range_constexpr(self.rank_waves): - source_rank = wave * self.rank_lanes + rank_lane - if fragment < self.fragments: - source_element = ( - (Int64(active_stage) * self.tp + Int64(source_rank)) - * self.capacity_m - + Int64(token) - ) * self.hidden + Int64(fragment) * VEC_BF16 - source_pointer = cute.make_ptr( - BFloat16, - (contribution_mailbox.iterator + source_element).llvm_ptr, - cute.AddressSpace.gmem, - assumed_align=16, - ) - packed = load_global_u32x4( - source_pointer, - volatile=True, - ) - dirty = dirty | fragment_has_negative_zero(packed) - for word in cutlass.range_constexpr(4): - lane_packed[wave, word] = packed[word] - - lane_sum = cute.make_rmem_tensor(cute.make_layout((VEC_BF16,)), Float32) - lane_sum.fill(Float32(0.0)) - for wave in cutlass.range_constexpr(self.rank_waves): - packed = cute.make_rmem_tensor(cute.make_layout((4,)), Uint32) - for word in cutlass.range_constexpr(4): - packed[word] = lane_packed[wave, word] - lane_sum.store( - lane_sum.load() + packed_u32x4_to_bf16x8(packed.load()).to(Float32) - ) - for offset in cutlass.range_constexpr(1, 5): - if cutlass.const_expr(offset < self.rank_lanes and offset in (1, 2, 4)): - for element in cutlass.range_constexpr(VEC_BF16): - lane_sum[element] = lane_sum[ - element - ] + cute.arch.shuffle_sync_bfly( - lane_sum[element], - offset=offset, - mask=-1, - mask_and_clamp=31, - ) - - if fragment < self.fragments and rank_lane == 0: - prenorm = lane_sum.load() - if cutlass.const_expr(self.add_residual): - prenorm = prenorm + residual_fragments[trip, None].load().to( - Float32 - ) - prenorm_bf16 = prenorm.to(BFloat16) - prenorm_fragments[trip, None].store(prenorm_bf16) - if cutlass.const_expr(self.write_residual_output): - output_element = ( - Int64(token) * self.hidden + Int64(fragment) * VEC_BF16 - ) - store_global_u32x4( - Int64((residual_output.iterator + output_element).toint()), - bf16x8_to_packed_u32x4(prenorm_bf16), - ) - - if token == 0 and cluster_rank == 0 and tidx == 0: - store_global_u32( - stage_state.iterator + NEXT_STAGE, - (active_stage + Uint32(1)) % Uint32(LAMPORT_GENERATIONS), - ) - if cutlass.const_expr(self.enable_pdl): - cute.arch.griddepcontrol_launch_dependents() - - for trip in cutlass.range_constexpr(self.trips): - fragment = base_fragment + trip * self.fragment_stride - for wave in cutlass.range_constexpr(self.rank_waves): - source_rank = wave * self.rank_lanes + rank_lane - if fragment < self.fragments: - source_element = ( - (Int64(active_stage) * self.tp + Int64(source_rank)) - * self.capacity_m - + Int64(token) - ) * self.hidden + Int64(fragment) * VEC_BF16 - store_lamport_sentinel_u32x4( - Int64((contribution_mailbox.iterator + source_element).toint()) - ) - - thread_sum = Float32(0.0) - for trip in cutlass.range_constexpr(self.trips): - fragment = base_fragment + trip * self.fragment_stride - if fragment < self.fragments and rank_lane == 0: - values = prenorm_fragments[trip, None].load().to(Float32) - thread_sum = thread_sum + (values * values).reduce( - cute.ReductionOp.ADD, - init_val=Float32(0.0), - reduction_profile=0, - ) - - smem = cutlass.utils.SmemAllocator() - warp_sums = smem.allocate_array(Float32, self.warps) - cluster_sums = smem.allocate_array(Float32, self.cluster_size) - cta_sum = _group_leader_block_sum( - thread_sum, - warp_sums, - self.warps, - self.rank_lanes, - ) - if tidx < self.cluster_size: - local_slot = cluster_sums + cluster_rank - remote_slot = map_shared_to_peer(local_slot, Int32(tidx)) - store_shared_cluster_f32(remote_slot, cta_sum) - cute.arch.cluster_arrive() - cute.arch.cluster_wait() - - full_sum = Float32(0.0) - for peer in cutlass.range_constexpr(self.cluster_size): - full_sum = full_sum + cute.arch.load( - (cluster_sums + peer).llvm_ptr, - Float32, - ) - inv_rms = cute.math.rsqrt( - full_sum / Float32(self.hidden) + Float32(self.rms_epsilon), - fastmath=True, - ) - for trip in cutlass.range_constexpr(self.trips): - fragment = base_fragment + trip * self.fragment_stride - if fragment < self.fragments and rank_lane == 0: - gamma_values = gamma_fragments[trip, None].load().to(Float32) - if cutlass.const_expr(self.weight_bias != 0.0): - gamma_values = gamma_values + Float32(self.weight_bias) - result = ( - prenorm_fragments[trip, None].load().to(Float32) - * inv_rms - * gamma_values - ).to(BFloat16) - output_element = Int64(token) * self.hidden + Int64(fragment) * VEC_BF16 - store_global_u32x4( - Int64((norm_output.iterator + output_element).toint()), - bf16x8_to_packed_u32x4(result), - ) diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ll/protocol.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ll/protocol.py deleted file mode 100644 index 3a26bb485..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/kernel_ll/protocol.py +++ /dev/null @@ -1,461 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Low-latency MNNVL protocol and its two operation paths.""" - -from __future__ import annotations - -from dataclasses import dataclass -from typing import Any, TypedDict, cast - -import cutlass.cute as cute -import torch -import torch.distributed as dist -from cutlass import BFloat16, Int32, Int64 -from cutlass.cute.runtime import make_fake_compact_tensor - -from ..cute_dsl_primitives import QUAD_BF16 -from ..runtime import ( - current_cu_stream, - make_fake_dynamic_compact_tensor, - to_cute, - to_cute_dynamic, -) -from ..symmetric_buffer import SymmetricBuffer -from .device_kernels import ( - LAMPORT_GENERATIONS, - _LamportResidualRMSNormDeviceKernel, - _QuadFinalizePublishDeviceKernel, - _ScalarFinalizePublishDeviceKernel, - _SharedOnlyPublishDeviceKernel, -) - - -@dataclass(frozen=True, slots=True) -class LLCollectiveTuning: - cluster_size: int = 8 - rank_lanes: int = 1 - threads: int = 128 - enable_pdl: bool = True - - -@dataclass(frozen=True, slots=True) -class LLFinalizeTuning: - elements_per_thread: int = 4 - threads: int = 128 - prefetch_group: int = 10 - load_shared_expert_before_pdl: bool = False - collective: LLCollectiveTuning = LLCollectiveTuning() - - -@dataclass(frozen=True, slots=True) -class LLAllReduceTuning: - publish_elements_per_thread: int = 8 - publish_threads: int = 128 - publish_release_before_store: bool = False - collective: LLCollectiveTuning = LLCollectiveTuning() - - -LL_FINALIZE_GB300_TP8_H8192_K10 = LLFinalizeTuning() -LL_FINALIZE_GB300_TP8_H8192_K10_M_GE_20 = LLFinalizeTuning( - collective=LLCollectiveTuning(cluster_size=16, rank_lanes=2) -) -LL_FINALIZE_GB300_TP16_H8192_K10 = LLFinalizeTuning( - collective=LLCollectiveTuning(cluster_size=16, rank_lanes=2) -) -LL_ALL_REDUCE_GB300_TP8_H8192_M_LE_4 = LLAllReduceTuning( - collective=LLCollectiveTuning(cluster_size=16, threads=64) -) -LL_ALL_REDUCE_GB300_TP8_H8192_M_GE_5 = LLAllReduceTuning() -LL_ALL_REDUCE_GB300_TP16_H8192_M_LE_10 = LLAllReduceTuning( - collective=LLCollectiveTuning(cluster_size=16, threads=64) -) -LL_ALL_REDUCE_GB300_TP16_H8192_M_11_TO_17 = LLAllReduceTuning() -LL_ALL_REDUCE_GB300_TP16_H8192_M_GE_18 = LLAllReduceTuning( - collective=LLCollectiveTuning(cluster_size=16, rank_lanes=2) -) -LL_ALL_REDUCE_GB300_TP8_H8192 = LL_ALL_REDUCE_GB300_TP8_H8192_M_GE_5 -LL_ALL_REDUCE_GB300_TP16_H8192 = LL_ALL_REDUCE_GB300_TP16_H8192_M_LE_10 - - -@dataclass(slots=True) -class LLProtocolState: - contribution_mailbox: SymmetricBuffer - stage_state: torch.Tensor - - -@dataclass(frozen=True, slots=True) -class _CompiledFinalize: - publish: Any - collective: Any - - -@dataclass(frozen=True, slots=True) -class _CompiledAllReduce: - publish: Any - collective: Any - - -class _PathKwargs(TypedDict): - hidden_size: int - top_k: int - capacity_m: int - write_residual_output: bool - - -class _LLPath: - def __init__( - self, - *, - hidden_size: int, - top_k: int, - capacity_m: int, - write_residual_output: bool, - ) -> None: - self.hidden_size = hidden_size - self.top_k = top_k - self.capacity_m = capacity_m - self.write_residual_output = write_residual_output - - def _outputs( - self, - m: int, - norm_output: torch.Tensor | None, - residual_output: torch.Tensor | None, - ) -> tuple[torch.Tensor, torch.Tensor | None]: - shape = (m, self.hidden_size) - device = torch.device("cuda", torch.cuda.current_device()) - if norm_output is None: - norm_output = torch.empty(shape, dtype=torch.bfloat16, device=device) - if self.write_residual_output and residual_output is None: - residual_output = torch.empty(shape, dtype=torch.bfloat16, device=device) - return norm_output, residual_output - - def _validate_state(self, state: LLProtocolState, m: int) -> None: - if not 1 <= m <= self.capacity_m: - raise ValueError(f"m must be in [1, {self.capacity_m}]") - address = state.contribution_mailbox.multicast_address - if address is None or address % 16: - raise ValueError( - "LL contribution mailbox requires a 16-byte-aligned multicast address" - ) - - def _launch_collective( - self, - collective, - residual_source: torch.Tensor | None, - gamma: torch.Tensor, - state: LLProtocolState, - norm_output: torch.Tensor, - residual_output: torch.Tensor | None, - m: int, - ) -> None: - residual_arg = residual_source if residual_source is not None else norm_output - residual_output_arg = ( - residual_output if residual_output is not None else norm_output - ) - collective( - to_cute(state.contribution_mailbox.tensor.flatten(), 16), - to_cute_dynamic(residual_arg.flatten(), 16, divisibility=self.hidden_size), - to_cute(gamma, 16), - to_cute_dynamic( - residual_output_arg.flatten(), - 16, - divisibility=self.hidden_size, - ), - to_cute_dynamic(norm_output.flatten(), 16, divisibility=self.hidden_size), - to_cute(state.stage_state, 4), - Int32(m), - current_cu_stream(), - ) - - -class FinalizeAllReduceRMSNormLLKernel(_LLPath): - def __init__(self, *, compiled: _CompiledFinalize, **kwargs) -> None: - super().__init__(**kwargs) - self._compiled = compiled - - def __call__( - self, - routed_output: torch.Tensor, - expert_weights: torch.Tensor, - permuted_indices: torch.Tensor, - shared_output: torch.Tensor | None, - residual_source: torch.Tensor | None, - gamma: torch.Tensor, - m: int, - *, - state: LLProtocolState, - norm_output: torch.Tensor | None = None, - residual_output: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor | None]: - self._validate_state(state, m) - norm_output, residual_output = self._outputs(m, norm_output, residual_output) - shared_arg = shared_output if shared_output is not None else norm_output - self._compiled.publish( - to_cute_dynamic(routed_output.flatten(), 16, divisibility=self.hidden_size), - to_cute_dynamic(expert_weights.flatten(), 2, divisibility=self.top_k), - to_cute_dynamic(permuted_indices.flatten(), 4, divisibility=self.top_k), - to_cute_dynamic(shared_arg.flatten(), 16, divisibility=self.hidden_size), - to_cute(state.stage_state, 4), - Int64(cast(int, state.contribution_mailbox.multicast_address)), - Int32(m), - current_cu_stream(), - ) - self._launch_collective( - self._compiled.collective, - residual_source, - gamma, - state, - norm_output, - residual_output, - m, - ) - return norm_output, residual_output - - -class AllReduceRMSNormLLKernel(_LLPath): - def __init__(self, *, compiled: _CompiledAllReduce, **kwargs) -> None: - super().__init__(**kwargs) - self._compiled = compiled - - def __call__( - self, - local_contribution: torch.Tensor, - residual_source: torch.Tensor | None, - gamma: torch.Tensor, - m: int, - *, - state: LLProtocolState, - norm_output: torch.Tensor | None = None, - residual_output: torch.Tensor | None = None, - ) -> tuple[torch.Tensor, torch.Tensor | None]: - self._validate_state(state, m) - norm_output, residual_output = self._outputs(m, norm_output, residual_output) - self._compiled.publish( - to_cute_dynamic( - local_contribution.flatten(), - 16, - divisibility=self.hidden_size, - ), - to_cute(state.stage_state, 4), - Int64(cast(int, state.contribution_mailbox.multicast_address)), - Int32(m), - current_cu_stream(), - ) - self._launch_collective( - self._compiled.collective, - residual_source, - gamma, - state, - norm_output, - residual_output, - m, - ) - return norm_output, residual_output - - -class LLProtocol: - """Own LL State and protocol-local compiled variants for both paths.""" - - def __init__( - self, - hidden_size: int, - top_k: int, - tp_size: int, - rank: int, - capacity_m: int, - rms_epsilon: float, - routed_scaling_factor: float, - weight_bias: float, - *, - include_shared_expert: bool, - add_residual: bool, - write_residual_output: bool, - finalize_tunings: tuple[LLFinalizeTuning, ...], - all_reduce_tunings: tuple[LLAllReduceTuning, ...], - group: dist.ProcessGroup, - ) -> None: - self.hidden_size = hidden_size - self.top_k = top_k - self.tp_size = tp_size - self.rank = rank - self.capacity_m = capacity_m - self.rms_epsilon = rms_epsilon - self.routed_scaling_factor = routed_scaling_factor - self.weight_bias = weight_bias - self.include_shared_expert = include_shared_expert - self.add_residual = add_residual - self.write_residual_output = write_residual_output - - collective_cache = { - tuning: self._compile_collective(tuning) - for tuning in { - *(item.collective for item in finalize_tunings), - *(item.collective for item in all_reduce_tunings), - } - } - self.finalize_kernels = { - tuning: FinalizeAllReduceRMSNormLLKernel( - compiled=_CompiledFinalize( - publish=self._compile_finalize(tuning), - collective=collective_cache[tuning.collective], - ), - **self._path_kwargs(), - ) - for tuning in dict.fromkeys(finalize_tunings) - } - self.all_reduce_kernels = { - tuning: AllReduceRMSNormLLKernel( - compiled=_CompiledAllReduce( - publish=self._compile_all_reduce_publish(tuning), - collective=collective_cache[tuning.collective], - ), - **self._path_kwargs(), - ) - for tuning in dict.fromkeys(all_reduce_tunings) - } - self.state = self._create_state(group) - - def _path_kwargs(self) -> _PathKwargs: - return { - "hidden_size": self.hidden_size, - "top_k": self.top_k, - "capacity_m": self.capacity_m, - "write_residual_output": self.write_residual_output, - } - - def _compile_finalize(self, tuning: LLFinalizeTuning): - if tuning.elements_per_thread not in (1, QUAD_BF16): - raise ValueError("LL finalize elements_per_thread must be 1 or 4") - kwargs: dict[str, Any] = { - "hidden": self.hidden_size, - "top_k": self.top_k, - "tp": self.tp_size, - "rank": self.rank, - "capacity_m": self.capacity_m, - "threads": tuning.threads, - "routed_scaling_factor": self.routed_scaling_factor, - "include_shared_expert": self.include_shared_expert, - "load_shared_expert_before_pdl": tuning.load_shared_expert_before_pdl, - "enable_pdl": tuning.collective.enable_pdl, - "prefetch_group": tuning.prefetch_group, - } - device_kernel = ( - _ScalarFinalizePublishDeviceKernel(**kwargs) - if tuning.elements_per_thread == 1 - else _QuadFinalizePublishDeviceKernel(**kwargs) - ) - args = ( - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_dynamic_compact_tensor( - BFloat16, alignment=2, divisibility=self.top_k - ), - make_fake_dynamic_compact_tensor( - Int32, alignment=4, divisibility=self.top_k - ), - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_compact_tensor(Int32, (2,), assumed_align=4), - Int64(0), - Int32(self.capacity_m), - current_cu_stream(), - ) - return cute.compile(device_kernel, *args) - - def _compile_all_reduce_publish(self, tuning: LLAllReduceTuning): - device_kernel = _SharedOnlyPublishDeviceKernel( - hidden=self.hidden_size, - tp=self.tp_size, - rank=self.rank, - capacity_m=self.capacity_m, - elements_per_thread=tuning.publish_elements_per_thread, - threads=tuning.publish_threads, - release_before_store=tuning.publish_release_before_store, - enable_pdl=tuning.collective.enable_pdl, - ) - args = ( - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_compact_tensor(Int32, (2,), assumed_align=4), - Int64(0), - Int32(self.capacity_m), - current_cu_stream(), - ) - return cute.compile(device_kernel, *args) - - def _compile_collective(self, tuning: LLCollectiveTuning): - device_kernel = _LamportResidualRMSNormDeviceKernel( - hidden=self.hidden_size, - tp=self.tp_size, - capacity_m=self.capacity_m, - cluster_size=tuning.cluster_size, - rank_lanes=tuning.rank_lanes, - threads=tuning.threads, - rms_epsilon=self.rms_epsilon, - weight_bias=self.weight_bias, - add_residual=self.add_residual, - write_residual_output=self.write_residual_output, - enable_pdl=tuning.enable_pdl, - ) - activation = self.capacity_m * self.hidden_size - args = ( - make_fake_compact_tensor( - BFloat16, - (LAMPORT_GENERATIONS * self.tp_size * activation,), - assumed_align=16, - ), - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_compact_tensor(BFloat16, (self.hidden_size,), assumed_align=16), - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_dynamic_compact_tensor( - BFloat16, alignment=16, divisibility=self.hidden_size - ), - make_fake_compact_tensor(Int32, (2,), assumed_align=4), - Int32(self.capacity_m), - current_cu_stream(), - ) - return cute.compile(device_kernel, *args) - - def _create_state(self, group: dist.ProcessGroup) -> LLProtocolState: - if dist.get_world_size(group) != self.tp_size: - raise ValueError("ProcessGroup size does not match tp_size") - if dist.get_rank(group) != self.rank: - raise ValueError("ProcessGroup rank does not match rank") - device = torch.device("cuda", torch.cuda.current_device()) - mailbox = SymmetricBuffer.allocate( - ( - LAMPORT_GENERATIONS, - self.tp_size, - self.capacity_m, - self.hidden_size, - ), - torch.bfloat16, - device, - group, - require_multicast=True, - ) - mailbox.tensor.view(torch.int16).fill_(-32768) - return LLProtocolState( - contribution_mailbox=mailbox, - stage_state=torch.zeros((2,), dtype=torch.int32, device=device), - ) diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/presets.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/presets.py deleted file mode 100644 index 8630a79eb..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/presets.py +++ /dev/null @@ -1,338 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Built-in routing configurations for the MNNVL CuTe DSL backend.""" - -import torch - -from .config import ( - KernelTarget, - MNNVLCuteDSLConfig, - MRangeDispatch, - ProtocolKind, - StaticProfile, -) -from .kernel_bt import ( - BT_ALL_REDUCE_GB300_TP8_H8192_PRESET_0, - BT_ALL_REDUCE_GB300_TP8_H8192_PRESET_1, - BT_ALL_REDUCE_GB300_TP16_H8192_PRESET_0, - BT_ALL_REDUCE_GB300_TP16_H8192_PRESET_1, - BT_FINALIZE_GB300_TP8_H8192_K10_PRESET_0, - BT_FINALIZE_GB300_TP8_H8192_K10_PRESET_1, - BT_FINALIZE_GB300_TP16_H8192_K10_PRESET_0, - BT_FINALIZE_GB300_TP16_H8192_K10_PRESET_1, -) -from .kernel_ht import ( - HT_ALL_REDUCE_GB300_TP8_H8192, - HT_ALL_REDUCE_GB300_TP16_H8192, - HT_FINALIZE_GB300_TP8_H8192_K10, - HT_FINALIZE_GB300_TP16_H8192_K10, -) -from .kernel_ll import ( - LL_ALL_REDUCE_GB300_TP8_H8192, - LL_ALL_REDUCE_GB300_TP16_H8192, - LL_FINALIZE_GB300_TP8_H8192_K10, - LL_FINALIZE_GB300_TP16_H8192_K10, -) - -__all__ = [ - "BT_ONLY_CONFIG", - "DEFAULT_CONFIG", - "HT_ONLY_CONFIG", - "LL_ONLY_CONFIG", -] - - -def _target(protocol: ProtocolKind, preset: object) -> KernelTarget[object]: - return KernelTarget(protocol=protocol, preset=preset) - - -LL_ONLY_CONFIG = MNNVLCuteDSLConfig( - profiles=( - StaticProfile( - tp_size=8, - hidden_size=8192, - top_k=10, - dtype=torch.bfloat16, - finalize_routes=MRangeDispatch( - upper_bounds=(None,), - targets=( - _target( - ProtocolKind.LL, - LL_FINALIZE_GB300_TP8_H8192_K10, - ), - ), - ), - all_reduce_routes=MRangeDispatch( - upper_bounds=(None,), - targets=( - _target( - ProtocolKind.LL, - LL_ALL_REDUCE_GB300_TP8_H8192, - ), - ), - ), - ), - StaticProfile( - tp_size=16, - hidden_size=8192, - top_k=10, - dtype=torch.bfloat16, - finalize_routes=MRangeDispatch( - upper_bounds=(None,), - targets=( - _target( - ProtocolKind.LL, - LL_FINALIZE_GB300_TP16_H8192_K10, - ), - ), - ), - all_reduce_routes=MRangeDispatch( - upper_bounds=(None,), - targets=( - _target( - ProtocolKind.LL, - LL_ALL_REDUCE_GB300_TP16_H8192, - ), - ), - ), - ), - ) -) - - -BT_ONLY_CONFIG = MNNVLCuteDSLConfig( - profiles=( - StaticProfile( - tp_size=8, - hidden_size=8192, - top_k=10, - dtype=torch.bfloat16, - finalize_routes=MRangeDispatch( - upper_bounds=(48, 1024), - targets=( - _target( - ProtocolKind.BT, - BT_FINALIZE_GB300_TP8_H8192_K10_PRESET_0, - ), - _target( - ProtocolKind.BT, - BT_FINALIZE_GB300_TP8_H8192_K10_PRESET_1, - ), - ), - ), - all_reduce_routes=MRangeDispatch( - upper_bounds=(256, 1024), - targets=( - _target( - ProtocolKind.BT, - BT_ALL_REDUCE_GB300_TP8_H8192_PRESET_0, - ), - _target( - ProtocolKind.BT, - BT_ALL_REDUCE_GB300_TP8_H8192_PRESET_1, - ), - ), - ), - ), - StaticProfile( - tp_size=16, - hidden_size=8192, - top_k=10, - dtype=torch.bfloat16, - finalize_routes=MRangeDispatch( - upper_bounds=(52, 1024), - targets=( - _target( - ProtocolKind.BT, - BT_FINALIZE_GB300_TP16_H8192_K10_PRESET_0, - ), - _target( - ProtocolKind.BT, - BT_FINALIZE_GB300_TP16_H8192_K10_PRESET_1, - ), - ), - ), - all_reduce_routes=MRangeDispatch( - upper_bounds=(512, 1024), - targets=( - _target( - ProtocolKind.BT, - BT_ALL_REDUCE_GB300_TP16_H8192_PRESET_0, - ), - _target( - ProtocolKind.BT, - BT_ALL_REDUCE_GB300_TP16_H8192_PRESET_1, - ), - ), - ), - ), - ) -) - - -HT_ONLY_CONFIG = MNNVLCuteDSLConfig( - profiles=( - StaticProfile( - tp_size=8, - hidden_size=8192, - top_k=10, - dtype=torch.bfloat16, - finalize_routes=MRangeDispatch( - upper_bounds=(None,), - targets=( - _target( - ProtocolKind.HT, - HT_FINALIZE_GB300_TP8_H8192_K10, - ), - ), - ), - all_reduce_routes=MRangeDispatch( - upper_bounds=(None,), - targets=( - _target( - ProtocolKind.HT, - HT_ALL_REDUCE_GB300_TP8_H8192, - ), - ), - ), - ), - StaticProfile( - tp_size=16, - hidden_size=8192, - top_k=10, - dtype=torch.bfloat16, - finalize_routes=MRangeDispatch( - upper_bounds=(None,), - targets=( - _target( - ProtocolKind.HT, - HT_FINALIZE_GB300_TP16_H8192_K10, - ), - ), - ), - all_reduce_routes=MRangeDispatch( - upper_bounds=(None,), - targets=( - _target( - ProtocolKind.HT, - HT_ALL_REDUCE_GB300_TP16_H8192, - ), - ), - ), - ), - ) -) - - -DEFAULT_CONFIG = MNNVLCuteDSLConfig( - profiles=( - StaticProfile( - tp_size=8, - hidden_size=8192, - top_k=10, - dtype=torch.bfloat16, - finalize_routes=MRangeDispatch( - upper_bounds=(23, 48, 703, None), - targets=( - _target( - ProtocolKind.LL, - LL_FINALIZE_GB300_TP8_H8192_K10, - ), - _target( - ProtocolKind.BT, - BT_FINALIZE_GB300_TP8_H8192_K10_PRESET_0, - ), - _target( - ProtocolKind.BT, - BT_FINALIZE_GB300_TP8_H8192_K10_PRESET_1, - ), - _target( - ProtocolKind.HT, - HT_FINALIZE_GB300_TP8_H8192_K10, - ), - ), - ), - all_reduce_routes=MRangeDispatch( - upper_bounds=(15, 256, 1024, None), - targets=( - _target( - ProtocolKind.LL, - LL_ALL_REDUCE_GB300_TP8_H8192, - ), - _target( - ProtocolKind.BT, - BT_ALL_REDUCE_GB300_TP8_H8192_PRESET_0, - ), - _target( - ProtocolKind.BT, - BT_ALL_REDUCE_GB300_TP8_H8192_PRESET_1, - ), - _target( - ProtocolKind.HT, - HT_ALL_REDUCE_GB300_TP8_H8192, - ), - ), - ), - ), - StaticProfile( - tp_size=16, - hidden_size=8192, - top_k=10, - dtype=torch.bfloat16, - finalize_routes=MRangeDispatch( - upper_bounds=(7, 52, 703, None), - targets=( - _target( - ProtocolKind.LL, - LL_FINALIZE_GB300_TP16_H8192_K10, - ), - _target( - ProtocolKind.BT, - BT_FINALIZE_GB300_TP16_H8192_K10_PRESET_0, - ), - _target( - ProtocolKind.BT, - BT_FINALIZE_GB300_TP16_H8192_K10_PRESET_1, - ), - _target( - ProtocolKind.HT, - HT_FINALIZE_GB300_TP16_H8192_K10, - ), - ), - ), - all_reduce_routes=MRangeDispatch( - upper_bounds=(5, 512, 959, None), - targets=( - _target( - ProtocolKind.LL, - LL_ALL_REDUCE_GB300_TP16_H8192, - ), - _target( - ProtocolKind.BT, - BT_ALL_REDUCE_GB300_TP16_H8192_PRESET_0, - ), - _target( - ProtocolKind.BT, - BT_ALL_REDUCE_GB300_TP16_H8192_PRESET_1, - ), - _target( - ProtocolKind.HT, - HT_ALL_REDUCE_GB300_TP16_H8192, - ), - ), - ), - ), - ) -) diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/runtime.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/runtime.py deleted file mode 100644 index a706f82ce..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/runtime.py +++ /dev/null @@ -1,70 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Framework-side facilities shared by production Kernel wrappers.""" - -import cuda.bindings.driver as cuda -import cutlass.cute as cute -import torch -from cutlass.cute.runtime import from_dlpack, make_fake_compact_tensor - - -class _GraphSafeDLPack: - __slots__ = ("tensor",) - - def __init__(self, tensor: torch.Tensor) -> None: - self.tensor = tensor - - def __dlpack__(self, stream=None): - # stream=-1 skips producer sync; CuTe launches on the current captured stream. - return self.tensor.__dlpack__(stream=-1) - - def __dlpack_device__(self): - return self.tensor.__dlpack_device__() - - -def to_cute(tensor: torch.Tensor, alignment: int) -> cute.Tensor: - return from_dlpack( - _GraphSafeDLPack(tensor.detach()), - assumed_align=alignment, - ) - - -def to_cute_dynamic( - tensor: torch.Tensor, - alignment: int, - *, - divisibility: int, -) -> cute.Tensor: - return to_cute(tensor, alignment).mark_compact_shape_dynamic( - mode=0, - divisibility=divisibility, - ) - - -def make_fake_dynamic_compact_tensor( - dtype, - *, - alignment: int, - divisibility: int, -) -> cute.Tensor: - return make_fake_compact_tensor( - dtype, - (cute.sym_int32(divisibility=divisibility),), - assumed_align=alignment, - ) - - -def current_cu_stream() -> cuda.CUstream: - return cuda.CUstream(torch.cuda.current_stream().cuda_stream) diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/symmetric_buffer.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl/symmetric_buffer.py deleted file mode 100644 index e4c4fd2a6..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl/symmetric_buffer.py +++ /dev/null @@ -1,107 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Typed ownership for one rendezvoused symmetric Tensor.""" - -from __future__ import annotations - -from collections.abc import Sequence -from dataclasses import dataclass, field - -import torch -import torch.distributed as dist -import torch.distributed._symmetric_memory as symm_mem - -# Upstream reads this helper through a package-relative import; every -# FlashInfer this tree supports ships it, so only the spelling differs here. -from flashinfer.comm.torch_symmetric_memory import _enable_symm_mem_for_group - - -@dataclass(frozen=True, slots=True) -class SymmetricBuffer: - """A symmetric Tensor and the mapping resources derived at rendezvous.""" - - tensor: torch.Tensor - # Keep the rendezvous mapping alive without exposing the backend handle as - # part of a Kernel State's public surface. - _handle: object = field(repr=False) - multicast_address: int | None = field(default=None, repr=False) - peer_addresses: torch.Tensor | None = field(default=None, repr=False) - - @classmethod - def allocate( - cls, - shape: Sequence[int], - dtype: torch.dtype, - device: torch.device, - group: dist.ProcessGroup, - *, - require_multicast: bool = False, - materialize_peer_addresses: bool = False, - ) -> SymmetricBuffer: - """Allocate with the current SymmMem backend and verify requested mappings.""" - if symm_mem.get_backend(device) is None: - raise RuntimeError( - "PyTorch Symmetric Memory has no backend for the current device" - ) - _enable_symm_mem_for_group(group.group_name) - return cls.rendezvous( - symm_mem.empty(shape, dtype=dtype, device=device), - group, - require_multicast=require_multicast, - materialize_peer_addresses=materialize_peer_addresses, - ) - - @classmethod - def rendezvous( - cls, - tensor: torch.Tensor, - group: dist.ProcessGroup, - *, - require_multicast: bool = False, - materialize_peer_addresses: bool = False, - ) -> SymmetricBuffer: - _enable_symm_mem_for_group(group.group_name) - handle = symm_mem.rendezvous(tensor, group) - multicast_address = None - if require_multicast: - multicast_address = int(handle.multicast_ptr or 0) - if not multicast_address: - raise RuntimeError("NVLink multicast mapping is unavailable") - - peer_addresses = None - if materialize_peer_addresses: - # Preserve the rendezvous offset for SymmMem Pool suballocations. - addresses = [ - handle.get_remote_tensor( - peer, - tensor.shape, - tensor.dtype, - ).data_ptr() - for peer in range(dist.get_world_size(group)) - ] - if any(not address for address in addresses): - raise RuntimeError("Symmetric peer mapping is unavailable") - peer_addresses = torch.tensor( - addresses, - dtype=torch.int64, - device=tensor.device, - ) - - return cls( - tensor=tensor, - _handle=handle, - multicast_address=multicast_address, - peer_addresses=peer_addresses, - ) diff --git a/python/sglang/kernels/ops/communication/mnnvl_cutedsl_ar.py b/python/sglang/kernels/ops/communication/mnnvl_cutedsl_ar.py deleted file mode 100644 index b07f2531a..000000000 --- a/python/sglang/kernels/ops/communication/mnnvl_cutedsl_ar.py +++ /dev/null @@ -1,584 +0,0 @@ -# Copyright (c) 2026 by FlashInfer team. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""MNNVL AllReduce fusion backend implemented with CuTe DSL. - -Ported from flashinfer-ai/flashinfer#4358 at main commit 906181e (with -sibling package mnnvl_cutedsl/), pending an installable FlashInfer release -that ships flashinfer.comm.mnnvl_cutedsl. The base communication -infrastructure (mnnvl probing, pattern enum, workspace ABC) still comes -from the installed FlashInfer; the pinned release provides it. -""" - -from __future__ import annotations - -import logging -from typing import Any, Optional, cast - -import torch -import torch.distributed as dist -import torch.distributed._symmetric_memory as symm_mem - -# Keep the copied backend and kernel package self-contained while reusing the -# stable communication infrastructure already supplied by the serving image. -from flashinfer.comm.mnnvl import is_multicast_supported -from flashinfer.comm.trtllm_ar import AllReduceFusionPattern -from flashinfer.comm.workspace_base import AllReduceFusionWorkspace -from torch.distributed import ProcessGroup - -from .mnnvl_cutedsl import DEFAULT_CONFIG, MNNVLCuteDSLConfig, ProtocolKind -from .mnnvl_cutedsl.config import StaticProfile -from .mnnvl_cutedsl.kernel_bt import BTAllReduceTuning, BTFinalizeTuning -from .mnnvl_cutedsl.kernel_bt.protocol import BTProtocol -from .mnnvl_cutedsl.kernel_ht import HTAllReduceTuning, HTFinalizeTuning -from .mnnvl_cutedsl.kernel_ht.protocol import HTProtocol -from .mnnvl_cutedsl.kernel_ll import LLAllReduceTuning, LLFinalizeTuning -from .mnnvl_cutedsl.kernel_ll.protocol import LLProtocol - -logger = logging.getLogger(__name__) - -__all__ = [ - "MNNVLCuteDSLAllReduceFusionWorkspace", - "mnnvl_cutedsl_allreduce_fusion", -] - - -def _check_tensor( - tensor: torch.Tensor, - name: str, - *, - shape: tuple[int | None, ...], - dtype: torch.dtype, - device: torch.device, - alignment: int, -) -> None: - if tensor.device != device: - raise ValueError(f"{name} must be on {device}") - if tensor.dtype != dtype: - raise ValueError(f"{name} must have dtype {dtype}") - if tensor.ndim != len(shape) or any( - expected is not None and actual != expected - for actual, expected in zip(tensor.shape, shape, strict=True) - ): - raise ValueError(f"{name} has an unsupported shape") - if not tensor.is_contiguous(): - raise ValueError(f"{name} must be contiguous") - if tensor.data_ptr() % alignment: - raise ValueError(f"{name} must be {alignment}-byte aligned") - - -def _warn_pdl_mismatch( - workspace: MNNVLCuteDSLAllReduceFusionWorkspace, - pattern: int, - m: int, - launch_with_pdl: bool, -) -> None: - preset_pdl = workspace._uses_pdl(pattern, m) - if launch_with_pdl != preset_pdl: - logger.warning( - "launch_with_pdl does not match the selected MNNVL CuTe DSL " - "preset; using enable_pdl=%s", - preset_pdl, - ) - - -class MNNVLCuteDSLAllReduceFusionWorkspace(AllReduceFusionWorkspace): - """Compiled LL, BT, and HT protocols for one static problem shape. - - Workspace construction compiles the selected kernels and must finish before - the first invocation. Calls using the same workspace must not overlap. - Feature-disabled tensor slots use internal placeholders that are not read. - """ - - _destroyed: bool - - def __init__( - self, - tp_size: int, - tp_rank: int, - max_token_num: int, - hidden_dim: int, - dtype: torch.dtype, - *, - group: Optional[ProcessGroup] = None, - top_k: int = 10, - rms_eps: float = 1e-6, - routed_scaling_factor: float = 1.0, - weight_bias: float = 0.0, - include_shared_expert: bool = True, - add_residual: bool = True, - write_residual_output: bool = True, - config: MNNVLCuteDSLConfig = DEFAULT_CONFIG, - ) -> None: - if tp_size not in (2, 4, 8, 16): - raise ValueError("tp_size must be 2, 4, 8, or 16") - if not 0 <= tp_rank < tp_size: - raise ValueError("tp_rank must be in [0, tp_size)") - if max_token_num <= 0: - raise ValueError("max_token_num must be positive") - if dtype != torch.bfloat16: - raise ValueError("MNNVL CuTe DSL kernels only support torch.bfloat16") - if not torch.cuda.is_available(): - raise RuntimeError("MNNVL CuTe DSL kernels require CUDA") - device = torch.device("cuda", torch.cuda.current_device()) - if torch.cuda.get_device_capability(device)[0] < 10: - raise RuntimeError("MNNVL CuTe DSL kernels require a Blackwell GPU") - if symm_mem.get_backend(device) is None: - raise RuntimeError("PyTorch Symmetric Memory is unavailable") - if not is_multicast_supported(device.index): - raise RuntimeError("NVLink multicast is unavailable") - if group is None: - if not dist.is_initialized(): - raise ValueError("A ProcessGroup is required before initialization") - group = dist.group.WORLD - if dist.get_world_size(group) != tp_size: - raise ValueError("ProcessGroup size does not match tp_size") - if dist.get_rank(group) != tp_rank: - raise ValueError("ProcessGroup rank does not match tp_rank") - - super().__init__(tp_size, tp_rank) - self._protocols: dict[ProtocolKind, LLProtocol | BTProtocol | HTProtocol] = {} - self.max_token_num = max_token_num - self.hidden_dim = hidden_dim - self.top_k = top_k - self.dtype = dtype - self.group = group - self.rms_eps = rms_eps - self.routed_scaling_factor = routed_scaling_factor - self.weight_bias = weight_bias - self.include_shared_expert = include_shared_expert - self.add_residual = add_residual - self.write_residual_output = write_residual_output - self.config = config - self.profile = config.resolve( - tp_size=tp_size, - hidden_size=hidden_dim, - top_k=top_k, - dtype=dtype, - capacity_m=max_token_num, - ) - - for protocol in (ProtocolKind.LL, ProtocolKind.BT, ProtocolKind.HT): - capacity = self.profile.protocol_capacity( - protocol, capacity_m=max_token_num - ) - if capacity is None: - continue - finalize_tunings = self._tunings( - self.profile, protocol, finalize=True, capacity_m=capacity - ) - all_reduce_tunings = self._tunings( - self.profile, protocol, finalize=False, capacity_m=capacity - ) - common = dict( - hidden_size=hidden_dim, - top_k=top_k, - tp_size=tp_size, - rank=tp_rank, - capacity_m=capacity, - rms_epsilon=rms_eps, - routed_scaling_factor=routed_scaling_factor, - weight_bias=weight_bias, - include_shared_expert=include_shared_expert, - add_residual=add_residual, - write_residual_output=write_residual_output, - group=group, - ) - instance: LLProtocol | BTProtocol | HTProtocol - if protocol is ProtocolKind.LL: - instance = LLProtocol( - **common, - finalize_tunings=finalize_tunings, - all_reduce_tunings=all_reduce_tunings, - ) - elif protocol is ProtocolKind.BT: - instance = BTProtocol( - **common, - finalize_tunings=finalize_tunings, - all_reduce_tunings=all_reduce_tunings, - ) - else: - instance = HTProtocol( - **common, - finalize_tunings=finalize_tunings, - all_reduce_tunings=all_reduce_tunings, - ) - self._protocols[protocol] = instance - - torch.cuda.synchronize(device) - dist.barrier(group=group) - - @staticmethod - def _tunings( - profile: StaticProfile, - protocol: ProtocolKind, - *, - finalize: bool, - capacity_m: int, - ) -> tuple: - routes = profile.finalize_routes if finalize else profile.all_reduce_routes - tunings = tuple( - dict.fromkeys( - target.preset - for target in routes.targets_for_capacity(capacity_m) - if target.protocol is protocol - ) - ) - expected_type = { - (ProtocolKind.LL, True): LLFinalizeTuning, - (ProtocolKind.LL, False): LLAllReduceTuning, - (ProtocolKind.BT, True): BTFinalizeTuning, - (ProtocolKind.BT, False): BTAllReduceTuning, - (ProtocolKind.HT, True): HTFinalizeTuning, - (ProtocolKind.HT, False): HTAllReduceTuning, - }[(protocol, finalize)] - if not all(isinstance(tuning, expected_type) for tuning in tunings): - path = "finalize" if finalize else "all-reduce" - raise TypeError(f"Invalid {protocol.value} {path} preset") - return tunings - - def _uses_pdl(self, pattern: int, m: int) -> bool: - if pattern == AllReduceFusionPattern.kMoEFinalizeARResidualRMSNorm: - target = self.profile.finalize_routes.select(m) - elif pattern == AllReduceFusionPattern.kARResidualRMSNorm: - target = self.profile.all_reduce_routes.select(m) - else: - raise NotImplementedError("Unsupported MNNVL CuTe DSL fusion pattern") - preset = cast(Any, target.preset) - enabled = getattr(preset, "enable_pdl", None) - if enabled is None: - enabled = preset.collective.enable_pdl - return bool(enabled) - - @property - def backend(self) -> str: - return "mnnvl-cutedsl" - - def is_buffer_size_sufficient( - self, - tp_size: int, - num_tokens: int, - hidden_dim: int, - dtype: torch.dtype, - use_oneshot=None, - ) -> bool: - del use_oneshot - return ( - tp_size == self.world_size - and num_tokens <= self.max_token_num - and hidden_dim == self.hidden_dim - and dtype == self.dtype - ) - - def _finalize_all_reduce_rms_norm( - self, - routed_output: torch.Tensor, - expert_weights: torch.Tensor, - permuted_indices: torch.Tensor, - shared_output: torch.Tensor | None, - residual_source: torch.Tensor | None, - gamma: torch.Tensor, - m: int, - *, - norm_output: torch.Tensor | None, - residual_output: torch.Tensor | None, - ) -> tuple[torch.Tensor, torch.Tensor | None]: - target = self.profile.finalize_routes.select(m) - protocol = cast(Any, self._protocols[target.protocol]) - kernel = protocol.finalize_kernels[target.preset] - return kernel( - routed_output, - expert_weights, - permuted_indices, - shared_output, - residual_source, - gamma, - m, - state=protocol.state, - norm_output=norm_output, - residual_output=residual_output, - ) - - def _all_reduce_rms_norm( - self, - local_contribution: torch.Tensor, - residual_source: torch.Tensor | None, - gamma: torch.Tensor, - m: int, - *, - norm_output: torch.Tensor | None, - residual_output: torch.Tensor | None, - ) -> tuple[torch.Tensor, torch.Tensor | None]: - target = self.profile.all_reduce_routes.select(m) - protocol = cast(Any, self._protocols[target.protocol]) - kernel = protocol.all_reduce_kernels[target.preset] - return kernel( - local_contribution, - residual_source, - gamma, - m, - state=protocol.state, - norm_output=norm_output, - residual_output=residual_output, - ) - - def destroy(self) -> None: - if self._destroyed: - return - self._protocols.clear() - self._destroyed = True - - -def _mnnvl_cutedsl_allreduce_fusion( - input: torch.Tensor, - workspace: MNNVLCuteDSLAllReduceFusionWorkspace, - pattern: int, - *, - launch_with_pdl: bool, - output: Optional[torch.Tensor] = None, - residual_in: Optional[torch.Tensor] = None, - residual_out: Optional[torch.Tensor] = None, - norm_out: Optional[torch.Tensor] = None, - quant_out: Optional[torch.Tensor] = None, - scale_out: Optional[torch.Tensor] = None, - rms_gamma: Optional[torch.Tensor] = None, - rms_eps: float = 1e-6, - scale_factor: Optional[torch.Tensor | float] = None, - layout_code: Optional[int] = None, - use_oneshot: Optional[bool] = None, - fp32_acc: bool = False, - moe_reduction_device_num_experts: Optional[int] = None, - moe_reduction_scale_input: Optional[torch.Tensor] = None, - moe_reduction_active_experts_token_input: Optional[torch.Tensor] = None, - moe_reduction_token_input: Optional[torch.Tensor] = None, - weight_bias: float = 0.0, - expanded_idx_to_permuted_idx: Optional[torch.Tensor] = None, - expert_scale_factor: Optional[torch.Tensor] = None, - shared_expert_output: Optional[torch.Tensor] = None, - block_quant_group_size: Optional[int] = None, -) -> torch.Tensor: - if workspace._destroyed: - raise RuntimeError( - "The MNNVLCuteDSLAllReduceFusionWorkspace has been destroyed" - ) - if pattern not in ( - AllReduceFusionPattern.kARResidualRMSNorm, - AllReduceFusionPattern.kMoEFinalizeARResidualRMSNorm, - ): - raise NotImplementedError("Unsupported MNNVL CuTe DSL fusion pattern") - unsupported = [ - name - for name, value in ( - ("output", output), - ("quant_out", quant_out), - ("scale_out", scale_out), - ("scale_factor", scale_factor), - ("layout_code", layout_code), - ("use_oneshot", use_oneshot), - ("block_quant_group_size", block_quant_group_size), - ("moe_reduction_scale_input", moe_reduction_scale_input), - ( - "moe_reduction_active_experts_token_input", - moe_reduction_active_experts_token_input, - ), - ("moe_reduction_token_input", moe_reduction_token_input), - ) - if value is not None - ] - if fp32_acc: - unsupported.append("fp32_acc") - if moe_reduction_device_num_experts is not None: - unsupported.append("moe_reduction_device_num_experts") - if unsupported: - raise ValueError("MNNVL CuTe DSL does not support: " + ", ".join(unsupported)) - - if rms_eps != workspace.rms_eps: - raise ValueError("rms_eps does not match the compiled workspace") - if weight_bias != workspace.weight_bias: - raise ValueError("weight_bias does not match the compiled workspace") - if rms_gamma is None: - raise ValueError("rms_gamma is required") - if workspace.add_residual and residual_in is None: - raise ValueError("residual_in is required by the compiled workspace") - if not workspace.add_residual and residual_in is not None: - raise ValueError("residual_in must be None for this compiled workspace") - if not workspace.write_residual_output and residual_out is not None: - raise ValueError("residual_out must be None for this compiled workspace") - - device = torch.device("cuda", torch.cuda.current_device()) - hidden = workspace.hidden_dim - _check_tensor( - rms_gamma, - "rms_gamma", - shape=(hidden,), - dtype=torch.bfloat16, - device=device, - alignment=16, - ) - - if pattern == AllReduceFusionPattern.kARResidualRMSNorm: - if any( - value is not None - for value in ( - expanded_idx_to_permuted_idx, - expert_scale_factor, - shared_expert_output, - ) - ): - raise ValueError("MoE finalize operands require the finalize pattern") - _check_tensor( - input, - "input", - shape=(None, hidden), - dtype=torch.bfloat16, - device=device, - alignment=16, - ) - m = input.shape[0] - if not 1 <= m <= workspace.max_token_num: - raise ValueError("input token count exceeds workspace capacity") - _warn_pdl_mismatch(workspace, pattern, m, launch_with_pdl) - if residual_in is not None: - _check_tensor( - residual_in, - "residual_in", - shape=(m, hidden), - dtype=torch.bfloat16, - device=device, - alignment=16, - ) - if norm_out is not None: - _check_tensor( - norm_out, - "norm_out", - shape=(m, hidden), - dtype=torch.bfloat16, - device=device, - alignment=16, - ) - if residual_out is not None: - _check_tensor( - residual_out, - "residual_out", - shape=(m, hidden), - dtype=torch.bfloat16, - device=device, - alignment=16, - ) - norm_out, _ = workspace._all_reduce_rms_norm( - input, - residual_in, - rms_gamma, - input.shape[0], - norm_output=norm_out, - residual_output=residual_out, - ) - return norm_out - - if pattern == AllReduceFusionPattern.kMoEFinalizeARResidualRMSNorm: - if expanded_idx_to_permuted_idx is None: - raise ValueError("expanded_idx_to_permuted_idx is required") - if expert_scale_factor is None: - raise ValueError("expert_scale_factor is required") - if workspace.include_shared_expert and shared_expert_output is None: - raise ValueError( - "shared_expert_output is required by the compiled workspace" - ) - if not workspace.include_shared_expert and shared_expert_output is not None: - raise ValueError( - "shared_expert_output must be None for this compiled workspace" - ) - m = expanded_idx_to_permuted_idx.shape[0] - if not 1 <= m <= workspace.max_token_num: - raise ValueError("input token count exceeds workspace capacity") - _warn_pdl_mismatch(workspace, pattern, m, launch_with_pdl) - _check_tensor( - input, - "input", - shape=(None, hidden), - dtype=torch.bfloat16, - device=device, - alignment=16, - ) - _check_tensor( - expert_scale_factor, - "expert_scale_factor", - shape=(m, workspace.top_k), - dtype=torch.bfloat16, - device=device, - alignment=2, - ) - _check_tensor( - expanded_idx_to_permuted_idx, - "expanded_idx_to_permuted_idx", - shape=(m, workspace.top_k), - dtype=torch.int32, - device=device, - alignment=4, - ) - if shared_expert_output is not None: - _check_tensor( - shared_expert_output, - "shared_expert_output", - shape=(m, hidden), - dtype=torch.bfloat16, - device=device, - alignment=16, - ) - if residual_in is not None: - _check_tensor( - residual_in, - "residual_in", - shape=(m, hidden), - dtype=torch.bfloat16, - device=device, - alignment=16, - ) - if norm_out is not None: - _check_tensor( - norm_out, - "norm_out", - shape=(m, hidden), - dtype=torch.bfloat16, - device=device, - alignment=16, - ) - if residual_out is not None: - _check_tensor( - residual_out, - "residual_out", - shape=(m, hidden), - dtype=torch.bfloat16, - device=device, - alignment=16, - ) - norm_out, _ = workspace._finalize_all_reduce_rms_norm( - input, - expert_scale_factor, - expanded_idx_to_permuted_idx, - shared_expert_output, - residual_in, - rms_gamma, - m, - norm_output=norm_out, - residual_output=residual_out, - ) - return norm_out - - raise AssertionError("unreachable") - - -# Keep the upstream private name intact while exposing SGLang's backend entry point; -# upstream exports it through flashinfer.comm.allreduce_fusion. -mnnvl_cutedsl_allreduce_fusion = _mnnvl_cutedsl_allreduce_fusion diff --git a/python/sglang/srt/layers/flashinfer_mnnvl_cutedsl.py b/python/sglang/srt/layers/flashinfer_mnnvl_cutedsl.py index 653197c84..da085bb43 100644 --- a/python/sglang/srt/layers/flashinfer_mnnvl_cutedsl.py +++ b/python/sglang/srt/layers/flashinfer_mnnvl_cutedsl.py @@ -2,7 +2,6 @@ from __future__ import annotations -import importlib.util import logging import threading from dataclasses import dataclass, replace @@ -21,34 +20,15 @@ logger = logging.getLogger(__name__) def _import_kernel_backend(): - try: - from flashinfer.comm import AllReduceFusionPattern - except ImportError as error: - raise RuntimeError( - "MNNVL CuTe DSL fusion requires FlashInfer's communication " - "infrastructure (flashinfer >= 0.6.16)" - ) from error - try: - from sglang.kernels.ops.communication.mnnvl_cutedsl import DEFAULT_CONFIG - from sglang.kernels.ops.communication.mnnvl_cutedsl_ar import ( - MNNVLCuteDSLAllReduceFusionWorkspace, - mnnvl_cutedsl_allreduce_fusion, - ) - except ImportError as error: - raise RuntimeError( - "SGLang's in-tree MNNVL CuTe DSL kernels failed to import; check " - "their dependencies, including nvidia-cutlass-dsl and cuda-python" - ) from error - if importlib.util.find_spec("flashinfer.comm.mnnvl_cutedsl") is not None: - logger.warning( - "The installed FlashInfer now ships flashinfer.comm.mnnvl_cutedsl; " - "SGLang is still running its in-tree port " - "(sglang.kernels.ops.communication.mnnvl_cutedsl), which can now " - "be retired in favor of the upstream backend." - ) + # Imported here rather than at module scope: the CuTe DSL backend drags in + # CUDA-only dependencies that CPU-side importers of this module never need. + from flashinfer.comm import AllReduceFusionPattern, allreduce_fusion + from flashinfer.comm.mnnvl_cutedsl import DEFAULT_CONFIG + from flashinfer.comm.mnnvl_cutedsl_ar import MNNVLCuteDSLAllReduceFusionWorkspace + return ( MNNVLCuteDSLAllReduceFusionWorkspace, - mnnvl_cutedsl_allreduce_fusion, + allreduce_fusion, AllReduceFusionPattern, DEFAULT_CONFIG, ) @@ -180,8 +160,8 @@ class FlashInferMNNVLCuteDSLARFusion: config=self.workspace_config, ) - # Publish only after the mailbox barrier; older FlashInfer workspace - # classes may not provide it and would desynchronize Lamport stages. + # Publish only after the mailbox barrier; without it the ranks + # would desynchronize their Lamport stages. torch.cuda.synchronize(self.device) dist.barrier(group=process_group)