Support RunAI loading for quantized checkpoints (#23850)
Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com> Co-authored-by: Sam Shleifer <sam@thinkingmachines.ai>
This commit is contained in:
co-authored by
Claude Opus 4.7
Sam Shleifer
parent
321298da75
commit
d41e8c459d
@@ -3238,11 +3238,13 @@ class RunaiModelStreamerLoader(BaseModelLoader):
|
|||||||
self.target_device_str = "cpu"
|
self.target_device_str = "cpu"
|
||||||
|
|
||||||
target_device = torch.device(device_config.device)
|
target_device = torch.device(device_config.device)
|
||||||
|
quant_config = _get_quantization_config(model_config, self.load_config)
|
||||||
with set_default_torch_dtype(model_config.dtype):
|
with set_default_torch_dtype(model_config.dtype):
|
||||||
with target_device:
|
with target_device:
|
||||||
model = _initialize_model(
|
model = _initialize_model(
|
||||||
model_config,
|
model_config,
|
||||||
self.load_config,
|
self.load_config,
|
||||||
|
quant_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
DefaultModelLoader.load_weights_and_postprocess(
|
DefaultModelLoader.load_weights_and_postprocess(
|
||||||
@@ -3260,7 +3262,16 @@ def get_model_loader(
|
|||||||
if load_config.load_format == LoadFormat.DUMMY:
|
if load_config.load_format == LoadFormat.DUMMY:
|
||||||
return DummyModelLoader(load_config)
|
return DummyModelLoader(load_config)
|
||||||
|
|
||||||
if model_config and (
|
# ModelOptModelLoader's local-copy quantize-and-export workflow doesn't apply
|
||||||
|
# to RUNAI_STREAMER, which streams weights directly from object storage.
|
||||||
|
# RUNAI_STREAMER loads always fall through to the unconditional branch at
|
||||||
|
# the bottom of this function. This also avoids calling _is_already_quantized()
|
||||||
|
# on RunAI streamer cache paths, where huggingface_hub raises HFValidationError.
|
||||||
|
model_optloader_allowed = (
|
||||||
|
model_config and load_config.load_format != LoadFormat.RUNAI_STREAMER
|
||||||
|
)
|
||||||
|
|
||||||
|
if model_optloader_allowed and (
|
||||||
(hasattr(model_config, "modelopt_quant") and model_config.modelopt_quant)
|
(hasattr(model_config, "modelopt_quant") and model_config.modelopt_quant)
|
||||||
or model_config.quantization
|
or model_config.quantization
|
||||||
in ["modelopt_fp8", "modelopt_fp4", "modelopt_mixed", "modelopt"]
|
in ["modelopt_fp8", "modelopt_fp4", "modelopt_mixed", "modelopt"]
|
||||||
@@ -3270,7 +3281,7 @@ def get_model_loader(
|
|||||||
|
|
||||||
# Use ModelOptModelLoader for unified quantization flags
|
# Use ModelOptModelLoader for unified quantization flags
|
||||||
if (
|
if (
|
||||||
model_config
|
model_optloader_allowed
|
||||||
and hasattr(model_config, "quantization")
|
and hasattr(model_config, "quantization")
|
||||||
and model_config.quantization
|
and model_config.quantization
|
||||||
in ["modelopt_fp8", "modelopt_fp4", "modelopt_mixed"]
|
in ["modelopt_fp8", "modelopt_fp4", "modelopt_mixed"]
|
||||||
|
|||||||
@@ -69,6 +69,8 @@ except ImportError as e:
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
RUNAI_STREAMER_TENSOR_ATTR = "_sglang_runai_streamer_tensor"
|
||||||
|
|
||||||
# Block size for sequential checkpoint prefetch reads (page cache warming).
|
# Block size for sequential checkpoint prefetch reads (page cache warming).
|
||||||
_PREFETCH_BLOCK_SIZE = None
|
_PREFETCH_BLOCK_SIZE = None
|
||||||
|
|
||||||
@@ -1317,7 +1319,9 @@ def runai_safetensors_weights_iterator(
|
|||||||
mininterval=2,
|
mininterval=2,
|
||||||
)
|
)
|
||||||
|
|
||||||
yield from tensor_iter
|
for name, tensor in tensor_iter:
|
||||||
|
setattr(tensor, RUNAI_STREAMER_TENSOR_ATTR, True)
|
||||||
|
yield name, tensor
|
||||||
|
|
||||||
|
|
||||||
def set_runai_streamer_env(load_config: LoadConfig):
|
def set_runai_streamer_env(load_config: LoadConfig):
|
||||||
|
|||||||
@@ -44,7 +44,10 @@ from sglang.srt.model_loader.utils import (
|
|||||||
should_async_load,
|
should_async_load,
|
||||||
should_deepgemm_weight_requant_ue8m0,
|
should_deepgemm_weight_requant_ue8m0,
|
||||||
)
|
)
|
||||||
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.deepseek_common.utils import (
|
from sglang.srt.models.deepseek_common.utils import (
|
||||||
_is_cuda,
|
_is_cuda,
|
||||||
_is_fp8_fnuz,
|
_is_fp8_fnuz,
|
||||||
@@ -67,6 +70,12 @@ logger = logging.getLogger(__name__)
|
|||||||
NVFP4_CKPT_FP8_ATTN_QUANT_MODULES = ["q_b_proj"]
|
NVFP4_CKPT_FP8_ATTN_QUANT_MODULES = ["q_b_proj"]
|
||||||
|
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class NextNEnabledConfig:
|
class NextNEnabledConfig:
|
||||||
num_nextn_layers: int
|
num_nextn_layers: int
|
||||||
@@ -267,7 +276,9 @@ class DeepseekV2WeightLoaderMixin:
|
|||||||
if fuse_qkv_a_proj and (
|
if fuse_qkv_a_proj and (
|
||||||
"q_a_proj" in name or "kv_a_proj_with_mqa" in name
|
"q_a_proj" in name or "kv_a_proj_with_mqa" in name
|
||||||
):
|
):
|
||||||
cached_a_proj[name] = loaded_weight
|
cached_a_proj[name] = _clone_if_runai_streamed_tensor(
|
||||||
|
loaded_weight
|
||||||
|
)
|
||||||
q_a_proj_name = (
|
q_a_proj_name = (
|
||||||
name
|
name
|
||||||
if "q_a_proj" in name
|
if "q_a_proj" in name
|
||||||
|
|||||||
@@ -743,42 +743,49 @@ class KimiK25ForConditionalGeneration(nn.Module):
|
|||||||
return hidden_states
|
return hidden_states
|
||||||
|
|
||||||
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
||||||
"""Load weights for the model, separating vision and language weights"""
|
"""Stream weights, loading vision weights inline and yielding language weights.
|
||||||
|
|
||||||
|
The streaming pattern (vs accumulating into lists) is required because RunAI's
|
||||||
|
iterator reuses backing buffers — collecting tensors before consuming them
|
||||||
|
would clobber prior tensors.
|
||||||
|
"""
|
||||||
mapper = getattr(self, "hf_to_sglang_mapper", None)
|
mapper = getattr(self, "hf_to_sglang_mapper", None)
|
||||||
if mapper is not None:
|
if mapper is not None:
|
||||||
weights = mapper.apply(weights)
|
weights = mapper.apply(weights)
|
||||||
|
|
||||||
# Separate vision tower weights and language model weights
|
vision_params = (
|
||||||
vision_weights = []
|
None
|
||||||
language_weights = []
|
if self.config.language_only
|
||||||
|
else dict(self.named_parameters(remove_duplicate=False))
|
||||||
|
)
|
||||||
|
|
||||||
for name, loaded_weight in weights:
|
def stream_language_weights():
|
||||||
if "vision_tower" in name or "mm_projector" in name:
|
for name, loaded_weight in weights:
|
||||||
name = name.replace(r"wqkv.", r"attn.qkv_proj.")
|
if "vision_tower" in name or "mm_projector" in name:
|
||||||
name = name.replace(r"wo.", r"attn.proj.")
|
if vision_params is None:
|
||||||
name = name.replace("mm_projector.proj.0", "mm_projector.linear_1")
|
continue
|
||||||
name = name.replace("mm_projector.proj.2", "mm_projector.linear_2")
|
vname = (
|
||||||
vision_weights.append((name, loaded_weight))
|
name.replace(r"wqkv.", r"attn.qkv_proj.")
|
||||||
else:
|
.replace(r"wo.", r"attn.proj.")
|
||||||
name = name.replace("language_model.", "")
|
.replace("mm_projector.proj.0", "mm_projector.linear_1")
|
||||||
# All other weights go to language model
|
.replace("mm_projector.proj.2", "mm_projector.linear_2")
|
||||||
language_weights.append((name, loaded_weight))
|
)
|
||||||
|
if vname not in vision_params:
|
||||||
|
raise ValueError(f"Weight {vname} not found in params_dict")
|
||||||
|
param = vision_params[vname]
|
||||||
|
weight_loader = getattr(
|
||||||
|
param, "weight_loader", default_weight_loader
|
||||||
|
)
|
||||||
|
weight_loader(param, loaded_weight)
|
||||||
|
continue
|
||||||
|
yield name.replace("language_model.", ""), loaded_weight
|
||||||
|
|
||||||
if not self.config.language_only:
|
if self.language_model is not None:
|
||||||
# Load vision tower weights
|
self.language_model.load_weights(stream_language_weights())
|
||||||
vision_state_dict = dict(vision_weights)
|
else:
|
||||||
params_dict = dict(self.named_parameters(remove_duplicate=False))
|
# encoder-only: drain the generator so inline vision-weight loading fires.
|
||||||
for name, loaded_weight in vision_state_dict.items():
|
for _ in stream_language_weights():
|
||||||
if name not in params_dict:
|
pass
|
||||||
raise ValueError(f"Weight {name} not found in params_dict")
|
|
||||||
param = params_dict[name]
|
|
||||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
|
||||||
# loaded_weight = self._pad_vit_attn_dummy_heads(name, loaded_weight)
|
|
||||||
weight_loader(param, loaded_weight)
|
|
||||||
|
|
||||||
# Load language model weights
|
|
||||||
if not self.config.encoder_only and language_weights:
|
|
||||||
self.language_model.load_weights(language_weights)
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_model_config_for_expert_location(cls, config: KimiK25Config):
|
def get_model_config_for_expert_location(cls, config: KimiK25Config):
|
||||||
|
|||||||
@@ -0,0 +1,128 @@
|
|||||||
|
import sys
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from typing import cast
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
import sglang.srt.model_loader.loader as loader_mod
|
||||||
|
import sglang.srt.model_loader.weight_utils as weight_utils
|
||||||
|
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.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=6, suite="stage-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeModel:
|
||||||
|
def eval(self):
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
class TestRunaiModelStreamerLoader(CustomTestCase):
|
||||||
|
def test_passes_quant_config_to_model_init(self):
|
||||||
|
quant_config = object()
|
||||||
|
fake_model = _FakeModel()
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
loader_mod,
|
||||||
|
"_get_quantization_config",
|
||||||
|
return_value=quant_config,
|
||||||
|
),
|
||||||
|
patch.object(loader_mod, "_initialize_model") as mock_initialize_model,
|
||||||
|
patch.object(
|
||||||
|
loader_mod.DefaultModelLoader,
|
||||||
|
"load_weights_and_postprocess",
|
||||||
|
) as mock_load_weights,
|
||||||
|
):
|
||||||
|
mock_initialize_model.return_value = fake_model
|
||||||
|
runai_loader = loader_mod.RunaiModelStreamerLoader(
|
||||||
|
LoadConfig(
|
||||||
|
load_format=LoadFormat.RUNAI_STREAMER,
|
||||||
|
model_loader_extra_config={},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
model_config = cast(
|
||||||
|
ModelConfig,
|
||||||
|
SimpleNamespace(dtype=torch.float16, modelopt_quant=False),
|
||||||
|
)
|
||||||
|
|
||||||
|
model = runai_loader.load_model(
|
||||||
|
model_config=model_config,
|
||||||
|
device_config=DeviceConfig("cpu"),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIs(model, fake_model)
|
||||||
|
self.assertIs(mock_load_weights.call_args.args[0], fake_model)
|
||||||
|
self.assertIs(mock_initialize_model.call_args.args[2], quant_config)
|
||||||
|
|
||||||
|
def test_marks_streamer_tensors(self):
|
||||||
|
source_tensor = torch.tensor([1], dtype=torch.int32)
|
||||||
|
|
||||||
|
class FakeStreamer:
|
||||||
|
def __enter__(self):
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __exit__(self, *_args):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def stream_files(self, *_args, **_kwargs):
|
||||||
|
self.files_to_tensors_metadata = {0: [object()]}
|
||||||
|
|
||||||
|
def get_tensors(self):
|
||||||
|
yield "weight", source_tensor
|
||||||
|
|
||||||
|
with patch.dict(
|
||||||
|
sys.modules,
|
||||||
|
{"runai_model_streamer": SimpleNamespace(SafetensorsStreamer=FakeStreamer)},
|
||||||
|
):
|
||||||
|
weights = list(
|
||||||
|
weight_utils.runai_safetensors_weights_iterator(["model.safetensors"])
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(weights[0][0], "weight")
|
||||||
|
self.assertTrue(getattr(weights[0][1], weight_utils.RUNAI_STREAMER_TENSOR_ATTR))
|
||||||
|
|
||||||
|
def test_deepseek_clone_only_clones_marked_tensors(self):
|
||||||
|
unmarked = torch.tensor([1], dtype=torch.int32)
|
||||||
|
|
||||||
|
self.assertIs(
|
||||||
|
deepseek_weight_loader._clone_if_runai_streamed_tensor(unmarked),
|
||||||
|
unmarked,
|
||||||
|
)
|
||||||
|
|
||||||
|
marked = torch.tensor([1], dtype=torch.int32)
|
||||||
|
setattr(marked, weight_utils.RUNAI_STREAMER_TENSOR_ATTR, True)
|
||||||
|
|
||||||
|
cloned = deepseek_weight_loader._clone_if_runai_streamed_tensor(marked)
|
||||||
|
|
||||||
|
self.assertIsNot(cloned, marked)
|
||||||
|
marked.fill_(2)
|
||||||
|
self.assertEqual(cloned.item(), 1)
|
||||||
|
|
||||||
|
def test_get_model_loader_uses_runai_for_prequantized_modelopt(self):
|
||||||
|
load_config = LoadConfig(
|
||||||
|
load_format=LoadFormat.RUNAI_STREAMER,
|
||||||
|
model_loader_extra_config={},
|
||||||
|
)
|
||||||
|
model_config = cast(
|
||||||
|
ModelConfig,
|
||||||
|
SimpleNamespace(
|
||||||
|
quantization="modelopt_fp4",
|
||||||
|
modelopt_quant=False,
|
||||||
|
_is_already_quantized=lambda: True,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
model_loader = loader_mod.get_model_loader(load_config, model_config)
|
||||||
|
|
||||||
|
self.assertIsInstance(model_loader, loader_mod.RunaiModelStreamerLoader)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user