[MoE Refactor][Test] FusedMoE layer test (#24675)

Signed-off-by: Bill Nell <bnell@redhat.com>
Co-authored-by: Robert Shaw <114415538+robertgshaw2-redhat@users.noreply.github.com>
This commit is contained in:
bnellnm
2026-04-06 13:17:23 -04:00
committed by GitHub
parent bfdc0a3a99
commit f01482408c
6 changed files with 1858 additions and 55 deletions

View File

@@ -0,0 +1,14 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import pytest
def pytest_addoption(parser):
parser.addoption(
"--subtests", action="store", type=str, default=None, help="subtest ids"
)
@pytest.fixture
def subtests(request):
return request.config.getoption("--subtests")

View File

@@ -11,7 +11,11 @@ from torch.multiprocessing import spawn # pyright: ignore[reportPrivateImportUs
from typing_extensions import ParamSpec
from vllm.config import VllmConfig, set_current_vllm_config
from vllm.distributed import init_distributed_environment, initialize_model_parallel
from vllm.distributed import (
cleanup_dist_env_and_memory,
init_distributed_environment,
initialize_model_parallel,
)
from vllm.utils.network_utils import get_open_port
## Parallel Processes Utils
@@ -36,10 +40,17 @@ def _set_vllm_config(
temp_file = tempfile.mkstemp()[1]
# When DP is enabled, processes are organized as:
# rank = dp_rank * tp_pp_world_size + tp_pp_rank
tp_pp_world_size = vllm_config.parallel_config.world_size
vllm_config.parallel_config.data_parallel_rank = rank // tp_pp_world_size
tp_pp_rank = rank % tp_pp_world_size
vllm_config.parallel_config.rank = tp_pp_rank
with set_current_vllm_config(vllm_config):
init_distributed_environment(
world_size=world_size,
rank=rank,
world_size=tp_pp_world_size,
rank=tp_pp_rank,
distributed_init_method=f"file://{temp_file}",
local_rank=local_rank,
backend="nccl",
@@ -59,11 +70,11 @@ def _worker_parallel_launch(
world_local_size: int,
node_rank: int,
init_method: str,
worker: Callable[Concatenate[ProcessGroupInfo, VllmConfig | None, Any, P], None],
worker: Callable[..., None],
vllm_config: VllmConfig | None,
env_dict: dict | None,
*args: P.args,
**kwargs: P.kwargs,
worker_kwargs: dict[str, Any],
*args: Any,
) -> None:
rank = node_rank * world_local_size + local_rank
torch.accelerator.set_device_index(local_rank)
@@ -98,14 +109,17 @@ def _worker_parallel_launch(
vllm_config,
cpu_group,
*args,
**kwargs,
**worker_kwargs,
)
except Exception as ex:
print(ex)
traceback.print_exc()
raise
finally:
torch.distributed.destroy_process_group()
if vllm_config is not None:
cleanup_dist_env_and_memory()
else:
torch.distributed.destroy_process_group()
def parallel_launch_with_config(
@@ -116,7 +130,6 @@ def parallel_launch_with_config(
*args: P.args,
**kwargs: P.kwargs,
) -> None:
assert not kwargs
spawn(
_worker_parallel_launch,
args=(
@@ -127,6 +140,7 @@ def parallel_launch_with_config(
worker,
vllm_config,
env_dict,
kwargs,
)
+ args,
nprocs=world_size,

File diff suppressed because it is too large Load Diff

View File

@@ -248,7 +248,7 @@ def make_quantized_test_activations(
return a, a_q, a_scale
def moe_quantize_weights(
def moe_quantize_weights_2d(
w: torch.Tensor,
w_s: torch.Tensor | None,
quant_dtype: torch.dtype | str | None,
@@ -293,6 +293,40 @@ def moe_quantize_weights(
return w, w_s, w_gs
def moe_quantize_weights(
w: torch.Tensor,
w_s: torch.Tensor | None,
quant_dtype: torch.dtype | str | None,
per_token_quant: bool,
block_shape: list[int] | None,
) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | None]:
assert w.dim() == 3
e, rows, cols = w.shape
w_l = [None] * e
w_s_l = [None] * e
w_gs_l = [None] * e
for idx in range(e):
w_l[idx], w_s_l[idx], w_gs_l[idx] = moe_quantize_weights_2d(
w[idx], None, quant_dtype, per_token_quant, block_shape
)
w = torch.stack(w_l)
w_s = torch.stack(w_s_l)
w_gs = torch.stack(w_gs_l) if e > 0 and w_gs_l[0] is not None else None
if w_s.ndim == 2:
assert w_s.shape[-1] == 1
w_s = w_s.view(-1, 1, 1)
if block_shape is not None:
block_n, block_k = block_shape
n_tiles = (rows + block_n - 1) // block_n
k_tiles = (cols + block_k - 1) // block_k
assert w_s.shape == (e, n_tiles, k_tiles)
return w, w_s, w_gs
def make_test_weight(
e: int,
rows: int,
@@ -303,30 +337,11 @@ def make_test_weight(
per_out_ch_quant: bool = False,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None, torch.Tensor | None]:
w_16 = torch.randn((e, rows, cols), device="cuda", dtype=in_dtype) / 15
w_gs = None
if quant_dtype is not None:
w_l = [None] * e
w_s_l = [None] * e
w_gs_l = [None] * e
for idx in range(e):
w_l[idx], w_s_l[idx], w_gs_l[idx] = moe_quantize_weights(
w_16[idx], None, quant_dtype, per_out_ch_quant, block_shape
)
w = torch.stack(w_l)
w_s = torch.stack(w_s_l)
if e > 0 and w_gs_l[0] is not None:
w_gs = torch.stack(w_gs_l)
if w_s.ndim == 2:
assert w_s.shape[-1] == 1
w_s = w_s.view(-1, 1, 1)
if block_shape is not None:
block_n, block_k = block_shape
n_tiles = (rows + block_n - 1) // block_n
k_tiles = (cols + block_k - 1) // block_k
assert w_s.shape == (e, n_tiles, k_tiles)
w, w_s, w_gs = moe_quantize_weights(
w_16, None, quant_dtype, per_out_ch_quant, block_shape
)
else:
w = w_16
w_s = None
@@ -454,7 +469,6 @@ def fused_moe(
)
# CustomOp?
class BaselineMM(torch.nn.Module):
def __init__(
self,
@@ -462,13 +476,22 @@ class BaselineMM(torch.nn.Module):
out_dtype: torch.dtype,
):
super().__init__()
self.b = b.to(dtype=torch.float32)
self.b = torch.nn.Parameter(b.to(dtype=torch.float32))
self.out_dtype = out_dtype
def forward(self, a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor | None]:
return torch.mm(a.to(dtype=torch.float32), self.b).to(self.out_dtype), None
class BaselineSiluAndMul(torch.nn.Module):
def __init__(self):
super().__init__()
def forward(self, x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
return torch.nn.functional.silu(x[..., :d]) * x[..., d:]
class TestMLP(torch.nn.Module):
def __init__(
self,
@@ -479,7 +502,7 @@ class TestMLP(torch.nn.Module):
super().__init__()
self.gate_up_proj = BaselineMM(w1, out_dtype)
self.down_proj = BaselineMM(w2, out_dtype)
self.act_fn = SiluAndMul()
self.act_fn = BaselineSiluAndMul()
def forward(self, x):
x, _ = self.gate_up_proj(x)
@@ -564,35 +587,24 @@ class RealMLP(torch.nn.Module):
return x
def make_shared_experts(
def make_shared_experts_with_weights(
N: int,
K: int,
in_dtype: torch.dtype = torch.bfloat16,
in_dtype: torch.dtype,
w1: torch.Tensor,
w2: torch.Tensor,
w1_s: torch.Tensor | None = None,
w2_s: torch.Tensor | None = None,
quant_dtype: torch.dtype | str | None = None,
) -> torch.nn.Module:
from vllm.model_executor.layers.quantization.fp8 import Fp8Config
(_, w1, w1_s, _), (_, w2, w2_s, _) = make_test_weights(
1,
N,
K,
in_dtype=in_dtype,
quant_dtype=quant_dtype,
)
old_dtype = torch.get_default_dtype()
try:
torch.set_default_dtype(in_dtype)
if quant_dtype == torch.float8_e4m3fn:
w1 = w1[0].transpose(0, 1)
w2 = w2[0].transpose(0, 1)
w1_s = w1_s[0].transpose(0, 1) if w1_s is not None else None
w2_s = w2_s[0].transpose(0, 1) if w2_s is not None else None
from vllm.model_executor.layers.quantization.fp8 import Fp8Config
quant_config = Fp8Config(True)
else:
w1 = w1[0]
w2 = w2[0]
w1_s = None
w2_s = None
quant_config = None
return RealMLP(K, N, w1, w2, "silu", quant_config, w1_s=w1_s, w2_s=w2_s)
@@ -614,3 +626,22 @@ def modular_triton_fused_moe(
TritonExperts(moe_config, quant_config),
inplace=False,
)
def make_shared_experts(
N: int,
K: int,
in_dtype: torch.dtype = torch.bfloat16,
quant_dtype: torch.dtype | str | None = None,
) -> torch.nn.Module:
(_, w1, w1_s, _), (_, w2, w2_s, _) = make_test_weights(
1,
N,
K,
in_dtype=in_dtype,
quant_dtype=quant_dtype,
)
return make_shared_experts_with_weights(
N, K, in_dtype, w1, w2, w1_s=w1_s, w2_s=w2_s, quant_dtype=quant_dtype
)