Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions tests/models/quantization/test_gpt_oss.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,11 @@ def test_gpt_oss_attention_quantization(

model_args = EvaluationConfig(model_name).get_model_args(tp_size)

# Emulation backend on MI300, MI250 is opt-in
# following https://github.com/vllm-project/vllm/pull/45896
if not on_gfx950():
model_args["moe_backend"] = "emulation"

Comment on lines +107 to +111

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could we put this logic into the oracle for a single source of truth?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Well, this used to be the case (automatic selection in the oracle on Instinct), but @BowenBao made this backend an opt-in https://github.com/vllm-project/vllm/pull/41436/changes#r3237495909 - with the motivation being that these backend are not super well optimized (e.g. unfused weight dequant / compute), so it really be a last resort. cc @BowenBao would it make sense to add this default back, as the last priority?

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Based on the number of recent changes towards this direction, I'm not strongly against it. If we are to change it let's put out an explicit warning.

extra_run_kwargs = {
"gen_kwargs": {"max_gen_toks": 8000},
"apply_chat_template": True,
Expand Down
12 changes: 12 additions & 0 deletions tests/quantization/test_quark.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,14 @@
)
from vllm.platforms import current_platform

if current_platform.is_rocm():
from vllm.platforms.rocm import on_gfx950
else:

def on_gfx950() -> bool:
return False


from .reference_mxfp4 import dq_mxfp4_torch, qdq_mxfp4_torch

# Minimum amd-quark version for MXFP4/OCP_MX tests (single source of truth).
Expand Down Expand Up @@ -213,6 +221,10 @@ def get_model_args(
if model_max_len is not None:
model_args["max_model_len"] = model_max_len

# Emulation backend on MI300, MI250 is opt-in following https://github.com/vllm-project/vllm/pull/45896
if not on_gfx950():
model_args["moe_backend"] = "emulation"

return model_args


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
from vllm.model_executor.layers.quantization.utils.ocp_mx_utils import (
OCP_MX_Scheme,
)
from vllm.platforms import current_platform

logger = init_logger(__name__)

Expand Down Expand Up @@ -83,8 +84,7 @@ def __init__(
OCP_MX_Scheme.w_mxfp4_a_fp8,
OCP_MX_Scheme.w_mxfp6_e3m2_a_fp8,
]:
# TODO: double check this one
self._quant_dtype = "mxfp8"
self._quant_dtype = current_platform.fp8_dtype()

@property
def quant_dtype(self) -> torch.dtype | str | None:
Expand Down
13 changes: 12 additions & 1 deletion vllm/model_executor/layers/fused_moe/experts/triton_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -332,12 +332,23 @@ def apply(
else:
lora_x = hidden_states

# TODO: The fallback to self.a1_scale was added for deferred static
# activation quantization in https://github.com/vllm-project/vllm/pull/40857.
# Activation emulation relies solely on `a1q_scale` output of
# `moe_kernel_quantize_input` - this should be adapted to
# always solely rely on `a1q_scale`.
input_scale = (
a1q_scale
if self.quantization_emulation
else (a1q_scale if a1q_scale is not None else self.a1_scale)
)

def _base_w13_fn():
invoke_fused_moe_triton_kernel(
hidden_states,
w1,
intermediate_cache1,
a1q_scale if a1q_scale is not None else self.a1_scale,
input_scale,
self.w1_scale,
None, # topk_weights
sorted_token_ids,
Expand Down
27 changes: 25 additions & 2 deletions vllm/model_executor/layers/fused_moe/oracle/mxfp4.py
Original file line number Diff line number Diff line change
Expand Up @@ -674,6 +674,8 @@ def convert_gpt_oss_weight_to_mxfp4_moe_kernel_format(
w2_weight_scale: torch.Tensor,
w13_bias: torch.Tensor | None = None,
w2_bias: torch.Tensor | None = None,
w13_input_scale: torch.Tensor | None = None,
w2_input_scale: torch.Tensor | None = None,
_cache_permute_indices: dict[torch.Size, torch.Tensor] | None = None,
) -> tuple[
torch.Tensor,
Expand Down Expand Up @@ -1191,8 +1193,29 @@ def swap_every_two_rows(x, axis=-1):
w2_bias,
)
elif mxfp4_backend == Mxfp4MoeBackend.EMULATION:
# No additional transformation needed for emulation backend,
# weights are dequantized on the fly in the experts class.
w13_has_per_expert_scale = (
w13_input_scale is not None
and w13_input_scale.ndim == 1
and not all_close_1d(w13_input_scale)
)
w2_has_per_expert_scale = (
w2_input_scale is not None
and w2_input_scale.ndim == 1
and not all_close_1d(w2_input_scale)
)
if w13_has_per_expert_scale or w2_has_per_expert_scale:
logger.warning_once(
"Found input_scales that are not equal for OCP MX MoE "
"emulation. Using the maximum across experts for each layer."
)
if w13_input_scale is not None:
layer.w13_input_scale = torch.nn.Parameter(
w13_input_scale.max().to(torch.float32), requires_grad=False
)
if w2_input_scale is not None:
layer.w2_input_scale = torch.nn.Parameter(
w2_input_scale.max().to(torch.float32), requires_grad=False
)
Comment on lines +1211 to +1218

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

is this not handled in quark quant method class already?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This method convert_gpt_oss_weight_to_mxfp4_moe_kernel_format is called from QuarkOCP_MX_MoEMethod.process_weights_after_loading indeed. It seems to be the one handling the per-backend necessary pre-processing.

See:

convert_gpt_oss_weight_to_mxfp4_moe_kernel_format(

return (
w13_weight,
w2_weight,
Expand Down
28 changes: 16 additions & 12 deletions vllm/model_executor/layers/fused_moe/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -200,6 +200,16 @@ def _mxfp4_quantize(
return A, None


def _fp8_quantize_dequantize(
A: torch.Tensor,
A_scale: torch.Tensor,
):
qA, qA_scale = ops.scaled_fp8_quant(A, A_scale, use_per_token_if_dynamic=False)
A = per_tensor_dequantize(qA, qA_scale).to(A.dtype)

return A, None


def _mxfp8_e4m3_quantize(
A: torch.Tensor,
A_scale: torch.Tensor | None,
Expand Down Expand Up @@ -268,23 +278,17 @@ def moe_kernel_quantize_input(
# purpose, because there is no native kernel for weight in ocp_mx_scheme
# and activation in FP8. The implementation is based on existing
# non-emulation ops.
qA, qA_scale = ops.scaled_fp8_quant(
A, A_scale, use_per_token_if_dynamic=False
)
A = per_tensor_dequantize(qA, qA_scale).to(A.dtype)
# After QDQ, we don't need further quantization
return A, None
# TODO: Remove this `ocp_mx_scheme is not None` block and rely solely
# on `quantization_emulation`.
return _fp8_quantize_dequantize(A, A_scale)
# else: For other schemes (e.g., *_a_mxfp6_e3m2, *_a_mxfp6_e2m3),
# weights are already dequantized, and we proceed with normal
# activation quantization below.

if quant_dtype == current_platform.fp8_dtype():
if quantization_emulation:
raise NotImplementedError(
f"moe_kernel_quantize_input does not support quant_dtype={quant_dtype}"
" MOE quantization emulation. Please open an issue."
)
return _fp8_quantize(A, A_scale, per_act_token_quant, block_shape)
return _fp8_quantize_dequantize(A, A_scale)
else:
return _fp8_quantize(A, A_scale, per_act_token_quant, block_shape)
elif quant_dtype == torch.int8:
if quantization_emulation:
raise NotImplementedError(
Expand Down
2 changes: 2 additions & 0 deletions vllm/model_executor/layers/quantization/quark/quark_moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -1211,6 +1211,8 @@ def _setup_kernel(self, layer: RoutedExperts):
w2_weight_scale=layer.w2_weight_scale,
w13_bias=w13_bias,
w2_bias=w2_bias,
w13_input_scale=layer.w13_input_scale,
w2_input_scale=layer.w2_input_scale,
Comment on lines +1214 to +1215

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

we had eval test coverage for w4a8 actually
https://github.com/vllm-project/vllm/blob/main/tests/evals/gpt_oss/configs/gpt-oss-20b-rocm-quark-mxfp4-fp8-triton.yaml
wonder why it does not catch these missed scales.

btw we should add emulation backend to these eval tests too, can be separate PRs

)
)

Expand Down
Loading