Support Flashinfer one-sided A2A + CuteDSL MoE for Nemotron Ultra (#28309)
Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
co-authored by
Brayden Zhong
parent
ac23be8d09
commit
edb2059139
@@ -33,6 +33,7 @@ def flashinfer_cutedsl_moe_masked(
|
|||||||
down_sm_count: Optional[int] = None,
|
down_sm_count: Optional[int] = None,
|
||||||
down_signals: Optional[torch.Tensor] = None,
|
down_signals: Optional[torch.Tensor] = None,
|
||||||
down_start_event: Optional[torch.cuda.Event] = None,
|
down_start_event: Optional[torch.cuda.Event] = None,
|
||||||
|
activation: str = "silu",
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Perform masked Mixture-of-Experts computation with FlashInfer's CuteDSL
|
Perform masked Mixture-of-Experts computation with FlashInfer's CuteDSL
|
||||||
@@ -103,7 +104,21 @@ def flashinfer_cutedsl_moe_masked(
|
|||||||
input_global_scale,
|
input_global_scale,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert w1.shape[-2] == 2 * n, f"w1 last-2 dim must be 2*n, got {w1.shape}"
|
if activation == "silu":
|
||||||
|
gated = True
|
||||||
|
elif activation == "relu2":
|
||||||
|
gated = False
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
f"CuteDSL masked MoE supports activation 'silu' (gated) or "
|
||||||
|
f"'relu2' (non-gated), got {activation!r}."
|
||||||
|
)
|
||||||
|
# Gated (silu_and_mul) GEMM1 emits [gate, up] so w1 has 2*n rows; non-gated
|
||||||
|
# relu2 emits a single projection of n rows.
|
||||||
|
gemm1_out_dim = 2 * n if gated else n
|
||||||
|
assert (
|
||||||
|
w1.shape[-2] == gemm1_out_dim
|
||||||
|
), f"w1 last-2 dim must be {gemm1_out_dim} (gated={gated}), got {w1.shape}"
|
||||||
assert (
|
assert (
|
||||||
w1.shape[-1] * 2 == k
|
w1.shape[-1] * 2 == k
|
||||||
), f"w1 last dim * 2 must equal k, got {w1.shape[-1]} vs k={k}"
|
), f"w1 last dim * 2 must equal k, got {w1.shape[-1]} vs k={k}"
|
||||||
@@ -123,7 +138,7 @@ def flashinfer_cutedsl_moe_masked(
|
|||||||
|
|
||||||
# TODO(kaixih@nvidia): dtype should be based on inputs.
|
# TODO(kaixih@nvidia): dtype should be based on inputs.
|
||||||
gateup_output = torch.empty(
|
gateup_output = torch.empty(
|
||||||
(num_experts, m, n * 2), dtype=torch.bfloat16, device=a_q.device
|
(num_experts, m, gemm1_out_dim), dtype=torch.bfloat16, device=a_q.device
|
||||||
)
|
)
|
||||||
gateup_output = gateup_output.permute(1, 2, 0) # requirement of kernel
|
gateup_output = gateup_output.permute(1, 2, 0) # requirement of kernel
|
||||||
sf_vec_size = 16
|
sf_vec_size = 16
|
||||||
@@ -147,12 +162,23 @@ def flashinfer_cutedsl_moe_masked(
|
|||||||
alpha_dtype=get_cute_dtype(w1_alpha),
|
alpha_dtype=get_cute_dtype(w1_alpha),
|
||||||
) # in logical [m, n, l]
|
) # in logical [m, n, l]
|
||||||
|
|
||||||
# SILU and quantization
|
# Activation + NVFP4 quantization of the GEMM2 input.
|
||||||
|
if gated:
|
||||||
|
# Fused silu(gate) * up + quantize; halves 2*n -> n.
|
||||||
diq, diq_sf = silu_and_mul_scaled_nvfp4_experts_quantize(
|
diq, diq_sf = silu_and_mul_scaled_nvfp4_experts_quantize(
|
||||||
gateup_output.permute(2, 0, 1),
|
gateup_output.permute(2, 0, 1),
|
||||||
masked_m,
|
masked_m,
|
||||||
a2_global_scale,
|
a2_global_scale,
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
# Non-gated relu^2: relu(x)^2 elementwise (no halving), then grouped
|
||||||
|
# NVFP4 quantize. scaled_fp4_grouped_quantize needs a per-expert (l,)
|
||||||
|
# global scale, so broadcast a scalar if one was supplied.
|
||||||
|
act = torch.relu(gateup_output.permute(2, 0, 1)).square().contiguous()
|
||||||
|
a2_gs = a2_global_scale
|
||||||
|
if a2_gs.numel() == 1:
|
||||||
|
a2_gs = a2_gs.expand(num_experts).contiguous()
|
||||||
|
diq, diq_sf = scaled_fp4_grouped_quantize(act, masked_m, a2_gs)
|
||||||
|
|
||||||
if down_start_event is not None:
|
if down_start_event is not None:
|
||||||
down_start_event.record()
|
down_start_event.record()
|
||||||
|
|||||||
@@ -284,6 +284,7 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None:
|
|||||||
local_expert_offset=layer.moe_ep_rank * layer.num_local_experts,
|
local_expert_offset=layer.moe_ep_rank * layer.num_local_experts,
|
||||||
output_dtype=layer.moe_runner_config.params_dtype,
|
output_dtype=layer.moe_runner_config.params_dtype,
|
||||||
device=str(layer.w13_weight.device),
|
device=str(layer.w13_weight.device),
|
||||||
|
activation=layer.moe_runner_config.activation,
|
||||||
)
|
)
|
||||||
|
|
||||||
w1_alpha, fc2_input_scale, w2_alpha, used_input_scale = (
|
w1_alpha, fc2_input_scale, w2_alpha, used_input_scale = (
|
||||||
@@ -355,7 +356,10 @@ def fused_experts_none_to_flashinfer_cutedsl_fp4(
|
|||||||
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
||||||
from sglang.srt.layers.quantization.fp4_utils import fp4_quantize
|
from sglang.srt.layers.quantization.fp4_utils import fp4_quantize
|
||||||
|
|
||||||
assert runner_config.activation == "silu", "Only silu is supported for CuteDSL MoE."
|
assert runner_config.activation in (
|
||||||
|
"silu",
|
||||||
|
"relu2",
|
||||||
|
), f"CuteDSL MoE supports 'silu' (gated) or 'relu2' (non-gated), got {runner_config.activation!r}."
|
||||||
assert quant_info.wrapper is not None, "CuteDSL v2 path requires CuteDslMoEWrapper."
|
assert quant_info.wrapper is not None, "CuteDSL v2 path requires CuteDslMoEWrapper."
|
||||||
|
|
||||||
hidden_states = dispatch_output.hidden_states
|
hidden_states = dispatch_output.hidden_states
|
||||||
@@ -409,7 +413,10 @@ def fused_experts_flashinfer_to_flashinfer_cutedsl_fp4(
|
|||||||
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
||||||
from sglang.srt.layers.quantization.fp4_utils import fp4_quantize
|
from sglang.srt.layers.quantization.fp4_utils import fp4_quantize
|
||||||
|
|
||||||
assert runner_config.activation == "silu", "Only silu is supported for CuteDSL MoE."
|
assert runner_config.activation in (
|
||||||
|
"silu",
|
||||||
|
"relu2",
|
||||||
|
), f"CuteDSL MoE supports 'silu' (gated) or 'relu2' (non-gated), got {runner_config.activation!r}."
|
||||||
assert quant_info.wrapper is not None, "CuteDSL v2 path requires CuteDslMoEWrapper."
|
assert quant_info.wrapper is not None, "CuteDSL v2 path requires CuteDslMoEWrapper."
|
||||||
|
|
||||||
hidden_states = dispatch_output.hidden_states
|
hidden_states = dispatch_output.hidden_states
|
||||||
@@ -468,7 +475,10 @@ def fused_experts_deepep_to_flashinfer_cutedsl_fp4(
|
|||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPLLCombineInput
|
from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPLLCombineInput
|
||||||
|
|
||||||
assert runner_config.activation == "silu", "Only silu is supported for CuteDSL MoE."
|
assert runner_config.activation in (
|
||||||
|
"silu",
|
||||||
|
"relu2",
|
||||||
|
), f"CuteDSL masked MoE supports 'silu' or 'relu2', got {runner_config.activation!r}."
|
||||||
assert (
|
assert (
|
||||||
not runner_config.apply_router_weight_on_input
|
not runner_config.apply_router_weight_on_input
|
||||||
), "apply_router_weight_on_input is not supported for Flashinfer"
|
), "apply_router_weight_on_input is not supported for Flashinfer"
|
||||||
@@ -504,6 +514,7 @@ def fused_experts_deepep_to_flashinfer_cutedsl_fp4(
|
|||||||
w2_blockscale=quant_info.w2_weight_sf,
|
w2_blockscale=quant_info.w2_weight_sf,
|
||||||
w2_alpha=quant_info.w2_alpha,
|
w2_alpha=quant_info.w2_alpha,
|
||||||
masked_m=masked_m,
|
masked_m=masked_m,
|
||||||
|
activation=runner_config.activation,
|
||||||
**(
|
**(
|
||||||
dict(
|
dict(
|
||||||
down_sm_count=overlap.num_sms,
|
down_sm_count=overlap.num_sms,
|
||||||
|
|||||||
@@ -935,18 +935,14 @@ def fused_experts_none_to_flashinfer_trtllm_fp4(
|
|||||||
topk_output = dispatch_output.topk_output
|
topk_output = dispatch_output.topk_output
|
||||||
|
|
||||||
# Quantize hidden states to FP4
|
# Quantize hidden states to FP4
|
||||||
hidden_states_scale = (
|
hidden_states_scale = dispatch_output.hidden_states_scale
|
||||||
dispatch_output.hidden_states_scale
|
|
||||||
if hasattr(dispatch_output, "hidden_states_scale")
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
per_token_scale = None
|
per_token_scale = None
|
||||||
if hidden_states_scale is not None:
|
if hidden_states_scale is not None:
|
||||||
# NVFP4 dispatch, inputs are already quantized.
|
# NVFP4 dispatch (flashinfer a2a): inputs are already FP4-quantized by
|
||||||
|
# the dispatcher, so pass them through unchanged.
|
||||||
hs_fp4 = hidden_states
|
hs_fp4 = hidden_states
|
||||||
hs_scale_linear = hidden_states_scale
|
hs_scale_linear = hidden_states_scale
|
||||||
elif quant_info.use_per_token_activation:
|
elif quant_info.use_per_token_activation:
|
||||||
# Enable FlashInfer TRTLLM per-token NVFP4 activation scaling; ignores checkpoint activation FP32 scale by treating it as
|
|
||||||
from flashinfer import SfLayout, nvfp4_quantize
|
from flashinfer import SfLayout, nvfp4_quantize
|
||||||
|
|
||||||
e4m3_max = 448.0
|
e4m3_max = 448.0
|
||||||
@@ -996,6 +992,8 @@ def fused_experts_none_to_flashinfer_trtllm_fp4(
|
|||||||
hidden_size = (
|
hidden_size = (
|
||||||
hs_fp4.shape[-1] * 2 if hs_fp4.dtype == torch.uint8 else hs_fp4.shape[-1]
|
hs_fp4.shape[-1] * 2 if hs_fp4.dtype == torch.uint8 else hs_fp4.shape[-1]
|
||||||
)
|
)
|
||||||
|
# When the dispatcher delivered pre-quantized FP4 (hidden_states is uint8),
|
||||||
|
# the MoE output is bf16 rather than the input dtype.
|
||||||
output_dtype = (
|
output_dtype = (
|
||||||
hidden_states.dtype if hidden_states_scale is None else torch.bfloat16
|
hidden_states.dtype if hidden_states_scale is None else torch.bfloat16
|
||||||
)
|
)
|
||||||
@@ -1018,7 +1016,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp4(
|
|||||||
else:
|
else:
|
||||||
with use_symmetric_memory(get_tp_group(), disabled=not _symm_required):
|
with use_symmetric_memory(get_tp_group(), disabled=not _symm_required):
|
||||||
symm_output = torch.empty(
|
symm_output = torch.empty(
|
||||||
hs_fp4.shape[0],
|
num_tokens,
|
||||||
hidden_size,
|
hidden_size,
|
||||||
dtype=output_dtype,
|
dtype=output_dtype,
|
||||||
device=hs_fp4.device,
|
device=hs_fp4.device,
|
||||||
|
|||||||
@@ -1509,7 +1509,7 @@ def biased_grouped_topk_gpu(
|
|||||||
and topk_group == 1
|
and topk_group == 1
|
||||||
and num_fused_shared_experts == 0
|
and num_fused_shared_experts == 0
|
||||||
and num_experts <= 512
|
and num_experts <= 512
|
||||||
and topk <= 8
|
and topk <= 32
|
||||||
):
|
):
|
||||||
# Ungrouped sigmoid (num_expert_group == 1): use the unified Triton
|
# Ungrouped sigmoid (num_expert_group == 1): use the unified Triton
|
||||||
# router, which subsumes the jit grouped_topk.cuh kernel here.
|
# router, which subsumes the jit grouped_topk.cuh kernel here.
|
||||||
|
|||||||
@@ -53,6 +53,7 @@ from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
|||||||
from sglang.srt.layers.moe.topk import TopK
|
from sglang.srt.layers.moe.topk import TopK
|
||||||
from sglang.srt.layers.moe.utils import (
|
from sglang.srt.layers.moe.utils import (
|
||||||
RoutingMethodType,
|
RoutingMethodType,
|
||||||
|
get_moe_a2a_backend,
|
||||||
should_skip_post_experts_all_reduce,
|
should_skip_post_experts_all_reduce,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.quantization import QuantizationConfig
|
from sglang.srt.layers.quantization import QuantizationConfig
|
||||||
@@ -65,6 +66,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
from sglang.srt.model_executor.forward_context import get_attn_backend
|
from sglang.srt.model_executor.forward_context import get_attn_backend
|
||||||
|
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||||||
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import (
|
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import (
|
||||||
eager_on_graph,
|
eager_on_graph,
|
||||||
is_in_breakable_cuda_graph,
|
is_in_breakable_cuda_graph,
|
||||||
@@ -114,6 +116,8 @@ class NemotronHMLP(nn.Module):
|
|||||||
quant_config: QuantizationConfig | None = None,
|
quant_config: QuantizationConfig | None = None,
|
||||||
bias: bool = False,
|
bias: bool = False,
|
||||||
reduce_results: bool = True,
|
reduce_results: bool = True,
|
||||||
|
tp_rank: int | None = None,
|
||||||
|
tp_size: int | None = None,
|
||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@@ -123,6 +127,8 @@ class NemotronHMLP(nn.Module):
|
|||||||
output_size=intermediate_size,
|
output_size=intermediate_size,
|
||||||
bias=bias,
|
bias=bias,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
|
tp_rank=tp_rank,
|
||||||
|
tp_size=tp_size,
|
||||||
prefix=f"{prefix}.up_proj",
|
prefix=f"{prefix}.up_proj",
|
||||||
)
|
)
|
||||||
self.down_proj = RowParallelLinear(
|
self.down_proj = RowParallelLinear(
|
||||||
@@ -131,6 +137,8 @@ class NemotronHMLP(nn.Module):
|
|||||||
bias=bias,
|
bias=bias,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
reduce_results=reduce_results,
|
reduce_results=reduce_results,
|
||||||
|
tp_rank=tp_rank,
|
||||||
|
tp_size=tp_size,
|
||||||
prefix=f"{prefix}.down_proj",
|
prefix=f"{prefix}.down_proj",
|
||||||
)
|
)
|
||||||
self.act_fn = ReLU2()
|
self.act_fn = ReLU2()
|
||||||
@@ -223,6 +231,12 @@ class NemotronHMoE(nn.Module):
|
|||||||
* config.n_shared_experts,
|
* config.n_shared_experts,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
reduce_results=False,
|
reduce_results=False,
|
||||||
|
**(
|
||||||
|
dict(tp_rank=0, tp_size=1)
|
||||||
|
if get_moe_a2a_backend().is_deepep()
|
||||||
|
or get_moe_a2a_backend().is_flashinfer()
|
||||||
|
else {}
|
||||||
|
),
|
||||||
prefix=f"{prefix}.shared_experts",
|
prefix=f"{prefix}.shared_experts",
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -263,10 +277,14 @@ class NemotronHMoE(nn.Module):
|
|||||||
self,
|
self,
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||||
# torch.compile cannot trace CUDA streams. Take the
|
overlap = _is_cuda and not torch.compiler.is_compiling()
|
||||||
# non-overlapping path only during dynamo tracing; replay can
|
if (
|
||||||
# use the overlapping fast path since dynamo is no longer active.
|
overlap
|
||||||
if _is_cuda and not torch.compiler.is_compiling():
|
and get_moe_a2a_backend().is_flashinfer()
|
||||||
|
and not get_is_capture_mode()
|
||||||
|
):
|
||||||
|
overlap = False
|
||||||
|
if overlap:
|
||||||
return self._forward_core_shared_routed_overlap(hidden_states)
|
return self._forward_core_shared_routed_overlap(hidden_states)
|
||||||
else:
|
else:
|
||||||
return self._forward_core_normal(hidden_states)
|
return self._forward_core_normal(hidden_states)
|
||||||
@@ -454,6 +472,7 @@ class NemotronHMoEDecoderLayer(NemotronHMLPLikeDecoderLayer):
|
|||||||
self.norm,
|
self.norm,
|
||||||
for_attn=False,
|
for_attn=False,
|
||||||
allow_reduce_scatter=True,
|
allow_reduce_scatter=True,
|
||||||
|
is_sparse=True,
|
||||||
is_last_layer=layer_idx == len(config.hybrid_override_pattern) - 1,
|
is_last_layer=layer_idx == len(config.hybrid_override_pattern) - 1,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -545,10 +564,7 @@ class NemotronHMambaDecoderLayer(NemotronHAttnLikeDecoderLayer):
|
|||||||
hidden_states, residual = self._dp_attn_input(
|
hidden_states, residual = self._dp_attn_input(
|
||||||
hidden_states, residual, forward_batch
|
hidden_states, residual, forward_batch
|
||||||
)
|
)
|
||||||
if (
|
if get_real_num_tokens(hidden_states, forward_batch) == 0:
|
||||||
forward_batch.forward_mode.is_idle()
|
|
||||||
or get_real_num_tokens(hidden_states, forward_batch) == 0
|
|
||||||
):
|
|
||||||
return torch.zeros_like(hidden_states), residual
|
return torch.zeros_like(hidden_states), residual
|
||||||
|
|
||||||
output = self._forward_mamba(hidden_states, forward_batch)
|
output = self._forward_mamba(hidden_states, forward_batch)
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ from sglang.srt.layers.communicator import (
|
|||||||
apply_flashinfer_allreduce_fusion,
|
apply_flashinfer_allreduce_fusion,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.layernorm import RMSNorm
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
|
from sglang.srt.layers.moe.utils import get_moe_a2a_backend
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
|
||||||
ATTN_LAYERS = (MAMBA, ATTENTION)
|
ATTN_LAYERS = (MAMBA, ATTENTION)
|
||||||
@@ -48,12 +49,17 @@ def pad_to_original_num_tokens(
|
|||||||
return padded
|
return padded
|
||||||
|
|
||||||
|
|
||||||
def _build_layer_scatter_modes() -> LayerScatterModes:
|
def _build_layer_scatter_modes(is_sparse: bool = False) -> LayerScatterModes:
|
||||||
|
scatter_mlp = is_sparse and not get_moe_a2a_backend().is_none()
|
||||||
|
mlp_mode = ScatterMode.SCATTERED if scatter_mlp else ScatterMode.FULL
|
||||||
|
middle_residual_mode = (
|
||||||
|
ScatterMode.SCATTERED if scatter_mlp else ScatterMode.TP_ATTN_FULL
|
||||||
|
)
|
||||||
return LayerScatterModes(
|
return LayerScatterModes(
|
||||||
layer_input_mode=ScatterMode.TP_ATTN_FULL,
|
layer_input_mode=ScatterMode.TP_ATTN_FULL,
|
||||||
attn_mode=ScatterMode.TP_ATTN_FULL,
|
attn_mode=ScatterMode.TP_ATTN_FULL,
|
||||||
mlp_mode=ScatterMode.FULL,
|
mlp_mode=mlp_mode,
|
||||||
middle_residual_mode=ScatterMode.TP_ATTN_FULL,
|
middle_residual_mode=middle_residual_mode,
|
||||||
layer_output_mode=ScatterMode.TP_ATTN_FULL,
|
layer_output_mode=ScatterMode.TP_ATTN_FULL,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -63,10 +69,11 @@ def make_layer_communicator(
|
|||||||
*,
|
*,
|
||||||
for_attn: bool,
|
for_attn: bool,
|
||||||
allow_reduce_scatter: bool = False,
|
allow_reduce_scatter: bool = False,
|
||||||
|
is_sparse: bool = False,
|
||||||
is_last_layer: bool = False,
|
is_last_layer: bool = False,
|
||||||
) -> LayerCommunicator:
|
) -> LayerCommunicator:
|
||||||
return LayerCommunicator(
|
return LayerCommunicator(
|
||||||
layer_scatter_modes=_build_layer_scatter_modes(),
|
layer_scatter_modes=_build_layer_scatter_modes(is_sparse),
|
||||||
input_layernorm=layer_norm if for_attn else nn.Identity(),
|
input_layernorm=layer_norm if for_attn else nn.Identity(),
|
||||||
post_attention_layernorm=nn.Identity() if for_attn else layer_norm,
|
post_attention_layernorm=nn.Identity() if for_attn else layer_norm,
|
||||||
force_layernorm_before_dp_gather=True,
|
force_layernorm_before_dp_gather=True,
|
||||||
|
|||||||
@@ -5414,9 +5414,10 @@ class ServerArgs:
|
|||||||
"fp8",
|
"fp8",
|
||||||
"mxfp8",
|
"mxfp8",
|
||||||
"modelopt_fp4",
|
"modelopt_fp4",
|
||||||
|
"modelopt_mixed",
|
||||||
"nvfp4_online",
|
"nvfp4_online",
|
||||||
None,
|
None,
|
||||||
], f"Invalid quantization '{view.quantization}'. \nFlashInfer TRTLLM routed MOE supports only: 'fp8', 'mxfp8', 'modelopt_fp4', 'nvfp4_online', or bfloat16 (None)."
|
], f"Invalid quantization '{view.quantization}'. \nFlashInfer TRTLLM routed MOE supports only: 'fp8', 'mxfp8', 'modelopt_fp4', 'modelopt_mixed', 'nvfp4_online', or bfloat16 (None)."
|
||||||
|
|
||||||
# The runner-driven shared-experts fusion disables moved to the
|
# The runner-driven shared-experts fusion disables moved to the
|
||||||
# pipeline (arg_groups/overrides.py: _moe_runner_fusion_disable),
|
# pipeline (arg_groups/overrides.py: _moe_runner_fusion_disable),
|
||||||
|
|||||||
Reference in New Issue
Block a user