[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:
co-authored by
Hecate0821
eternally-z
Wilboludriver
Wilbolu
Ke Bao
parent
b0a25d0913
commit
21028b5507
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user