Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
56 changes: 56 additions & 0 deletions python/sglang/srt/models/qwen3_5_mtp.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,9 @@
"""Inference-only Qwen3_5 MTP model."""

import copy
import json
import logging
import os
from contextlib import ExitStack
from typing import Iterable, Optional, Tuple

Expand Down Expand Up @@ -82,6 +84,41 @@ def _mtp_quant_config(quant_config):
return quant_config


def _load_mtp_w8a8_dequant_scales(model_path: str) -> dict:
"""Per-row int8 dequant scales for W8A8 MTP weights, keyed by the
original checkpoint weight name.

On NPU the MTP draft runs unquantized (quant_config=None) while
ModelSlim checkpoints may store the MTP projections as W8A8 (int8 +
per-row weight_scale). Without dequantization the raw int8 codes are
copied into the bf16 parameters and the draft model is numerically
broken (spec-decode accept rate collapses to ~0).
"""
scales: dict = {}
try:
index_file = os.path.join(
model_path, "quant_model_weights.safetensors.index.json"
)
if not os.path.exists(index_file):
return scales
with open(index_file) as f:
weight_map = json.load(f)["weight_map"]
from safetensors import safe_open

for key in weight_map:
if not key.startswith("mtp.") or not key.endswith(".weight_scale"):
continue
with safe_open(
os.path.join(model_path, weight_map[key]),
framework="pt",
device="cpu",
) as f:
scales[key[: -len("_scale")]] = f.get_tensor(key)
except Exception as e:
logger.warning("MTP draft: failed to load W8A8 scales: %s", e)
return scales


