[Test] Replace NVFP4 MoE runner backend e2e matrix with a layer-level unit test (#33611)
This commit is contained in:
@@ -1,18 +1,7 @@
|
|||||||
"""Backend tests for CuteDSL MoE (FusedMoE + moe_runner, moe_a2a=none).
|
"""CuteDSL MoE e2e with EP=TP=4 (moe_a2a=none): each GPU holds 1/4 of the
|
||||||
|
experts at full intermediate width, partial results combined via all-reduce.
|
||||||
Exercises the CuteDSL moe_runner path with ModelOpt FP4 by launching a
|
Kept for the distributed-EP dimension; single-GPU backend numerics live in
|
||||||
server with --moe-runner-backend flashinfer_cutedsl.
|
unit/layers/quantization/test_nvfp4_moe_backends.py."""
|
||||||
|
|
||||||
Two configurations are tested:
|
|
||||||
- EP=1, TP=4: each GPU holds all experts with TP-sharded intermediate dim
|
|
||||||
- EP=4, TP=4: each GPU holds 1/4 of experts at full intermediate width,
|
|
||||||
partial results combined via all-reduce (no A2A dispatch)
|
|
||||||
|
|
||||||
Requires 4 GPUs. Run from repo root with:
|
|
||||||
python -m pytest test/registered/backends/test_deepseek_v3_fp4_cutedsl_moe.py -v -s
|
|
||||||
Or via the nightly suite:
|
|
||||||
python test/run_suite.py --hw cuda --suite nightly-4-gpu-b200 --nightly
|
|
||||||
"""
|
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
@@ -28,68 +17,13 @@ from sglang.test.test_utils import (
|
|||||||
write_github_step_summary,
|
write_github_step_summary,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=900, suite="nightly-4-gpu-b200", nightly=True)
|
register_cuda_ci(est_time=450, suite="nightly-4-gpu-b200", nightly=True)
|
||||||
|
|
||||||
FULL_DEEPSEEK_V3_FP4_MODEL_PATH = "nvidia/DeepSeek-V3-0324-FP4"
|
FULL_DEEPSEEK_V3_FP4_MODEL_PATH = "nvidia/DeepSeek-V3-0324-FP4"
|
||||||
SERVER_LAUNCH_TIMEOUT = 1000
|
SERVER_LAUNCH_TIMEOUT = 1000
|
||||||
GSM8K_ACCURACY_THRESHOLD = 0.935
|
GSM8K_ACCURACY_THRESHOLD = 0.935
|
||||||
|
|
||||||
|
|
||||||
class TestDeepseekV3FP4CuteDSLMoE(CustomTestCase):
|
|
||||||
"""CuteDSL standard moe_runner path: flashinfer_cutedsl + modelopt_fp4, EP=1."""
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def setUpClass(cls):
|
|
||||||
cls.model = FULL_DEEPSEEK_V3_FP4_MODEL_PATH
|
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
|
||||||
other_args = [
|
|
||||||
"--tp",
|
|
||||||
"4",
|
|
||||||
"--ep",
|
|
||||||
"1",
|
|
||||||
"--mem-fraction-static",
|
|
||||||
"0.75",
|
|
||||||
"--attention-backend",
|
|
||||||
"trtllm_mla",
|
|
||||||
"--moe-runner-backend",
|
|
||||||
"flashinfer_cutedsl",
|
|
||||||
"--quantization",
|
|
||||||
"modelopt_fp4",
|
|
||||||
"--model-loader-extra-config",
|
|
||||||
'{"enable_multithread_load": true}',
|
|
||||||
]
|
|
||||||
cls.process = popen_launch_server(
|
|
||||||
cls.model,
|
|
||||||
cls.base_url,
|
|
||||||
timeout=SERVER_LAUNCH_TIMEOUT,
|
|
||||||
other_args=other_args,
|
|
||||||
)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def tearDownClass(cls):
|
|
||||||
kill_process_tree(cls.process.pid)
|
|
||||||
|
|
||||||
def test_a_gsm8k(
|
|
||||||
self,
|
|
||||||
): # Append an "a" to make this test run first (alphabetically) to warm up the server
|
|
||||||
args = SimpleNamespace(
|
|
||||||
num_shots=8,
|
|
||||||
data_path=None,
|
|
||||||
num_questions=1319,
|
|
||||||
parallel=1319,
|
|
||||||
max_new_tokens=512,
|
|
||||||
host="http://127.0.0.1",
|
|
||||||
port=int(self.base_url.split(":")[-1]),
|
|
||||||
)
|
|
||||||
metrics = run_eval_few_shot_gsm8k(args)
|
|
||||||
if is_in_ci():
|
|
||||||
write_github_step_summary(
|
|
||||||
f"### test_gsm8k (deepseek-v3-fp4-cutedsl-moe)\n"
|
|
||||||
f'{metrics["accuracy"]=:.3f}\n'
|
|
||||||
)
|
|
||||||
self.assertGreater(metrics["accuracy"], GSM8K_ACCURACY_THRESHOLD)
|
|
||||||
|
|
||||||
|
|
||||||
class TestDeepseekV3FP4CuteDSLMoEEP4(CustomTestCase):
|
class TestDeepseekV3FP4CuteDSLMoEEP4(CustomTestCase):
|
||||||
"""CuteDSL standard moe_runner path: flashinfer_cutedsl + modelopt_fp4, EP=TP=4."""
|
"""CuteDSL standard moe_runner path: flashinfer_cutedsl + modelopt_fp4, EP=TP=4."""
|
||||||
|
|
||||||
|
|||||||
@@ -1,76 +0,0 @@
|
|||||||
import unittest
|
|
||||||
from types import SimpleNamespace
|
|
||||||
|
|
||||||
from sglang.srt.utils import kill_process_tree
|
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
|
||||||
from sglang.test.run_eval import run_eval
|
|
||||||
from sglang.test.test_utils import (
|
|
||||||
DEFAULT_URL_FOR_TEST,
|
|
||||||
CustomTestCase,
|
|
||||||
is_in_ci,
|
|
||||||
popen_launch_server,
|
|
||||||
write_github_step_summary,
|
|
||||||
)
|
|
||||||
|
|
||||||
register_cuda_ci(est_time=900, suite="nightly-4-gpu-b200", nightly=True)
|
|
||||||
|
|
||||||
FULL_DEEPSEEK_V3_FP4_MODEL_PATH = "nvidia/DeepSeek-V3-0324-FP4"
|
|
||||||
SERVER_LAUNCH_TIMEOUT = 1000
|
|
||||||
|
|
||||||
|
|
||||||
class TestDeepseekV3FP4CutlassMoE(CustomTestCase):
|
|
||||||
@classmethod
|
|
||||||
def setUpClass(cls):
|
|
||||||
cls.model = FULL_DEEPSEEK_V3_FP4_MODEL_PATH
|
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
|
||||||
other_args = [
|
|
||||||
"--tp",
|
|
||||||
"4",
|
|
||||||
"--ep",
|
|
||||||
"4",
|
|
||||||
"--attention-backend",
|
|
||||||
"trtllm_mla",
|
|
||||||
"--moe-runner-backend",
|
|
||||||
"flashinfer_cutlass",
|
|
||||||
"--quantization",
|
|
||||||
"modelopt_fp4",
|
|
||||||
"--model-loader-extra-config",
|
|
||||||
'{"enable_multithread_load": true}',
|
|
||||||
]
|
|
||||||
cls.process = popen_launch_server(
|
|
||||||
cls.model,
|
|
||||||
cls.base_url,
|
|
||||||
timeout=SERVER_LAUNCH_TIMEOUT,
|
|
||||||
other_args=other_args,
|
|
||||||
)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def tearDownClass(cls):
|
|
||||||
kill_process_tree(cls.process.pid)
|
|
||||||
|
|
||||||
def test_a_gsm8k(
|
|
||||||
self,
|
|
||||||
): # Append an "a" to make this test run first (alphabetically) to warm up the server
|
|
||||||
args = SimpleNamespace(
|
|
||||||
base_url=self.base_url,
|
|
||||||
model=self.model,
|
|
||||||
eval_name="gsm8k",
|
|
||||||
api="completion",
|
|
||||||
max_tokens=512,
|
|
||||||
num_examples=1319,
|
|
||||||
num_threads=1319,
|
|
||||||
num_shots=8,
|
|
||||||
)
|
|
||||||
metrics = run_eval(args)
|
|
||||||
print(f"{metrics=}")
|
|
||||||
|
|
||||||
if is_in_ci():
|
|
||||||
write_github_step_summary(
|
|
||||||
f"### test_gsm8k (deepseek-v3-fp4-cutlass-moe)\n"
|
|
||||||
f'{metrics["score"]=:.3f}\n'
|
|
||||||
)
|
|
||||||
self.assertGreater(metrics["score"], 0.935)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
unittest.main()
|
|
||||||
@@ -1,87 +0,0 @@
|
|||||||
"""Extra: DeepSeek-V3 FP4 with FlashInfer Cutlass MoE backend.
|
|
||||||
|
|
||||||
Sibling per-commit file (test_deepseek_v3_fp4_4gpu.py) keeps the
|
|
||||||
SymmetricMemory variant.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import os
|
|
||||||
import unittest
|
|
||||||
from types import SimpleNamespace
|
|
||||||
|
|
||||||
from sglang.srt.utils import kill_process_tree
|
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
|
||||||
from sglang.test.run_eval import run_eval
|
|
||||||
from sglang.test.test_utils import (
|
|
||||||
DEFAULT_URL_FOR_TEST,
|
|
||||||
CustomTestCase,
|
|
||||||
is_in_ci,
|
|
||||||
popen_launch_server,
|
|
||||||
write_github_step_summary,
|
|
||||||
)
|
|
||||||
|
|
||||||
register_cuda_ci(est_time=960, stage="extra-b", runner_config="4-gpu-b200")
|
|
||||||
|
|
||||||
FULL_DEEPSEEK_V3_FP4_MODEL_PATH = "nvidia/DeepSeek-V3-0324-FP4"
|
|
||||||
SERVER_LAUNCH_TIMEOUT = 1200
|
|
||||||
|
|
||||||
|
|
||||||
class TestDeepseekV3FP4CutlassMoE(CustomTestCase):
|
|
||||||
@classmethod
|
|
||||||
def setUpClass(cls):
|
|
||||||
cls.model = FULL_DEEPSEEK_V3_FP4_MODEL_PATH
|
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
|
||||||
other_args = [
|
|
||||||
"--tp",
|
|
||||||
"4",
|
|
||||||
"--ep",
|
|
||||||
"4",
|
|
||||||
"--attention-backend",
|
|
||||||
"trtllm_mla",
|
|
||||||
"--moe-runner-backend",
|
|
||||||
"flashinfer_cutlass",
|
|
||||||
"--quantization",
|
|
||||||
"modelopt_fp4",
|
|
||||||
"--model-loader-extra-config",
|
|
||||||
'{"enable_multithread_load": true}',
|
|
||||||
]
|
|
||||||
cls.process = popen_launch_server(
|
|
||||||
cls.model,
|
|
||||||
cls.base_url,
|
|
||||||
timeout=SERVER_LAUNCH_TIMEOUT,
|
|
||||||
other_args=other_args,
|
|
||||||
env={
|
|
||||||
**os.environ,
|
|
||||||
"SGLANG_MOE_NVFP4_DISPATCH": "1", # Enable nvfp4 all gather
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def tearDownClass(cls):
|
|
||||||
kill_process_tree(cls.process.pid)
|
|
||||||
|
|
||||||
def test_a_gsm8k(
|
|
||||||
self,
|
|
||||||
): # Append an "a" to make this test run first (alphabetically) to warm up the server
|
|
||||||
args = SimpleNamespace(
|
|
||||||
base_url=self.base_url,
|
|
||||||
model=self.model,
|
|
||||||
eval_name="gsm8k",
|
|
||||||
api="completion",
|
|
||||||
max_tokens=512,
|
|
||||||
num_examples=1319,
|
|
||||||
num_threads=1319,
|
|
||||||
num_shots=8,
|
|
||||||
)
|
|
||||||
metrics = run_eval(args)
|
|
||||||
print(f"{metrics=}")
|
|
||||||
|
|
||||||
if is_in_ci():
|
|
||||||
write_github_step_summary(
|
|
||||||
f"### test_gsm8k (deepseek-v3-fp4-cutlass-moe)\n"
|
|
||||||
f'{metrics["score"]=:.3f}\n'
|
|
||||||
)
|
|
||||||
self.assertGreater(metrics["score"], 0.93)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
unittest.main()
|
|
||||||
@@ -0,0 +1,228 @@
|
|||||||
|
"""Numerics for the NVFP4 FusedMoE runner backends (--moe-runner-backend).
|
||||||
|
|
||||||
|
Runs the real FusedMoE layer path (construct -> fill NVFP4 checkpoint-format
|
||||||
|
weights -> process_weights_after_loading -> forward) per backend against a
|
||||||
|
dequantized torch MoE reference, covering the per-backend weight preparation
|
||||||
|
(TRTLLM shuffle / CUTLASS swizzle / CuteDSL v2 interleave + MMA blockscales)
|
||||||
|
and the MoE runner dispatch. Single GPU, tp=ep=1.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.moe.utils import MoeRunnerBackend
|
||||||
|
from sglang.srt.layers.quantization.modelopt_quant import ModelOptFp4Config
|
||||||
|
from sglang.srt.runtime_context import get_context, get_flags, get_parallel
|
||||||
|
from sglang.srt.utils import get_device_sm
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=120, stage="base-b", runner_config="4-gpu-b200")
|
||||||
|
|
||||||
|
E, H, I, TOPK, M = 8, 1024, 1024, 2, 32
|
||||||
|
FLOAT8_E4M3_MAX = 448.0
|
||||||
|
FLOAT4_E2M1_MAX = 6.0
|
||||||
|
|
||||||
|
kE2M1ToFloat = torch.tensor(
|
||||||
|
[0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=torch.float32
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _init_single_process_dist():
|
||||||
|
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
|
||||||
|
os.environ.setdefault("MASTER_PORT", "29631")
|
||||||
|
os.environ.setdefault("RANK", "0")
|
||||||
|
os.environ.setdefault("WORLD_SIZE", "1")
|
||||||
|
os.environ.setdefault("LOCAL_RANK", "0")
|
||||||
|
from sglang.srt.distributed.parallel_state import (
|
||||||
|
init_distributed_environment,
|
||||||
|
initialize_model_parallel,
|
||||||
|
model_parallel_is_initialized,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not torch.distributed.is_initialized():
|
||||||
|
init_distributed_environment(world_size=1, rank=0, local_rank=0, backend="gloo")
|
||||||
|
if not model_parallel_is_initialized():
|
||||||
|
initialize_model_parallel(
|
||||||
|
tensor_model_parallel_size=1,
|
||||||
|
expert_model_parallel_size=1,
|
||||||
|
pipeline_model_parallel_size=1,
|
||||||
|
backend="gloo",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def convert_swizzled_to_linear(a_sf_swizzled, m, k, block_size=16):
|
||||||
|
m_tiles = (m + 128 - 1) // 128
|
||||||
|
f = block_size * 4
|
||||||
|
k_tiles = (k + f - 1) // f
|
||||||
|
tmp = torch.reshape(a_sf_swizzled, (1, m_tiles, k_tiles, 32, 4, 4))
|
||||||
|
tmp = torch.permute(tmp, (0, 1, 4, 3, 2, 5))
|
||||||
|
out = tmp.reshape(m_tiles * 128, k_tiles * f // block_size)
|
||||||
|
return out[0:m, 0 : k // block_size]
|
||||||
|
|
||||||
|
|
||||||
|
def break_fp4_bytes(a):
|
||||||
|
m, n = a.shape
|
||||||
|
a_flat = a.flatten()
|
||||||
|
high = (a_flat & 0xF0) >> 4
|
||||||
|
low = a_flat & 0x0F
|
||||||
|
combined = torch.stack((low, high), dim=1).flatten()
|
||||||
|
signs = (combined & 0x08).to(torch.bool)
|
||||||
|
abs_vals = (combined & 0x07).to(torch.long)
|
||||||
|
kE2M1 = kE2M1ToFloat.to(device=a.device)
|
||||||
|
values = kE2M1[abs_vals] * torch.where(signs, -1.0, 1.0)
|
||||||
|
return values.reshape(m, n * 2).to(dtype=torch.float32)
|
||||||
|
|
||||||
|
|
||||||
|
def dequant_nvfp4(w_q, sf_swizzled, gs, n, k):
|
||||||
|
w_f32 = break_fp4_bytes(w_q).reshape(n, k // 16, 16)
|
||||||
|
sf = convert_swizzled_to_linear(sf_swizzled.view(torch.float8_e4m3fn), n, k)
|
||||||
|
return (w_f32 * (sf.float() / gs).unsqueeze(-1)).reshape(n, k)
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipIf(get_device_sm() < 100, "NVFP4 MoE backends require SM100+")
|
||||||
|
class TestNvFp4MoeBackends(CustomTestCase):
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
_init_single_process_dist()
|
||||||
|
torch.set_default_device("cuda")
|
||||||
|
|
||||||
|
def _run_backend(self, backend: str):
|
||||||
|
from flashinfer import fp4_quantize
|
||||||
|
|
||||||
|
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
||||||
|
from sglang.srt.layers.moe.topk import StandardTopKOutput
|
||||||
|
|
||||||
|
torch.manual_seed(7)
|
||||||
|
quant_config = ModelOptFp4Config(
|
||||||
|
is_checkpoint_nvfp4_serialized=True, group_size=16
|
||||||
|
)
|
||||||
|
with get_context().override_server_args(
|
||||||
|
model_path="dummy"
|
||||||
|
), get_flags().moe.override(
|
||||||
|
runner_backend=MoeRunnerBackend(backend)
|
||||||
|
), get_parallel().override(
|
||||||
|
moe_ep_size=1,
|
||||||
|
moe_ep_rank=0,
|
||||||
|
moe_tp_size=1,
|
||||||
|
moe_tp_rank=0,
|
||||||
|
tp_size=1,
|
||||||
|
tp_rank=0,
|
||||||
|
):
|
||||||
|
layer = FusedMoE(
|
||||||
|
num_experts=E,
|
||||||
|
hidden_size=H,
|
||||||
|
intermediate_size=I,
|
||||||
|
layer_id=0,
|
||||||
|
top_k=TOPK,
|
||||||
|
params_dtype=torch.bfloat16,
|
||||||
|
quant_config=quant_config,
|
||||||
|
gate_up_interleaved=False,
|
||||||
|
).cuda()
|
||||||
|
|
||||||
|
w13_ref = torch.zeros(E, 2 * I, H, dtype=torch.float32, device="cuda")
|
||||||
|
w2_ref = torch.zeros(E, H, I, dtype=torch.float32, device="cuda")
|
||||||
|
for e in range(E):
|
||||||
|
w13 = torch.randn(2 * I, H, dtype=torch.bfloat16, device="cuda") / 10
|
||||||
|
w2 = torch.randn(H, I, dtype=torch.bfloat16, device="cuda") / 10
|
||||||
|
w13_gs = FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX / w13.abs().max().float()
|
||||||
|
w2_gs = FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX / w2.abs().max().float()
|
||||||
|
w13_q, w13_sf = fp4_quantize(w13, w13_gs)
|
||||||
|
w2_q, w2_sf = fp4_quantize(w2, w2_gs)
|
||||||
|
layer.w13_weight.data[e].copy_(w13_q)
|
||||||
|
layer.w2_weight.data[e].copy_(w2_q)
|
||||||
|
layer.w13_weight_scale.data[e].copy_(
|
||||||
|
convert_swizzled_to_linear(
|
||||||
|
w13_sf.view(torch.float8_e4m3fn), 2 * I, H
|
||||||
|
)
|
||||||
|
)
|
||||||
|
layer.w2_weight_scale.data[e].copy_(
|
||||||
|
convert_swizzled_to_linear(w2_sf.view(torch.float8_e4m3fn), H, I)
|
||||||
|
)
|
||||||
|
layer.w13_weight_scale_2.data[e].fill_(1.0 / w13_gs)
|
||||||
|
layer.w2_weight_scale_2.data[e].fill_(1.0 / w2_gs)
|
||||||
|
w13_ref[e] = dequant_nvfp4(w13_q, w13_sf, w13_gs, 2 * I, H)
|
||||||
|
w2_ref[e] = dequant_nvfp4(w2_q, w2_sf, w2_gs, H, I)
|
||||||
|
act_scale = 1.0 / (FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX)
|
||||||
|
layer.w13_input_scale.data.fill_(act_scale)
|
||||||
|
layer.w2_input_scale.data.fill_(act_scale)
|
||||||
|
|
||||||
|
layer.quant_method.process_weights_after_loading(layer)
|
||||||
|
|
||||||
|
x = torch.randn(M, H, dtype=torch.bfloat16, device="cuda") / 10
|
||||||
|
router_logits = torch.randn(M, E, dtype=torch.float32, device="cuda")
|
||||||
|
weights = torch.softmax(router_logits, dim=-1)
|
||||||
|
topk_weights, topk_ids = torch.topk(weights, TOPK, dim=-1)
|
||||||
|
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||||
|
|
||||||
|
out = layer.forward(
|
||||||
|
x,
|
||||||
|
StandardTopKOutput(
|
||||||
|
topk_weights=topk_weights,
|
||||||
|
topk_ids=topk_ids.to(torch.int32),
|
||||||
|
router_logits=router_logits,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
if not isinstance(out, torch.Tensor):
|
||||||
|
out = out[0] if isinstance(out, tuple) else out.hidden_states
|
||||||
|
|
||||||
|
ref = self._torch_moe_reference(
|
||||||
|
layer, x, topk_weights, topk_ids, w13_ref, w2_ref
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(out.shape, (M, H))
|
||||||
|
cos = torch.nn.functional.cosine_similarity(
|
||||||
|
out.float().flatten(), ref.flatten(), dim=0
|
||||||
|
).item()
|
||||||
|
self.assertGreater(cos, 0.99)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _torch_moe_reference(layer, x, topk_weights, topk_ids, w13_ref, w2_ref):
|
||||||
|
from flashinfer import fp4_quantize
|
||||||
|
|
||||||
|
def quant_roundtrip(t2d, gs):
|
||||||
|
q, sf = fp4_quantize(t2d.to(torch.bfloat16), gs)
|
||||||
|
return dequant_nvfp4(q, sf, gs, t2d.shape[0], t2d.shape[1])
|
||||||
|
|
||||||
|
# The kernels quantize the input and the GEMM1->GEMM2 intermediate to
|
||||||
|
# NVFP4; mirror both round trips or the comparison carries ~7% noise.
|
||||||
|
act_gs = torch.tensor(
|
||||||
|
FLOAT8_E4M3_MAX * FLOAT4_E2M1_MAX, dtype=torch.float32, device="cuda"
|
||||||
|
)
|
||||||
|
# TRTLLM consumes w13 as [up; gate] (GEMM1 scales are applied on that
|
||||||
|
# assumption); CUTLASS / CuteDSL-v2 load up first as well via
|
||||||
|
# load_up_proj_weight_first.
|
||||||
|
up_first = (
|
||||||
|
layer.quant_method.load_up_proj_weight_first
|
||||||
|
or layer.quant_method.enable_flashinfer_trtllm_moe
|
||||||
|
)
|
||||||
|
x_dq = quant_roundtrip(x.float(), act_gs)
|
||||||
|
m, h = x.shape
|
||||||
|
ref = torch.zeros(m, h, dtype=torch.float32, device="cuda")
|
||||||
|
for t in range(m):
|
||||||
|
for j in range(topk_ids.shape[1]):
|
||||||
|
e = int(topk_ids[t, j])
|
||||||
|
gu = x_dq[t] @ w13_ref[e].T
|
||||||
|
if up_first:
|
||||||
|
up, gate = gu[:I], gu[I:]
|
||||||
|
else:
|
||||||
|
gate, up = gu[:I], gu[I:]
|
||||||
|
act = torch.nn.functional.silu(gate) * up
|
||||||
|
act_dq = quant_roundtrip(act.unsqueeze(0), act_gs)[0]
|
||||||
|
ref[t] += float(topk_weights[t, j]) * (act_dq @ w2_ref[e].T)
|
||||||
|
return ref
|
||||||
|
|
||||||
|
def test_flashinfer_cutlass(self):
|
||||||
|
self._run_backend("flashinfer_cutlass")
|
||||||
|
|
||||||
|
def test_flashinfer_trtllm(self):
|
||||||
|
self._run_backend("flashinfer_trtllm")
|
||||||
|
|
||||||
|
def test_flashinfer_cutedsl(self):
|
||||||
|
self._run_backend("flashinfer_cutedsl")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user