From b61cb5f9de87f07770ba216fab4fba1e6496f4e3 Mon Sep 17 00:00:00 2001 From: Broduker <30527731+Broduker@users.noreply.github.com> Date: Thu, 30 Jul 2026 23:03:34 +0800 Subject: [PATCH] Fix DeepSeek V4 loading with RunAI Model Streamer. (#30240) --- python/sglang/srt/models/deepseek_v4.py | 70 ++++++++-- .../sglang/srt/models/deepseek_v4_dspark.py | 6 +- .../test_runai_model_streamer_loader.py | 123 ++++++++++++++++++ 3 files changed, 185 insertions(+), 14 deletions(-) diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index a21936d73..a81b5de4d 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -119,7 +119,10 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo get_tc_piecewise_forward_context, ) from sglang.srt.model_loader.utils import maybe_executor_submit, should_async_load -from sglang.srt.model_loader.weight_utils import default_weight_loader +from sglang.srt.model_loader.weight_utils import ( + RUNAI_STREAMER_TENSOR_ATTR, + default_weight_loader, +) from sglang.srt.models.dbrx import ReplicatedLinear from sglang.srt.models.deepseek_common.amd.deepseek_v4_fused_mhc import ( try_fused_hc_post_pre, @@ -2814,13 +2817,7 @@ class DeepseekV4ForCausalLM(nn.Module): raise ValueError("num_nextn_predict_layers is not in the config") if not envs.SGLANG_OPT_FP8_WO_A_GEMM.get(): - weights = list(weights) - exists_wo_a_scale = any(n.endswith(".wo_a.scale") for n, t in weights) - if exists_wo_a_scale: - logger.info("Execute dequant fp8 wo_a") - weights = _dequant_fp8_wo_a(weights) - else: - logger.info("Skip dequant fp8 wo_a") + weights = _dequant_fp8_wo_a_streaming(weights) stacked_params_mapping = DEEPSEEK_V4_STACKED_PARAMS_MAPPING @@ -3026,7 +3023,7 @@ class DeepseekV4ForCausalLM(nn.Module): if key not in cache_compressor_weight: cache_compressor_weight[key] = ( is_kv, - loaded_weight, + _clone_if_runai_streamed_tensor(loaded_weight), ) else: assert key in cache_compressor_weight @@ -3064,7 +3061,9 @@ class DeepseekV4ForCausalLM(nn.Module): assert ( shard_key not in bucket ), f"duplicate shard {shard_key} for {param_name}" - bucket[shard_key] = loaded_weight + bucket[shard_key] = _clone_if_runai_streamed_tensor( + loaded_weight + ) if len(bucket) == 2: fused_weight = torch.cat( [bucket["q"], bucket["kv"]], dim=0 @@ -3194,6 +3193,57 @@ def _dequant_fp8(weight: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: return result.to(torch.bfloat16) +def _clone_if_runai_streamed_tensor(tensor: torch.Tensor) -> torch.Tensor: + if getattr(tensor, RUNAI_STREAMER_TENSOR_ATTR, False): + return tensor.clone().detach() + return tensor + + +def _dequant_fp8_wo_a_streaming( + weights: Iterable[Tuple[str, torch.Tensor]], +) -> Iterable[Tuple[str, torch.Tensor]]: + pending: dict[str, dict[str, torch.Tensor]] = {} + saw_wo_a_scale = False + emitted = False + + for name, tensor in weights: + if name.endswith(".wo_a.weight"): + prefix = name[: -len(".weight")] + bucket = pending.setdefault(prefix, {}) + scale = bucket.pop("scale", None) + if scale is not None: + pending.pop(prefix, None) + emitted = True + yield name, _dequant_fp8(tensor, scale) + else: + bucket["weight"] = _clone_if_runai_streamed_tensor(tensor) + continue + + if name.endswith(".wo_a.scale"): + saw_wo_a_scale = True + prefix = name[: -len(".scale")] + bucket = pending.setdefault(prefix, {}) + weight = bucket.pop("weight", None) + if weight is not None: + pending.pop(prefix, None) + emitted = True + yield prefix + ".weight", _dequant_fp8(weight, tensor) + else: + bucket["scale"] = _clone_if_runai_streamed_tensor(tensor) + continue + + yield name, tensor + + if emitted: + logger.info("Finished streaming dequant fp8 wo_a") + for prefix, bucket in pending.items(): + if "weight" in bucket: + assert not saw_wo_a_scale, f"{prefix}.scale is missing" + yield prefix + ".weight", bucket["weight"] + if "scale" in bucket: + yield prefix + ".scale", bucket["scale"] + + def _dequant_fp8_wo_a( weights: Iterable[Tuple[str, torch.Tensor]], ) -> Iterable[Tuple[str, torch.Tensor]]: diff --git a/python/sglang/srt/models/deepseek_v4_dspark.py b/python/sglang/srt/models/deepseek_v4_dspark.py index dc2bdbf19..2b01ec2e2 100644 --- a/python/sglang/srt/models/deepseek_v4_dspark.py +++ b/python/sglang/srt/models/deepseek_v4_dspark.py @@ -30,7 +30,7 @@ from sglang.srt.models.deepseek_v4 import ( DEEPSEEK_V4_STACKED_PARAMS_MAPPING, DeepseekV4DecoderLayer, MqaAttentionBase, - _dequant_fp8_wo_a, + _dequant_fp8_wo_a_streaming, hc_head_torch, make_hc_head_params, ) @@ -761,9 +761,7 @@ class DeepseekV4ForCausalLMDSpark(nn.Module): params_dict = dict(self.named_parameters()) loaded_params = set() - weights = list(weights) - if any(name.endswith(".wo_a.scale") for name, _ in weights): - weights = list(_dequant_fp8_wo_a(weights)) + weights = _dequant_fp8_wo_a_streaming(weights) stacked_params_mapping = DEEPSEEK_V4_STACKED_PARAMS_MAPPING from sglang.srt.layers.moe.fused_moe_triton import FusedMoE diff --git a/test/registered/unit/model_loader/test_runai_model_streamer_loader.py b/test/registered/unit/model_loader/test_runai_model_streamer_loader.py index f7215d1ed..d4fc161b8 100644 --- a/test/registered/unit/model_loader/test_runai_model_streamer_loader.py +++ b/test/registered/unit/model_loader/test_runai_model_streamer_loader.py @@ -12,6 +12,11 @@ from sglang.srt.configs.device_config import DeviceConfig from sglang.srt.configs.load_config import LoadConfig, LoadFormat from sglang.srt.configs.model_config import ModelConfig from sglang.srt.models.deepseek_common import deepseek_weight_loader +from sglang.srt.models.deepseek_v4 import ( + _dequant_fp8_wo_a, + _dequant_fp8_wo_a_streaming, +) +from sglang.srt.models.deepseek_v4_dspark import DeepseekV4ForCausalLMDSpark from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -105,6 +110,124 @@ class TestRunaiModelStreamerLoader(CustomTestCase): marked.fill_(2) self.assertEqual(cloned.item(), 1) + def test_deepseek_v4_streaming_dequant_fp8_wo_a_pairs_weight_and_scale(self): + weight = torch.eye(128, dtype=torch.float32).to(torch.float8_e4m3fn) + scale = torch.ones((1, 1), dtype=torch.float32) + + for weights in ( + [ + ("layers.0.attn.wo_a.scale", scale), + ("layers.0.attn.wo_a.weight", weight), + ("layers.0.attn.wq.weight", torch.tensor([3])), + ], + [ + ("layers.0.attn.wo_a.weight", weight), + ("layers.0.attn.wq.weight", torch.tensor([3])), + ("layers.0.attn.wo_a.scale", scale), + ], + ): + converted = list(_dequant_fp8_wo_a_streaming(weights)) + + converted_names = [name for name, _ in converted] + self.assertIn("layers.0.attn.wo_a.weight", converted_names) + self.assertNotIn("layers.0.attn.wo_a.scale", converted_names) + converted_weight = dict(converted)["layers.0.attn.wo_a.weight"] + self.assertEqual(converted_weight.dtype, torch.bfloat16) + + def test_deepseek_v4_streaming_dequant_matches_legacy_by_name(self): + weight = torch.eye(128, dtype=torch.float32).to(torch.float8_e4m3fn) + scale = torch.ones((1, 1), dtype=torch.float32) + ordinary = torch.tensor([3]) + weights = [ + ("layers.0.attn.wo_a.weight", weight), + ("layers.0.attn.wq.weight", ordinary), + ("layers.0.attn.wo_a.scale", scale), + ] + + legacy = list(_dequant_fp8_wo_a(weights)) + streaming = list(_dequant_fp8_wo_a_streaming(weights)) + + self.assertNotEqual( + [name for name, _ in legacy], [name for name, _ in streaming] + ) + self.assertEqual(set(dict(legacy)), set(dict(streaming))) + for name, legacy_tensor in dict(legacy).items(): + torch.testing.assert_close(legacy_tensor, dict(streaming)[name]) + + def test_deepseek_v4_streaming_dequant_clones_pending_runai_tensors(self): + weight = torch.eye(128, dtype=torch.float32).to(torch.float8_e4m3fn) + scale = torch.ones((1, 1), dtype=torch.float32) + setattr(scale, weight_utils.RUNAI_STREAMER_TENSOR_ATTR, True) + + def weights(): + yield "layers.0.attn.wo_a.scale", scale + scale.fill_(0) + yield "layers.0.attn.wo_a.weight", weight + + converted = dict(_dequant_fp8_wo_a_streaming(weights())) + + converted_weight = converted["layers.0.attn.wo_a.weight"] + self.assertGreater(converted_weight.abs().sum().item(), 0) + + def test_deepseek_v4_dspark_load_weights_streams_wo_a_dequant(self): + weight = torch.eye(128, dtype=torch.float32).to(torch.float8_e4m3fn) + scale = torch.ones((1, 1), dtype=torch.float32) + setattr(scale, weight_utils.RUNAI_STREAMER_TENSOR_ATTR, True) + loaded_weights = [] + + def weight_loader(_param, loaded_weight): + loaded_weights.append(loaded_weight) + + param = SimpleNamespace(weight_loader=weight_loader) + remapper = SimpleNamespace(confidence_head=None) + model = SimpleNamespace( + config=SimpleNamespace(n_routed_experts=1), + named_parameters=lambda: [ + ("stages.0.self_attn.wo_a.weight", param), + ], + _remap_dspark_weight_name=lambda name: ( + DeepseekV4ForCausalLMDSpark._remap_dspark_weight_name(remapper, name) + ), + _assert_confidence_head_loaded=lambda **_kwargs: None, + ) + + def weights(): + yield "mtp.0.attn.wo_a.scale", scale + scale.fill_(0) + yield "mtp.0.attn.wo_a.weight", weight + + DeepseekV4ForCausalLMDSpark.load_weights(model, weights()) + + self.assertEqual(len(loaded_weights), 1) + self.assertEqual(loaded_weights[0].dtype, torch.bfloat16) + self.assertGreater(loaded_weights[0].abs().sum().item(), 0) + + def test_deepseek_v4_streaming_dequant_preserves_missing_scale_behavior(self): + weight = torch.eye(128, dtype=torch.float32).to(torch.float8_e4m3fn) + ordinary = torch.tensor([3]) + + converted = dict( + _dequant_fp8_wo_a_streaming( + [ + ("layers.0.attn.wo_a.weight", weight), + ("layers.0.attn.wq.weight", ordinary), + ] + ) + ) + + self.assertIs(converted["layers.0.attn.wo_a.weight"], weight) + self.assertIs(converted["layers.0.attn.wq.weight"], ordinary) + + with self.assertRaises(AssertionError): + list( + _dequant_fp8_wo_a_streaming( + [ + ("layers.0.attn.wo_a.weight", weight), + ("layers.1.attn.wo_a.scale", torch.ones((1, 1))), + ] + ) + ) + def test_get_model_loader_uses_runai_for_prequantized_modelopt(self): load_config = LoadConfig( load_format=LoadFormat.RUNAI_STREAMER,