class Qwen3_5ForCausalLMMTP(nn.Module):
@staticmethod
def shared_experts_fusion_disable_reason(hf_config, quant_config):
Expand Down Expand Up @@ -321,6 +358,13 @@ def load_fused_expert_weights(

params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
mtp_w8a8_scales = (
_load_mtp_w8a8_dequant_scales(
get_spec().speculative_draft_model_path or get_model().model_path
)
if self.quant_config is None
else {}
)

for name, loaded_weight in weights:
# The last-stage MTP draft cannot share the target embedding on PP0.
Expand All @@ -347,6 +391,18 @@ def load_fused_expert_weights(
if "mtp" not in name:
continue

orig_ckpt_name = name
if (
self.quant_config is None
and loaded_weight.dtype == torch.int8
and orig_ckpt_name in mtp_w8a8_scales
):
w8a8_scale = mtp_w8a8_scales[orig_ckpt_name]
if w8a8_scale.shape[0] == loaded_weight.shape[0]:
loaded_weight = (
loaded_weight.to(torch.float32) * w8a8_scale.to(torch.float32)
).to(getattr(self.config, "torch_dtype", torch.bfloat16))

if name.startswith("mtp."):
# Remove the mtp. prefix for processing
name = name.replace("mtp.", "model.")
Expand Down
159 changes: 159 additions & 0 deletions test/registered/unit/models/test_qwen3_5_mtp_w8a8_dequant.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,159 @@
"""CPU coverage for W8A8 dequantization of MTP weights on the unquantized
(NPU) draft path.

On NPU the MTP draft is forced to ``quant_config=None`` while ModelSlim
checkpoints may store the MTP projections as W8A8 (int8 + per-row
``weight_scale``). Without dequantization the raw int8 codes are copied
into the bf16 parameters and the draft model is numerically broken
(spec-decode accept rate collapses to ~0). These tests cover scale
discovery and the ``load_weights`` dequant branch without an NPU.
"""

import json
from types import SimpleNamespace

import torch
from safetensors.torch import save_file

from sglang.srt.models import qwen3_5_mtp
from sglang.srt.models.qwen3_5_mtp import (
Qwen3_5ForCausalLMMTP,
_load_mtp_w8a8_dequant_scales,
)
from sglang.test.ci.ci_register import register_cpu_ci

register_cpu_ci(est_time=10, suite="base-a-test-cpu")

Q_PROJ = "mtp.layers.0.self_attn.q_proj"


def _write_ckpt(tmp_path, tensors):
save_file(tensors, str(tmp_path / "shard_00001.safetensors"))
index = {
"metadata": {"total_size": 0},
"weight_map": {name: "shard_00001.safetensors" for name in tensors},
}
(tmp_path / "quant_model_weights.safetensors.index.json").write_text(
json.dumps(index)
)
return str(tmp_path)


def _stub_model(monkeypatch, ckpt, quant_config=None):
"""Bypass __init__ (needs runtime context); set only what load_weights reads."""
model = Qwen3_5ForCausalLMMTP.__new__(Qwen3_5ForCausalLMMTP)
torch.nn.Module.__init__(model)
model.quant_config = quant_config
model.config = SimpleNamespace(torch_dtype=torch.bfloat16)
monkeypatch.setattr(
qwen3_5_mtp,
"get_spec",
lambda: SimpleNamespace(speculative_draft_model_path=ckpt),
)
monkeypatch.setattr(
qwen3_5_mtp, "get_model", lambda: SimpleNamespace(model_path=ckpt)
)
return model


class _CaptureParam:
def __init__(self):
self.calls = []

def weight_loader(self, param, loaded_weight, shard_id):
self.calls.append((shard_id, loaded_weight))


def test_scale_discovery_keys_by_weight_name(tmp_path):
q = torch.randint(-127, 128, (16, 8), dtype=torch.int8)
s = (torch.rand(16, 1) * 0.01 + 0.001).float()
ckpt = _write_ckpt(
tmp_path,
{
f"{Q_PROJ}.weight": q,
f"{Q_PROJ}.weight_scale": s,
# FLOAT (unquantized) MTP tensors carry no scale
"mtp.fc.weight": torch.randn(8, 32, dtype=torch.bfloat16),
# non-mtp scales must be ignored
"model.language_model.layers.0.mlp.gate_proj.weight_scale": torch.rand(
4, 1
),
},
)
scales = _load_mtp_w8a8_dequant_scales(ckpt)
assert list(scales) == [f"{Q_PROJ}.weight"]
assert torch.equal(scales[f"{Q_PROJ}.weight"], s)


def test_scale_discovery_without_index(tmp_path):
assert _load_mtp_w8a8_dequant_scales(str(tmp_path)) == {}


def test_scale_discovery_missing_shard_does_not_raise(tmp_path):
(tmp_path / "quant_model_weights.safetensors.index.json").write_text(
json.dumps(
{"weight_map": {f"{Q_PROJ}.weight_scale": "missing.safetensors"}}
)
)
assert _load_mtp_w8a8_dequant_scales(str(tmp_path)) == {}


def test_load_weights_dequantizes_int8_mtp_projection(monkeypatch, tmp_path):
q = torch.randint(-127, 128, (16, 8), dtype=torch.int8)
s = (torch.rand(16, 1) * 0.01 + 0.001).float()
ckpt = _write_ckpt(
tmp_path, {f"{Q_PROJ}.weight": q, f"{Q_PROJ}.weight_scale": s}
)
model = _stub_model(monkeypatch, ckpt)
capture = _CaptureParam()
model.named_parameters = lambda *a, **k: [
("model.layers.0.qkv_proj.weight", capture)
]

model.load_weights(iter([(f"{Q_PROJ}.weight", q.clone())]))

assert len(capture.calls) == 1
shard_id, loaded = capture.calls[0]
assert shard_id == "q"
assert loaded.dtype == torch.bfloat16
expected = (q.float() * s.float()).to(torch.bfloat16)
assert torch.equal(loaded, expected)


def test_load_weights_passes_int8_through_when_quantized(monkeypatch, tmp_path):
q = torch.randint(-127, 128, (16, 8), dtype=torch.int8)
s = (torch.rand(16, 1) * 0.01 + 0.001).float()
ckpt = _write_ckpt(
tmp_path, {f"{Q_PROJ}.weight": q, f"{Q_PROJ}.weight_scale": s}
)
model = _stub_model(monkeypatch, ckpt, quant_config=SimpleNamespace())
capture = _CaptureParam()
model.named_parameters = lambda *a, **k: [
("model.layers.0.qkv_proj.weight", capture)
]

model.load_weights(iter([(f"{Q_PROJ}.weight", q.clone())]))

shard_id, loaded = capture.calls[0]
assert shard_id == "q"
assert loaded.dtype == torch.int8
assert torch.equal(loaded, q)


def test_load_weights_scale_shape_mismatch_is_noop(monkeypatch, tmp_path):
q = torch.randint(-127, 128, (16, 8), dtype=torch.int8)
s = (torch.rand(4, 1) * 0.01 + 0.001).float() # wrong row count
ckpt = _write_ckpt(
tmp_path, {f"{Q_PROJ}.weight": q, f"{Q_PROJ}.weight_scale": s}
)
model = _stub_model(monkeypatch, ckpt)
capture = _CaptureParam()
model.named_parameters = lambda *a, **k: [
("model.layers.0.qkv_proj.weight", capture)
]

model.load_weights(iter([(f"{Q_PROJ}.weight", q.clone())]))

_, loaded = capture.calls[0]
assert loaded.dtype == torch.int8
assert torch.equal(loaded, q)
Loading