Fix weight sync selector for frozen speculative drafts (#1926)

Co-authored-by: Zhiyao Jiang <jessicajiang324@gmail.com>
This commit is contained in:
Xinyu Jiang
2026-07-29 17:34:54 -07:00
committed by GitHub
co-authored by Zhiyao Jiang
parent fa255567f7
commit 310ec07c0a
5 changed files with 40 additions and 9 deletions
@@ -406,9 +406,20 @@ def collect_named_tensors_for_weight_transfer(
yield name, tensor
def begin_weight_update(rollout_engines: Sequence[ActorHandle]):
"""Open a weight-update session on all rollout engines (restore packed weights)."""
ray.get([engine.begin_weight_update.remote() for engine in rollout_engines])
def begin_weight_update(rollout_engines: Sequence[ActorHandle], selector: str = "all"):
"""Open a weight-update session on the selected rollout engines (restore packed weights)."""
ray.get([engine.begin_weight_update.remote(selector=selector) for engine in rollout_engines])
def weight_update_selector(args) -> str:
"""Exclude the draft only when the trainer provably has no MTP block to send it."""
if (
getattr(args, "sglang_speculative_algorithm", None)
and not getattr(args, "mtp_num_layers", None)
and getattr(args, "megatron_to_hf_mode", "raw") != "bridge"
):
return "target"
return "all"
def end_weight_update(rollout_engines: Sequence[ActorHandle]):
@@ -115,6 +115,7 @@ class UpdateWeightFromDistributed(DistBucketedWeightUpdateMixin):
self.weight_version,
self.rollout_engines,
converted_named_tensors,
selector=self._weight_update_selector,
)
ray.get(refs)
converted_named_tensors.clear()
@@ -258,6 +259,7 @@ def update_weights_from_distributed(
weight_version: int,
rollout_engines: Sequence[ActorHandle],
converted_named_tensors: Sequence[tuple[str, torch.Tensor]],
selector: str = "all",
) -> list[ObjectRef]:
"""
Send metadata (Ray), broadcast tensors (NCCL rank 0 → engines).
@@ -267,6 +269,7 @@ def update_weights_from_distributed(
names=[name for name, _ in converted_named_tensors],
dtypes=[param.dtype for _, param in converted_named_tensors],
shapes=[param.shape for _, param in converted_named_tensors],
selector=selector,
group_name=group_name,
weight_version=str(weight_version),
)
@@ -21,6 +21,7 @@ from ..common import (
end_weight_update,
get_atomic_update_groups,
get_named_value_update_units,
weight_update_selector,
)
from ..hf_weight_iterator_base import HfWeightIteratorBase
@@ -306,13 +307,14 @@ class DistBucketedWeightUpdateMixin:
def _pause_and_prepare_engines(self) -> None:
"""Pause rollout engines, flush cache, and open the weight-update session."""
self._weight_update_selector = weight_update_selector(self.args)
if dist.get_rank() == 0:
mode = self.args.pause_generation_mode
ray.get([engine.pause_generation.remote(mode=mode) for engine in self.rollout_engines])
if mode != "in_place":
ray.get([engine.flush_cache.remote() for engine in self.rollout_engines])
begin_weight_update(self.rollout_engines)
begin_weight_update(self.rollout_engines, self._weight_update_selector)
def _finalize_and_resume_engines(self) -> None:
"""Close the weight-update session and resume rollout engines."""
@@ -20,7 +20,7 @@ from miles.utils.distributed_utils import get_gloo_group
from miles.utils.lora import LORA_ADAPTER_NAME
from ..sglang import FlattenedTensorBucket, MultiprocessingSerializer
from .common import _check_weight_sync_results, begin_weight_update, end_weight_update
from .common import _check_weight_sync_results, begin_weight_update, end_weight_update, weight_update_selector
from .hf_weight_iterator_base import HfWeightIteratorBase
from .update_weight_from_distributed.broadcast import (
connect_rollout_engines_from_distributed,
@@ -215,7 +215,7 @@ class UpdateWeightFromTensor:
ray.get([engine.pause_generation.remote(mode=mode) for engine in self.rollout_engines])
ray.get([engine.flush_cache.remote() for engine in self.rollout_engines])
if not skip_base_sync:
begin_weight_update(self.rollout_engines)
begin_weight_update(self.rollout_engines, weight_update_selector(self.args))
dist.barrier(group=get_gloo_group())
megatron_local_weights = self.weights_getter()
@@ -269,6 +269,7 @@ class UpdateWeightFromTensor:
ipc_engine=self._ipc_engine,
ipc_gather_src=self._ipc_gather_src,
ipc_gather_group=self._ipc_gather_group,
selector=weight_update_selector(self.args),
weight_version=self.weight_version,
)
if self.use_distribute and self._is_distributed_src_rank:
@@ -278,6 +279,7 @@ class UpdateWeightFromTensor:
self.weight_version,
self.distributed_rollout_engines,
hf_named_tensors,
selector=weight_update_selector(self.args),
)
if refs_distributed:
refs = (refs or []) + refs_distributed
@@ -297,6 +299,7 @@ class UpdateWeightFromTensor:
ipc_engine=self._ipc_engine,
ipc_gather_src=self._ipc_gather_src,
ipc_gather_group=self._ipc_gather_group,
selector=weight_update_selector(self.args),
lora_config=self._lora_config,
lora_name=LORA_ADAPTER_NAME,
lora_loaded=self._lora_loaded,
@@ -317,6 +320,7 @@ def _send_to_colocated_engine(
lora_name: str | None = None,
lora_loaded: bool = False,
check_equal: bool = False,
selector: str = "all",
) -> tuple[list[ObjectRef], Any]:
# Placeholder ranks (GPU slots reserved but no engine) have no gather group.
# gather_object is only collective among group members, so we skip entirely.
@@ -394,6 +398,7 @@ def _send_to_colocated_engine(
"serialized_named_tensors": [tensors[i] for tensors in serialized_named_tensors],
"load_format": "flattened_bucket",
"weight_version": str(weight_version),
"selector": selector,
}
refs.append(ipc_engine.update_weights_from_tensor.remote(**kwargs))
+13 -3
View File
@@ -319,6 +319,7 @@ class SGLangEngine(RayActor):
load_format: str | None = None,
flush_cache: bool = False,
weight_version: str | None = None,
selector: str = "all",
):
"""
Update model weights from tensor data. The HTTP server will only post meta data, and the real weights will be copied directly from GPUs.
@@ -330,6 +331,7 @@ class SGLangEngine(RayActor):
"serialized_named_tensors": serialized_named_tensors,
"load_format": load_format,
"flush_cache": flush_cache,
"selector": selector,
}
if weight_version is not None:
payload["weight_version"] = weight_version
@@ -584,7 +586,14 @@ class SGLangEngine(RayActor):
pass
def update_weights_from_distributed(
self, names, dtypes, shapes, group_name, flush_cache=False, weight_version: str | None = None
self,
names,
dtypes,
shapes,
group_name,
flush_cache=False,
weight_version: str | None = None,
selector: str = "all",
):
payload = {
"names": names,
@@ -592,6 +601,7 @@ class SGLangEngine(RayActor):
"shapes": shapes,
"group_name": group_name,
"flush_cache": flush_cache,
"selector": selector,
}
if weight_version is not None:
payload["weight_version"] = weight_version
@@ -613,9 +623,9 @@ class SGLangEngine(RayActor):
response.raise_for_status()
return response
def begin_weight_update(self):
def begin_weight_update(self, selector: str = "all"):
"""Open a weight-update session on the engine (restores packed weights for loading)."""
return self._make_request("begin_weight_update", {})
return self._make_request("begin_weight_update", {"selector": selector})
def end_weight_update(self):
"""Close the weight-update session (post-load + quant post-process on the full model)."""