[RL] support weight reload for low-bit rollout (#9650)

Co-authored-by: Hecate0821 <hec4te0821@gmail.com>
Co-authored-by: eternally-z <zzywzj@gmail.com>
Co-authored-by: Wilboludriver <wilbolu@outlook.com>
Co-authored-by: Wilbolu <81792854+Wilboludriver@users.noreply.github.com>
Co-authored-by: Ke Bao <ispobaoke@gmail.com>
This commit is contained in:
Peng Zhang
2025-12-10 15:44:01 +08:00
committed by GitHub
co-authored by Hecate0821 eternally-z Wilboludriver Wilbolu Ke Bao
parent b0a25d0913
commit 21028b5507
7 changed files with 581 additions and 4 deletions
+23 -1
View File
@@ -495,7 +495,8 @@ class Qwen3ForCausalLM(nn.Module):
def end_layer(self):
return self.model.end_layer
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
def _load_weights_impl(self, weights: Iterable[Tuple[str, torch.Tensor]]):
"""Internal implementation of weight loading without reload scenario handling."""
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
("qkv_proj", "q_proj", "q"),
@@ -506,6 +507,7 @@ class Qwen3ForCausalLM(nn.Module):
]
params_dict = dict(self.named_parameters())
updated_params = set()
for name, loaded_weight in weights:
if "Embedding" in self.config.name_or_path:
name = add_prefix(name, "model")
@@ -552,6 +554,7 @@ class Qwen3ForCausalLM(nn.Module):
param = params_dict[name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
updated_params.add(name)
break
else:
# Skip loading extra bias for GPTQ models.
@@ -564,9 +567,28 @@ class Qwen3ForCausalLM(nn.Module):
param, "weight_loader", default_weight_loader
)
weight_loader(param, loaded_weight)
updated_params.add(name)
else:
logger.warning(f"Parameter {name} not found in params_dict")
return updated_params
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
"""Load weights into the model, with support for RL training reload scenarios."""
from sglang.srt.model_loader.loader import QuantizedRLModelLoader
# Check if this is a reload scenario for RL training with quantized models
is_reload = QuantizedRLModelLoader.is_reload_scenario(self)
if is_reload:
# Use the fast path for RL training reloads
logger.info("[QuantizedRL] Using fast path reload in load_weights")
QuantizedRLModelLoader.rebinding_and_load_weights(
self, self._load_weights_impl, weights
)
else:
# Standard weight loading path
self._load_weights_impl(weights)
def get_embed_and_head(self):
return self.model.embed_tokens.weight, self.lm_head.weight