2024-01-31 14:34:17 -08:00
|
|
|
"""Tests for the MOE layers.
|
|
|
|
|
|
|
|
Run `pytest tests/kernels/test_moe.py`.
|
|
|
|
"""
|
2024-09-09 23:02:52 -04:00
|
|
|
from typing import List
|
|
|
|
|
2024-01-31 14:34:17 -08:00
|
|
|
import pytest
|
|
|
|
import torch
|
|
|
|
from transformers import MixtralConfig
|
|
|
|
from transformers.models.mixtral.modeling_mixtral import MixtralSparseMoeBlock
|
|
|
|
|
2024-09-25 10:35:52 -04:00
|
|
|
from tests.kernels.utils import opcheck
|
|
|
|
from vllm import _custom_ops as ops
|
2024-01-31 14:34:17 -08:00
|
|
|
from vllm.model_executor.layers.activation import SiluAndMul
|
2024-03-25 23:59:47 +09:00
|
|
|
from vllm.model_executor.layers.fused_moe import fused_moe
|
2024-09-09 23:02:52 -04:00
|
|
|
from vllm.model_executor.layers.fused_moe.fused_marlin_moe import (
|
|
|
|
fused_marlin_moe, single_marlin_moe)
|
2024-09-25 10:35:52 -04:00
|
|
|
from vllm.model_executor.layers.fused_moe.fused_moe import (
|
|
|
|
fused_topk, moe_align_block_size)
|
2024-09-09 23:02:52 -04:00
|
|
|
from vllm.model_executor.layers.quantization.utils.marlin_utils_test import (
|
|
|
|
marlin_quantize)
|
2024-01-31 14:34:17 -08:00
|
|
|
from vllm.model_executor.models.mixtral import MixtralMoE
|
2024-09-09 23:02:52 -04:00
|
|
|
from vllm.scalar_type import scalar_types
|
2024-09-18 18:38:11 +08:00
|
|
|
from vllm.utils import seed_everything
|
2024-01-31 14:34:17 -08:00
|
|
|
|
|
|
|
|
2024-02-05 17:38:02 -08:00
|
|
|
def torch_moe(a, w1, w2, score, topk):
|
2024-01-31 14:34:17 -08:00
|
|
|
B, D = a.shape
|
2024-02-05 17:38:02 -08:00
|
|
|
a = a.view(B, -1, D).repeat(1, topk, 1).reshape(-1, D)
|
|
|
|
out = torch.zeros(B * topk, w2.shape[1], dtype=a.dtype, device=a.device)
|
|
|
|
score = torch.softmax(score, dim=-1, dtype=torch.float32)
|
|
|
|
topk_weight, topk_ids = torch.topk(score, topk)
|
2024-01-31 14:34:17 -08:00
|
|
|
topk_weight = topk_weight.view(-1)
|
2024-02-05 17:38:02 -08:00
|
|
|
topk_ids = topk_ids.view(-1)
|
2024-01-31 14:34:17 -08:00
|
|
|
for i in range(w1.shape[0]):
|
|
|
|
mask = topk_ids == i
|
|
|
|
if mask.sum():
|
|
|
|
out[mask] = SiluAndMul()(
|
|
|
|
a[mask] @ w1[i].transpose(0, 1)) @ w2[i].transpose(0, 1)
|
|
|
|
return (out.view(B, -1, w2.shape[1]) *
|
2024-02-05 17:38:02 -08:00
|
|
|
topk_weight.view(B, -1, 1).to(out.dtype)).sum(dim=1)
|
2024-01-31 14:34:17 -08:00
|
|
|
|
|
|
|
|
2024-09-09 23:02:52 -04:00
|
|
|
def torch_moe_single(a, w, score, topk):
|
|
|
|
B, D = a.shape
|
|
|
|
a = a.view(B, -1, D).repeat(1, topk, 1).reshape(-1, D)
|
|
|
|
out = torch.zeros(B * topk, w.shape[1], dtype=a.dtype, device=a.device)
|
|
|
|
score = torch.softmax(score, dim=-1, dtype=torch.float32)
|
|
|
|
_, topk_ids = torch.topk(score, topk)
|
|
|
|
topk_ids = topk_ids.view(-1)
|
|
|
|
for i in range(w.shape[0]):
|
|
|
|
mask = topk_ids == i
|
|
|
|
if mask.sum():
|
|
|
|
out[mask] = a[mask] @ w[i].transpose(0, 1)
|
|
|
|
return (out.view(B, -1, w.shape[1])).sum(dim=1)
|
|
|
|
|
|
|
|
|
2024-07-02 00:08:29 +03:00
|
|
|
@pytest.mark.parametrize("m", [1024 * 128, 512, 222, 33, 1])
|
2024-01-31 14:34:17 -08:00
|
|
|
@pytest.mark.parametrize("n", [2048, 256, 1024])
|
|
|
|
@pytest.mark.parametrize("k", [128, 511, 1024])
|
|
|
|
@pytest.mark.parametrize("e", [8, 64])
|
|
|
|
@pytest.mark.parametrize("topk", [2, 6])
|
|
|
|
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
|
|
|
def test_fused_moe(
|
|
|
|
m: int,
|
|
|
|
n: int,
|
|
|
|
k: int,
|
|
|
|
e: int,
|
|
|
|
topk: int,
|
|
|
|
dtype: torch.dtype,
|
|
|
|
):
|
2024-09-09 23:02:52 -04:00
|
|
|
a = torch.randn((m, k), device="cuda", dtype=dtype) / 10
|
|
|
|
w1 = torch.randn((e, 2 * n, k), device="cuda", dtype=dtype) / 10
|
|
|
|
w2 = torch.randn((e, k, n), device="cuda", dtype=dtype) / 10
|
2024-01-31 14:34:17 -08:00
|
|
|
|
2024-09-09 23:02:52 -04:00
|
|
|
score = torch.randn((m, e), device="cuda", dtype=dtype)
|
2024-02-05 17:38:02 -08:00
|
|
|
triton_output = fused_moe(a, w1, w2, score, topk, renormalize=False)
|
|
|
|
torch_output = torch_moe(a, w1, w2, score, topk)
|
2024-08-15 21:24:04 -07:00
|
|
|
torch.testing.assert_close(triton_output, torch_output, atol=1e-2, rtol=0)
|
2024-01-31 14:34:17 -08:00
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.parametrize("dtype",
|
|
|
|
[torch.float32, torch.float16, torch.bfloat16])
|
|
|
|
@torch.inference_mode()
|
|
|
|
def test_mixtral_moe(dtype: torch.dtype):
|
2024-03-10 19:49:14 -07:00
|
|
|
"""Make sure our Mixtral MoE implementation agrees with the one from
|
|
|
|
huggingface."""
|
2024-01-31 14:34:17 -08:00
|
|
|
|
|
|
|
# Instantiate our and huggingface's MoE blocks
|
|
|
|
config = MixtralConfig()
|
|
|
|
hf_moe = MixtralSparseMoeBlock(config).to(dtype).to("cuda")
|
|
|
|
vllm_moe = MixtralMoE(
|
|
|
|
num_experts=config.num_local_experts,
|
|
|
|
top_k=config.num_experts_per_tok,
|
|
|
|
hidden_size=config.hidden_size,
|
|
|
|
intermediate_size=config.intermediate_size,
|
|
|
|
params_dtype=dtype,
|
|
|
|
tp_size=1,
|
2024-02-05 17:38:02 -08:00
|
|
|
).cuda()
|
2024-01-31 14:34:17 -08:00
|
|
|
|
|
|
|
# Load the weights
|
2024-04-11 13:35:51 -07:00
|
|
|
vllm_moe.gate.weight.data[:] = hf_moe.gate.weight.data
|
2024-01-31 14:34:17 -08:00
|
|
|
for i in range(config.num_local_experts):
|
|
|
|
weights = (hf_moe.experts[i].w1.weight.data,
|
|
|
|
hf_moe.experts[i].w3.weight.data)
|
2024-07-02 17:54:35 -04:00
|
|
|
vllm_moe.experts.w13_weight[i][:] = torch.cat(weights, dim=0)
|
|
|
|
vllm_moe.experts.w2_weight[i][:] = hf_moe.experts[i].w2.weight.data
|
2024-01-31 14:34:17 -08:00
|
|
|
|
|
|
|
# Generate input batch of dimensions [batch_size, seq_len, hidden_dim]
|
2024-03-24 16:00:16 -07:00
|
|
|
hf_inputs = torch.randn((1, 64, config.hidden_size)).to(dtype).to("cuda")
|
|
|
|
# vLLM uses 1D query [num_tokens, hidden_dim]
|
|
|
|
vllm_inputs = hf_inputs.flatten(0, 1)
|
2024-01-31 14:34:17 -08:00
|
|
|
|
|
|
|
# Run forward passes for both MoE blocks
|
2024-03-24 16:00:16 -07:00
|
|
|
hf_states, _ = hf_moe.forward(hf_inputs)
|
|
|
|
vllm_states = vllm_moe.forward(vllm_inputs)
|
2024-01-31 14:34:17 -08:00
|
|
|
|
|
|
|
mixtral_moe_tol = {
|
|
|
|
torch.float32: 1e-3,
|
|
|
|
torch.float16: 1e-3,
|
|
|
|
torch.bfloat16: 1e-2,
|
|
|
|
}
|
|
|
|
|
2024-08-15 21:24:04 -07:00
|
|
|
torch.testing.assert_close(hf_states.flatten(0, 1),
|
|
|
|
vllm_states,
|
|
|
|
rtol=mixtral_moe_tol[dtype],
|
|
|
|
atol=mixtral_moe_tol[dtype])
|
2024-09-09 23:02:52 -04:00
|
|
|
|
|
|
|
|
|
|
|
def stack_and_dev(tensors: List[torch.Tensor]):
|
|
|
|
dev = tensors[0].device
|
|
|
|
return torch.stack(tensors, dim=0).to(dev)
|
|
|
|
|
|
|
|
|
|
|
|
def compute_max_diff(output, output_ref):
|
|
|
|
return torch.mean(torch.abs(output - output_ref)) / torch.mean(
|
|
|
|
torch.abs(output_ref))
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.parametrize("m", [64, 512, 222, 33, 1])
|
|
|
|
@pytest.mark.parametrize("n", [128, 2048, 256, 1024])
|
|
|
|
@pytest.mark.parametrize("k", [128, 1024, 512])
|
|
|
|
@pytest.mark.parametrize("e", [4, 8, 64])
|
|
|
|
@pytest.mark.parametrize("topk", [2, 6])
|
|
|
|
@pytest.mark.parametrize("group_size", [-1, 32, 64, 128])
|
|
|
|
@pytest.mark.parametrize("act_order", [True, False])
|
2024-09-16 17:47:19 +02:00
|
|
|
@pytest.mark.parametrize("num_bits", [4, 8])
|
2024-09-29 03:19:40 +02:00
|
|
|
@pytest.mark.parametrize("is_k_full", [True, False])
|
2024-09-09 23:02:52 -04:00
|
|
|
def test_fused_marlin_moe(
|
|
|
|
m: int,
|
|
|
|
n: int,
|
|
|
|
k: int,
|
|
|
|
e: int,
|
|
|
|
topk: int,
|
|
|
|
group_size: int,
|
|
|
|
act_order: bool,
|
2024-09-16 17:47:19 +02:00
|
|
|
num_bits: int,
|
2024-09-29 03:19:40 +02:00
|
|
|
is_k_full: bool,
|
2024-09-09 23:02:52 -04:00
|
|
|
):
|
2024-09-18 18:38:11 +08:00
|
|
|
seed_everything(7)
|
2024-09-09 23:02:52 -04:00
|
|
|
|
|
|
|
if topk > e:
|
|
|
|
return
|
|
|
|
|
|
|
|
# Filter act_order
|
|
|
|
if act_order:
|
|
|
|
if group_size == -1:
|
|
|
|
return
|
|
|
|
if group_size in (k, n):
|
|
|
|
return
|
2024-09-29 03:19:40 +02:00
|
|
|
else:
|
|
|
|
if not is_k_full:
|
|
|
|
return
|
2024-09-09 23:02:52 -04:00
|
|
|
|
2024-09-16 17:47:19 +02:00
|
|
|
quant_type = (scalar_types.uint4b8
|
|
|
|
if num_bits == 4 else scalar_types.uint8b128)
|
2024-09-09 23:02:52 -04:00
|
|
|
dtype = torch.float16
|
|
|
|
a = torch.randn((m, k), device="cuda", dtype=dtype) / 10
|
|
|
|
w1 = torch.randn((e, 2 * n, k), device="cuda", dtype=dtype) / 10
|
|
|
|
w2 = torch.randn((e, k, n), device="cuda", dtype=dtype) / 10
|
|
|
|
|
|
|
|
w_ref1_l = []
|
|
|
|
qweight1_l = []
|
|
|
|
scales1_l = []
|
|
|
|
g_idx1_l = []
|
|
|
|
sort_indices1_l = []
|
|
|
|
|
|
|
|
for i in range(w1.shape[0]):
|
|
|
|
test_perm = torch.randperm(k)
|
|
|
|
w_ref1, qweight1, scales1, g_idx1, sort_indices1, _ = marlin_quantize(
|
|
|
|
w1[i].transpose(1, 0), quant_type, group_size, act_order,
|
|
|
|
test_perm)
|
|
|
|
w_ref1_l.append(w_ref1)
|
|
|
|
qweight1_l.append(qweight1)
|
|
|
|
scales1_l.append(scales1)
|
|
|
|
g_idx1_l.append(g_idx1)
|
|
|
|
sort_indices1_l.append(sort_indices1)
|
|
|
|
|
|
|
|
w_ref1 = stack_and_dev(w_ref1_l)
|
|
|
|
qweight1 = stack_and_dev(qweight1_l).contiguous()
|
|
|
|
scales1 = stack_and_dev(scales1_l)
|
|
|
|
g_idx1 = stack_and_dev(g_idx1_l)
|
|
|
|
sort_indices1 = stack_and_dev(sort_indices1_l)
|
|
|
|
|
|
|
|
w_ref2_l = []
|
|
|
|
qweight2_l = []
|
|
|
|
scales2_l = []
|
|
|
|
g_idx2_l = []
|
|
|
|
sort_indices2_l = []
|
|
|
|
|
|
|
|
for i in range(w2.shape[0]):
|
|
|
|
test_perm = torch.randperm(n)
|
|
|
|
w_ref2, qweight2, scales2, g_idx2, sort_indices2, _ = marlin_quantize(
|
|
|
|
w2[i].transpose(1, 0), quant_type, group_size, act_order,
|
|
|
|
test_perm)
|
|
|
|
w_ref2_l.append(w_ref2)
|
|
|
|
qweight2_l.append(qweight2)
|
|
|
|
scales2_l.append(scales2)
|
|
|
|
g_idx2_l.append(g_idx2)
|
|
|
|
sort_indices2_l.append(sort_indices2)
|
|
|
|
|
|
|
|
w_ref2 = stack_and_dev(w_ref2_l)
|
|
|
|
qweight2 = stack_and_dev(qweight2_l).contiguous()
|
|
|
|
scales2 = stack_and_dev(scales2_l)
|
|
|
|
g_idx2 = stack_and_dev(g_idx2_l)
|
|
|
|
sort_indices2 = stack_and_dev(sort_indices2_l)
|
|
|
|
|
|
|
|
score = torch.randn((m, e), device="cuda", dtype=dtype)
|
|
|
|
|
|
|
|
topk_weights, topk_ids = fused_topk(a, score, topk, False)
|
|
|
|
|
|
|
|
triton_output = fused_moe(
|
|
|
|
a,
|
|
|
|
w_ref1.transpose(1, 2).contiguous(),
|
|
|
|
w_ref2.transpose(1, 2).contiguous(),
|
|
|
|
score,
|
|
|
|
topk,
|
|
|
|
renormalize=False,
|
|
|
|
)
|
|
|
|
marlin_output = fused_marlin_moe(
|
|
|
|
a,
|
|
|
|
qweight1,
|
|
|
|
qweight2,
|
|
|
|
score,
|
|
|
|
g_idx1,
|
|
|
|
g_idx2,
|
|
|
|
sort_indices1,
|
|
|
|
sort_indices2,
|
|
|
|
topk_weights,
|
|
|
|
topk_ids,
|
|
|
|
w1_scale=scales1,
|
|
|
|
w2_scale=scales2,
|
2024-09-16 17:47:19 +02:00
|
|
|
num_bits=num_bits,
|
2024-09-29 03:19:40 +02:00
|
|
|
is_k_full=is_k_full,
|
2024-09-09 23:02:52 -04:00
|
|
|
)
|
|
|
|
|
|
|
|
assert compute_max_diff(marlin_output, triton_output) < 4e-2
|
|
|
|
|
2024-09-25 10:35:52 -04:00
|
|
|
if ops.supports_moe_ops:
|
|
|
|
token_expert_indicies = torch.empty(m,
|
|
|
|
topk,
|
|
|
|
dtype=torch.int32,
|
|
|
|
device=a.device)
|
|
|
|
|
|
|
|
opcheck(torch.ops._moe_C.topk_softmax, (
|
|
|
|
topk_weights,
|
|
|
|
topk_ids,
|
|
|
|
token_expert_indicies,
|
|
|
|
score.float(),
|
|
|
|
))
|
|
|
|
|
|
|
|
block_size_m = 4
|
|
|
|
|
|
|
|
sorted_token_ids, _, _ = moe_align_block_size(topk_ids, block_size_m,
|
|
|
|
e)
|
|
|
|
|
|
|
|
max_workspace_size = ((m + 255) // 256) * (max(2 * n, k) // 64) * 16
|
|
|
|
workspace = torch.zeros(max_workspace_size,
|
|
|
|
dtype=torch.int,
|
|
|
|
device="cuda",
|
|
|
|
requires_grad=False)
|
|
|
|
|
|
|
|
opcheck(torch.ops._moe_C.marlin_gemm_moe,
|
|
|
|
(a, qweight1, sorted_token_ids, topk_weights, topk_ids,
|
|
|
|
scales1, g_idx1, sort_indices1, workspace, quant_type, m,
|
|
|
|
2 * n, k, True, e, topk, block_size_m, True, False))
|
|
|
|
|
2024-09-09 23:02:52 -04:00
|
|
|
|
|
|
|
@pytest.mark.skip("This test is here for the sake of debugging, "
|
|
|
|
"don't run it in automated tests.")
|
|
|
|
@pytest.mark.parametrize("m", [64, 512, 222, 33, 1])
|
|
|
|
@pytest.mark.parametrize("n", [128, 2048, 256, 1024])
|
|
|
|
@pytest.mark.parametrize("k", [128, 1024, 512])
|
|
|
|
@pytest.mark.parametrize("e", [4, 8, 64])
|
|
|
|
@pytest.mark.parametrize("topk", [2, 6])
|
|
|
|
@pytest.mark.parametrize("group_size", [-1, 32, 64, 128])
|
|
|
|
@pytest.mark.parametrize("act_order", [True, False])
|
2024-09-16 17:47:19 +02:00
|
|
|
@pytest.mark.parametrize("num_bits", [4, 8])
|
2024-09-29 03:19:40 +02:00
|
|
|
@pytest.mark.parametrize("is_k_full", [True, False])
|
2024-09-16 17:47:19 +02:00
|
|
|
def test_single_marlin_moe_multiply(
|
2024-09-09 23:02:52 -04:00
|
|
|
m: int,
|
|
|
|
n: int,
|
|
|
|
k: int,
|
|
|
|
e: int,
|
|
|
|
topk: int,
|
|
|
|
group_size: int,
|
|
|
|
act_order: bool,
|
2024-09-16 17:47:19 +02:00
|
|
|
num_bits: int,
|
2024-09-29 03:19:40 +02:00
|
|
|
is_k_full: bool,
|
2024-09-09 23:02:52 -04:00
|
|
|
):
|
|
|
|
if topk > e:
|
|
|
|
return
|
|
|
|
|
|
|
|
# Filter act_order
|
|
|
|
if act_order:
|
|
|
|
if group_size == -1:
|
|
|
|
return
|
|
|
|
if group_size == k:
|
|
|
|
return
|
2024-09-29 03:19:40 +02:00
|
|
|
else:
|
|
|
|
if not is_k_full:
|
|
|
|
return
|
2024-09-09 23:02:52 -04:00
|
|
|
|
2024-09-16 17:47:19 +02:00
|
|
|
quant_type = (scalar_types.uint4b8
|
|
|
|
if num_bits == 4 else scalar_types.uint8b128)
|
2024-09-09 23:02:52 -04:00
|
|
|
dtype = torch.float16
|
|
|
|
a = torch.randn((m, k), device="cuda", dtype=dtype) / 10
|
|
|
|
w = torch.randn((e, n, k), device="cuda", dtype=dtype) / 10
|
|
|
|
|
|
|
|
w_ref_l = []
|
|
|
|
qweights_l = []
|
|
|
|
scales_l = []
|
|
|
|
g_idx_l = []
|
|
|
|
sort_indices_l = []
|
|
|
|
|
|
|
|
for i in range(w.shape[0]):
|
|
|
|
test_perm = torch.randperm(k)
|
|
|
|
w_ref, qweight, scales, g_idx, sort_indices, _ = marlin_quantize(
|
|
|
|
w[i].transpose(1, 0), quant_type, group_size, act_order, test_perm)
|
|
|
|
w_ref_l.append(w_ref)
|
|
|
|
qweights_l.append(qweight)
|
|
|
|
scales_l.append(scales)
|
|
|
|
g_idx_l.append(g_idx)
|
|
|
|
sort_indices_l.append(sort_indices)
|
|
|
|
|
|
|
|
w_ref = stack_and_dev(w_ref_l)
|
|
|
|
qweight = stack_and_dev(qweights_l).contiguous()
|
|
|
|
scales = stack_and_dev(scales_l)
|
|
|
|
g_idx = stack_and_dev(g_idx_l)
|
|
|
|
sort_indices = stack_and_dev(sort_indices_l)
|
|
|
|
|
|
|
|
score = torch.randn((m, e), device="cuda", dtype=dtype)
|
2024-09-29 03:19:40 +02:00
|
|
|
marlin_output = single_marlin_moe(
|
|
|
|
a,
|
|
|
|
qweight,
|
|
|
|
scales,
|
|
|
|
score,
|
|
|
|
g_idx,
|
|
|
|
sort_indices,
|
|
|
|
topk,
|
|
|
|
renormalize=False,
|
|
|
|
num_bits=num_bits,
|
|
|
|
is_k_full=is_k_full,
|
|
|
|
)
|
2024-09-09 23:02:52 -04:00
|
|
|
torch_output = torch_moe_single(a, w_ref.transpose(1, 2), score, topk)
|
|
|
|
|
|
|
|
assert compute_max_diff(marlin_output, torch_output) < 1e-2
|
2024-09-25 10:35:52 -04:00
|
|
|
|
|
|
|
|
|
|
|
def test_moe_align_block_size_opcheck():
|
|
|
|
num_experts = 4
|
|
|
|
block_size = 4
|
|
|
|
topk_ids = torch.randint(0,
|
|
|
|
num_experts, (3, 4),
|
|
|
|
dtype=torch.int32,
|
|
|
|
device='cuda')
|
|
|
|
|
|
|
|
max_num_tokens_padded = topk_ids.numel() + num_experts * (block_size - 1)
|
|
|
|
sorted_ids = torch.empty((max_num_tokens_padded, ),
|
|
|
|
dtype=torch.int32,
|
|
|
|
device=topk_ids.device)
|
|
|
|
sorted_ids.fill_(topk_ids.numel())
|
|
|
|
max_num_m_blocks = max_num_tokens_padded // block_size
|
|
|
|
expert_ids = torch.empty((max_num_m_blocks, ),
|
|
|
|
dtype=torch.int32,
|
|
|
|
device=topk_ids.device)
|
|
|
|
num_tokens_post_pad = torch.empty((1),
|
|
|
|
dtype=torch.int32,
|
|
|
|
device=topk_ids.device)
|
|
|
|
|
|
|
|
opcheck(torch.ops._C.moe_align_block_size,
|
|
|
|
(topk_ids, num_experts, block_size, sorted_ids, expert_ids,
|
|
|
|
num_tokens_post_pad))
|