diff --git a/docker/rocm.Dockerfile b/docker/rocm.Dockerfile index f19f11e57..380bc5df9 100644 --- a/docker/rocm.Dockerfile +++ b/docker/rocm.Dockerfile @@ -104,7 +104,7 @@ ARG ENABLE_MORI=0 ARG NIC_BACKEND=none ARG MORI_REPO="https://github.com/ROCm/mori.git" -ARG MORI_COMMIT="v1.1.1" +ARG MORI_COMMIT="96ffa169710f214e76e07abe5008d686fe54522b" # AMD AINIC apt repo settings ARG AINIC_VERSION=1.117.5-a-38 diff --git a/python/sglang/srt/layers/moe/token_dispatcher/moriep.py b/python/sglang/srt/layers/moe/token_dispatcher/moriep.py index 013ed92f0..a77db62c8 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/moriep.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/moriep.py @@ -40,6 +40,10 @@ from sglang.srt.distributed import ( ) from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype +# Blockwise quantization group sizes: number of elements sharing one scale factor +FP8_BLOCK_SIZE = 128 +MXFP4_BLOCK_SIZE = 32 + _is_hip = is_hip() _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip @@ -128,6 +132,18 @@ class EpMode(Enum): LOW_LATENCY = "low_latency" +class DispatchDtype(Enum): + bf16 = "bfloat16" + fp8 = "float8_blockwise" + fp4 = "mxfp4_blockwise" + + +class CombineDtype(Enum): + bf16 = "bfloat16" + fp8 = "float8_blockwise" + fp8_direct_cast = "float8_direct_cast" + + @dataclass(frozen=True) class EpDispatchConfig: kernel_type: mori.ops.EpDispatchCombineKernelType @@ -188,8 +204,8 @@ def init_mori_op( num_max_dispatch_tokens_per_rank, deepep_mode, instance_id=0, - fp8_dispatch=False, - fp4_dispatch=False, + dispatch_dtype=DispatchDtype.bf16, + combine_dtype=CombineDtype.bf16, enable_sdma=False, ): @@ -235,15 +251,15 @@ def init_mori_op( data_type = fp8_dtype scale_type_size = torch.float32.itemsize - if fp8_dispatch: - scale_dim = hidden_size // 128 - elif fp4_dispatch: + if dispatch_dtype == DispatchDtype.fp8: + scale_dim = hidden_size // FP8_BLOCK_SIZE + elif dispatch_dtype == DispatchDtype.fp4: # FP4 kernel still takes the original hidden size and do quantization # internally, so hidden_dim is not reduced. The reason is that for FP4 # quantization, we need to keep the original hidden size to calculate # the quantization scale correctly. Don't use packed hidden size for FP4 kernel. hidden_dim = hidden_size - scale_dim = hidden_size // 32 + scale_dim = hidden_size // MXFP4_BLOCK_SIZE data_type = torch.float4_e2m1fn_x2 scale_type_size = torch.float8_e8m0fnu.itemsize @@ -255,15 +271,19 @@ def init_mori_op( block_num = 256 warp_num_per_block = 16 + # Fp8 blockwise combine uses its own internal scale_dim driven which can be + # overridden by env ``MORI_FP8_COMBINE_SCALE_DIM`` (default 56) + # See https://github.com/ROCm/mori/blob/96ffa169710f214e76e07abe5008d686fe54522b/python/mori/ops/dispatch_combine.py#L81-L84 combine_quant_type = "none" - if get_bool_env_var("SGLANG_MORI_FP8_COMB", "False"): + if combine_dtype == CombineDtype.fp8: + combine_quant_type = "fp8_blockwise" + elif combine_dtype == CombineDtype.fp8_direct_cast: combine_quant_type = "fp8_direct_cast" logger.info( f"[MORI init] {world_size=} {rank=} {hidden_size=} {params_dtype=} " f"{num_max_dispatch_tokens_per_rank=} {num_local_experts=} " - f"{router_topk=} {mode=} {fp8_dispatch=} {fp4_dispatch=} " - f"{combine_quant_type=}" + f"{router_topk=} {mode=} {dispatch_dtype=} {combine_dtype=} " ) def check_mori_compatibility(kwargs: dict) -> None: @@ -365,8 +385,8 @@ class _MoriEPDispatcherImplBase: self.enable_sdma = get_bool_env_var("MORI_ENABLE_SDMA", "false") self._mori_op = None - self.fp8_dispatch = False - self.fp4_dispatch = False + self.dispatch_dtype = DispatchDtype.bf16 + self.combine_dtype = CombineDtype.bf16 self.quant_config: Optional[dict] = None @@ -389,19 +409,23 @@ class _MoriEPDispatcherImplBase: self.num_max_dispatch_tokens_per_rank, self.deepep_mode, self.instance_id, - self.fp8_dispatch, - self.fp4_dispatch, + self.dispatch_dtype, + self.combine_dtype, self.enable_sdma, ) return self._mori_op def _apply_dispatch_dtype_override(self): - """Apply env var override to fp8_dispatch/fp4_dispatch flags.""" + """Apply env var override to fp8_dispatch/fp4_dispatch/fp8_combine flags.""" if "SGLANG_MORI_DISPATCH_DTYPE" in os.environ: dispatch_dtype = os.environ["SGLANG_MORI_DISPATCH_DTYPE"].lower() if dispatch_dtype != "auto": - self.fp8_dispatch = dispatch_dtype == "fp8" - self.fp4_dispatch = dispatch_dtype == "fp4" + if dispatch_dtype == "bf16": + self.dispatch_dtype = DispatchDtype.bf16 + elif dispatch_dtype == "fp8": + self.dispatch_dtype = DispatchDtype.fp8 + elif dispatch_dtype == "fp4": + self.dispatch_dtype = DispatchDtype.fp4 elif ( "SGLANG_MORI_FP8_DISP" in os.environ or "SGLANG_MORI_FP4_DISP" in os.environ ): @@ -411,8 +435,29 @@ class _MoriEPDispatcherImplBase: "and will be removed in a future release. " "Use SGLANG_MORI_DISPATCH_DTYPE=auto|bf16|fp8|fp4 instead." ) - self.fp8_dispatch = get_bool_env_var("SGLANG_MORI_FP8_DISP", "False") - self.fp4_dispatch = get_bool_env_var("SGLANG_MORI_FP4_DISP", "False") + if get_bool_env_var("SGLANG_MORI_FP8_DISP", "False"): + self.dispatch_dtype = DispatchDtype.fp8 + if get_bool_env_var("SGLANG_MORI_FP4_DISP", "False"): + self.dispatch_dtype = DispatchDtype.fp4 + + if "SGLANG_MORI_COMBINE_DTYPE" in os.environ: + combine_dtype = os.environ["SGLANG_MORI_COMBINE_DTYPE"].lower() + if combine_dtype != "auto": + if combine_dtype == "fp8": + self.combine_dtype = CombineDtype.fp8 + elif combine_dtype == "bf16": + self.combine_dtype = CombineDtype.bf16 + elif combine_dtype == "fp8_direct_cast": + self.combine_dtype = CombineDtype.fp8_direct_cast + elif "SGLANG_MORI_FP8_COMB" in os.environ: + # Deprecated: will be removed in a future release + logger.warning_once( + "SGLANG_MORI_FP8_COMB is deprecated " + "and will be removed in a future release. " + "Use SGLANG_MORI_COMBINE_DTYPE=auto|bf16|fp8|fp8_direct_cast instead." + ) + if get_bool_env_var("SGLANG_MORI_FP8_COMB", "False"): + self.combine_dtype = CombineDtype.fp8 def dispatch_a( self, @@ -440,14 +485,14 @@ class _MoriEPDispatcherImplBase: # Auto-detect dispatch quantization from weight dtype weight_dtype = quant_config.get("weight_dtype", None) if weight_dtype in (torch.float8_e4m3fn, torch.float8_e4m3fnuz): - self.fp8_dispatch = True - self.fp4_dispatch = False + self.dispatch_dtype = DispatchDtype.fp8 + self.combine_dtype = CombineDtype.bf16 elif weight_dtype == torch.float4_e2m1fn_x2: - self.fp8_dispatch = False - self.fp4_dispatch = True + self.dispatch_dtype = DispatchDtype.fp4 + self.combine_dtype = CombineDtype.fp8 else: - self.fp8_dispatch = False - self.fp4_dispatch = False + self.dispatch_dtype = DispatchDtype.bf16 + self.combine_dtype = CombineDtype.bf16 # Apply env var override immediately so dispatch_a sees correct flags self._apply_dispatch_dtype_override() @@ -494,9 +539,7 @@ class _MoriEPDispatcherImplNormal(_MoriEPDispatcherImplBase): output_dtype = hidden_states.dtype scale = None - fp8_dispatch, fp4_dispatch = self.fp8_dispatch, self.fp4_dispatch - - if fp8_dispatch: + if self.dispatch_dtype == DispatchDtype.fp8: # FP8 quant if num_token > 0: # NOTE: aiter is able to handle token=0 case in UT. But for some @@ -509,12 +552,12 @@ class _MoriEPDispatcherImplNormal(_MoriEPDispatcherImplBase): hidden_states.shape, dtype=fp8_dtype, device=hidden_states.device ) scale = torch.empty( - (0, self.hidden_size // 128), + (0, self.hidden_size // FP8_BLOCK_SIZE), dtype=torch.float32, device=hidden_states.device, ) - elif fp4_dispatch: + elif self.dispatch_dtype == DispatchDtype.fp4: # FP4 quant if num_token > 0: hidden_states, scale = self.fp4_quant_func(hidden_states, shuffle=False) @@ -525,7 +568,7 @@ class _MoriEPDispatcherImplNormal(_MoriEPDispatcherImplBase): device=hidden_states.device, ) scale = torch.empty( - (0, self.hidden_size // 32), + (0, self.hidden_size // MXFP4_BLOCK_SIZE), dtype=torch.float8_e8m0fnu, device=hidden_states.device, ) @@ -752,9 +795,7 @@ class _MoriEPDispatcherImplLowLatency(_MoriEPDispatcherImplBase): output_dtype = hidden_states.dtype scale = None - fp8_dispatch, fp4_dispatch = self.fp8_dispatch, self.fp4_dispatch - - if fp8_dispatch: + if self.dispatch_dtype == DispatchDtype.fp8: # FP8 quant if num_tokens > 0: # NOTE: aiter is able to handle token=0 case in UT. But for some @@ -767,12 +808,12 @@ class _MoriEPDispatcherImplLowLatency(_MoriEPDispatcherImplBase): hidden_states.shape, dtype=fp8_dtype, device=hidden_states.device ) scale = torch.empty( - (0, self.hidden_size // 128), + (0, self.hidden_size // FP8_BLOCK_SIZE), dtype=torch.float32, device=hidden_states.device, ) - elif fp4_dispatch: + elif self.dispatch_dtype == DispatchDtype.fp4: # FP4 quant if num_tokens > 0: hidden_states, scale = self.fp4_quant_func(hidden_states, shuffle=False) @@ -783,7 +824,7 @@ class _MoriEPDispatcherImplLowLatency(_MoriEPDispatcherImplBase): device=hidden_states.device, ) scale = torch.empty( - (0, self.hidden_size // 32), + (0, self.hidden_size // MXFP4_BLOCK_SIZE), dtype=torch.float8_e8m0fnu, device=hidden_states.device, )