[NVIDIA] Fix broken fp8 MoE of deepseek v3 (#13264)

Co-authored-by: Baizhou Zhang <sobereddiezhang@gmail.com>
Co-authored-by: Kangyan-Zhou <zky314343421@gmail.com>
This commit is contained in:
Kaixi Hou
2025-11-17 16:13:28 -08:00
committed by GitHub
parent aac07bf7fd
commit e389f91dec
5 changed files with 117 additions and 4 deletions

View File

@@ -182,6 +182,27 @@ jobs:
cd test/srt
python3 run_suite.py --suite nightly-8-gpu-h20 --continue-on-error
nightly-test-8-gpu-b200:
if: github.repository == 'sgl-project/sglang'
runs-on: 8-gpu-b200
env:
RUNNER_LABELS: 8-gpu-b200
strategy:
fail-fast: false
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Install dependencies
run: |
IS_BLACKWELL=1 bash scripts/ci/ci_install_dependency.sh
- name: Run test
timeout-minutes: 45
run: |
cd test/srt
python3 run_suite.py --suite nightly-8-gpu-b200 --auto-partition-id 0 --auto-partition-size 1 --timeout-per-file 3600
check-all-jobs:
if: always() && (github.repository == 'sgl-project/sglang' || github.event_name == 'workflow_dispatch')
needs:
@@ -193,6 +214,7 @@ jobs:
- nightly-test-4-gpu
- nightly-test-8-gpu-h200
- nightly-test-8-gpu-h20
- nightly-test-8-gpu-b200
runs-on: ubuntu-latest
steps:
- name: Check if any job failed

View File

@@ -1210,9 +1210,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
else topk_config.correction_bias.to(x.dtype)
)
routing_method_type = getattr(
layer, "routing_method_type", RoutingMethodType.DeepSeekV3
)
routing_method_type = getattr(layer, "routing_method_type")
with use_symmetric_memory(
get_tp_group(), disabled=not is_allocation_symmetric()

View File

@@ -93,6 +93,7 @@ from sglang.srt.layers.moe.ep_moe.layer import DeepEPMoE, get_moe_impl_class
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.layers.moe.kt_ep_wrapper import KTEPWrapperMethod
from sglang.srt.layers.moe.topk import TopK, TopKOutputFormat
from sglang.srt.layers.moe.utils import RoutingMethodType
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.quantization.fp8 import Fp8Config
from sglang.srt.layers.quantization.fp8_kernel import (
@@ -674,6 +675,7 @@ class DeepseekV2MoE(nn.Module):
layer_id=self.layer_id,
quant_config=quant_config,
routed_scaling_factor=self.routed_scaling_factor,
routing_method_type=RoutingMethodType.DeepSeekV3,
prefix=add_prefix("experts", prefix),
)

View File

@@ -176,6 +176,7 @@ suites = {
# TODO: Add it back after the bug is fixed
# TestFile("test_eagle_infer_beta_dp_attention.py", 200),
],
"per-commit-8-gpu-b200": [],
"per-commit-4-gpu-gb200": [
TestFile("test_cutedsl_moe.py", 300),
TestFile("test_deepseek_v3_fp4_4gpu.py", 1800),
@@ -215,7 +216,9 @@ suites = {
TestFile("test_deepseek_v3_fp4_cutlass_moe.py", 900),
TestFile("test_fp4_moe.py", 300),
],
"nightly-8-gpu-b200": [],
"nightly-8-gpu-b200": [
TestFile("test_deepseek_r1_fp8_trtllm_backend.py", 3600),
],
"nightly-4-gpu": [
TestFile("test_qwen3_next_deterministic.py", 200),
],

View File

@@ -0,0 +1,88 @@
import os
import unittest
from types import SimpleNamespace
from sglang.srt.utils import kill_process_tree
from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
try_cached_model,
)
FULL_DEEPSEEK_V3_MODEL_PATH = "deepseek-ai/DeepSeek-V3-0324"
class TestDeepseekR1Fp8Flashinfer(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = try_cached_model(FULL_DEEPSEEK_V3_MODEL_PATH)
cls.base_url = DEFAULT_URL_FOR_TEST
other_args = [
"--trust-remote-code",
"--disable-radix-cache",
"--max-running-requests",
"512",
"--chunked-prefill-size",
"8192",
"--mem-fraction-static",
"0.9",
"--cuda-graph-max-bs",
"128",
"--max-prefill-tokens",
"8192",
"--kv-cache-dtype",
"fp8_e4m3",
"--quantization",
"fp8",
"--tensor-parallel-size",
"8",
"--data-parallel-size",
"1",
"--expert-parallel-size",
"1",
"--scheduler-recv-interval",
"10",
"--stream-interval",
"10",
"--attention-backend",
"trtllm_mla",
"--moe-runner-backend",
"flashinfer_trtllm",
"--enable-symm-mem",
]
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=other_args,
env={
**os.environ,
"SGLANG_ENABLE_FLASHINFER_FP8_GEMM": "1",
},
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
def test_gsm8k(self):
args = SimpleNamespace(
num_shots=5,
data_path=None,
num_questions=512,
parallel=512,
max_new_tokens=512,
host="http://127.0.0.1",
port=int(self.base_url.split(":")[-1]),
)
metrics = run_eval_few_shot_gsm8k(args)
print(f"Eval accuracy of GSM8K: {metrics=}")
self.assertGreater(metrics["accuracy"], 0.92)
if __name__ == "__main__":
unittest.main()