[AMD] Fix garbled unquantized Qwen3-30B-A3B output on ROCm/aiter where the aiter CK fused-MoE falls back to Triton with pre-shuffled weights (#28244)

Co-authored-by: Xinyu Jiang <xinyuj2@andrew.cmu.edu>
This commit is contained in:
Zhiyao Jiang
2026-06-20 01:25:01 -07:00
committed by GitHub
co-authored by Xinyu Jiang
parent c1416bb3ee
commit 1115373668
2 changed files with 265 additions and 18 deletions
@@ -238,8 +238,13 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
set_weight_attrs(w2_weight_bias, extra_weight_attrs)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
_should_use_aiter_moe = _use_aiter and (
get_moe_runner_backend().is_auto() or get_moe_runner_backend().is_aiter()
_should_use_aiter_moe = (
_use_aiter
and (
get_moe_runner_backend().is_auto()
or get_moe_runner_backend().is_aiter()
)
and self._aiter_ck_moe_supported(layer)
)
if _should_use_aiter_moe:
copy_or_rebind_param(
@@ -392,6 +397,11 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
param.data = param.data.reshape(expected_shape)
def _aiter_ck_moe_supported(self, layer) -> bool:
# aiter CK fused-MoE requires intermediate_size_per_partition to be 128-aligned
# (GemmSpec=Default; otherwise CK raises "not support this GEMM problem").
return layer.intermediate_size_per_partition % 128 == 0
def create_moe_runner(
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
):
@@ -410,7 +420,7 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
backend = MoeRunnerBackend.TRITON
self.runner = MoeRunner(backend, moe_runner_config)
# Separate runner so CK-shape errors fall back to self.runner on every call.
# aiter CK fused-MoE only supports 128-aligned shapes; otherwise use triton.
self._aiter_runner: Optional[MoeRunner] = None
if (
_use_aiter
@@ -420,7 +430,22 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
)
and get_moe_a2a_backend().supports_aiter()
):
self._aiter_runner = MoeRunner(MoeRunnerBackend.AITER, moe_runner_config)
if self._aiter_ck_moe_supported(layer):
self._aiter_runner = MoeRunner(
MoeRunnerBackend.AITER, moe_runner_config
)
elif get_moe_runner_backend().is_aiter():
raise ValueError(
"moe_runner_backend=aiter is not supported for "
f"intermediate_size_per_partition={layer.intermediate_size_per_partition}; "
"use --moe-runner-backend triton."
)
else:
logger.warning_once(
"aiter CK fused-MoE does not support "
f"intermediate_size_per_partition={layer.intermediate_size_per_partition}; "
"using triton MoE runner."
)
@property
def load_up_proj_weight_first(self) -> bool:
@@ -522,20 +547,12 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
AiterMoeQuantInfo,
)
try:
quant_info = AiterMoeQuantInfo(
w13_weight=layer.w13_weight,
w2_weight=layer.w2_weight,
expert_mask=layer.dispatcher.expert_mask_gpu,
)
return self._aiter_runner.run(dispatch_output, quant_info)
except RuntimeError as e:
# AITER CK fused_moe may not support all GEMM dimensions
# (e.g. Gemma4 MoE with 128 experts x 704 intermediate size)
logger.warning_once(
f"AITER CK fused_moe failed ({e}), "
"falling back to Triton MoE runner."
)
quant_info = AiterMoeQuantInfo(
w13_weight=layer.w13_weight,
w2_weight=layer.w2_weight,
expert_mask=layer.dispatcher.expert_mask_gpu,
)
return self._aiter_runner.run(dispatch_output, quant_info)
quant_info = TritonMoeQuantInfo(
w13_weight=layer.w13_weight,