Fix gpt-oss RunAI streamer weight ownership (#38908)
Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5.1
parent
bd45cd50ca
commit
fd32226706
@@ -66,7 +66,10 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
|
||||
get_tc_piecewise_forward_context,
|
||||
is_in_tc_piecewise_cuda_graph,
|
||||
)
|
||||
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.utils import (
|
||||
create_fused_set_kv_buffer_arg,
|
||||
enable_fused_set_kv_buffer,
|
||||
@@ -953,20 +956,25 @@ class GptOssForCausalLM(nn.Module):
|
||||
)
|
||||
|
||||
def _load_weights_mxfp4(self, weights, is_nextn, weight_name_mapping):
|
||||
mxfp4_weights = []
|
||||
normal_weights = []
|
||||
|
||||
for name, weight in weights:
|
||||
if (
|
||||
".experts" in name
|
||||
and self.quant_config is not None
|
||||
and self.quant_config.get_name() == "mxfp4"
|
||||
):
|
||||
mxfp4_weights.append((name, weight))
|
||||
else:
|
||||
normal_weights.append((name, weight))
|
||||
def experts(weights):
|
||||
# The RunAI streamer reuses one staging buffer across tensors, so a
|
||||
# tensor read after later ones arrive can be read back as garbage.
|
||||
# Expert weights are copied into their parameter as they are
|
||||
# yielded; the rest are held until afterwards and need their own
|
||||
# memory.
|
||||
for name, weight in weights:
|
||||
if (
|
||||
".experts" in name
|
||||
and self.quant_config is not None
|
||||
and self.quant_config.get_name() == "mxfp4"
|
||||
):
|
||||
yield name, weight
|
||||
else:
|
||||
normal_weights.append((name, _own_if_runai_streamed(weight)))
|
||||
|
||||
mxfp4_loaded_params = self._load_mxfp4_experts_weights(mxfp4_weights)
|
||||
mxfp4_loaded_params = self._load_mxfp4_experts_weights(experts(weights))
|
||||
self._load_normal_weights(
|
||||
normal_weights,
|
||||
is_nextn=is_nextn,
|
||||
@@ -1379,6 +1387,18 @@ class GptOssForCausalLM(nn.Module):
|
||||
return get_attention_sliding_window_size(self.config)
|
||||
|
||||
|
||||
def _own_if_runai_streamed(tensor: torch.Tensor) -> torch.Tensor:
|
||||
"""Take a copy the streamer cannot overwrite.
|
||||
|
||||
The copy lands on the host: distributed streaming yields device tensors,
|
||||
and these are held until the whole checkpoint has streamed, so cloning
|
||||
them in place would add their own GiB to peak GPU usage.
|
||||
"""
|
||||
if getattr(tensor, RUNAI_STREAMER_TENSOR_ATTR, False):
|
||||
return tensor.detach().to("cpu", copy=True)
|
||||
return tensor
|
||||
|
||||
|
||||
def _canonicalize_weights(config, weights_in: Iterable[Tuple[str, torch.Tensor]]):
|
||||
weights_out_dict = dict(weights_in)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user