Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
109 changes: 95 additions & 14 deletions verl_speco/integration/vllm_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -1976,6 +1976,13 @@ def _speco_add_vllm_spec_decode_extra_fields(

async def launch_server(self, *args, **kwargs):
self._speco_vllm_spec_decode_pending_stats = _new_vllm_spec_decode_stats()
drafter_cfg = _load_env_drafter_config()
self._speco_initial_draft_weights_required = bool(
drafter_cfg.get("enable")
and _speculative_method_from_drafter(drafter_cfg) in {"dflash", "dspark"}
)
self._speco_initial_draft_weights_ready = False
self._speco_initial_draft_weights_lock = None
install_vllm_runtime_observability()
_ensure_vllm_drafter_speculative_config_from_env(self.config)
return await super().launch_server(*args, **kwargs)
Expand Down Expand Up @@ -2015,7 +2022,34 @@ def from_vllm_config_with_speco_stats(cls, *call_args, **call_kwargs):
finally:
AsyncLLM.from_vllm_config = original_from_vllm_config_attr

async def _speco_ensure_initial_draft_weights(self) -> None:
"""Initialize the serving drafter before admitting the first request."""
if not bool(getattr(self, "_speco_initial_draft_weights_required", False)):
return
if bool(getattr(self, "_speco_initial_draft_weights_ready", False)):
return

import asyncio

lock = getattr(self, "_speco_initial_draft_weights_lock", None)
if lock is None:
lock = asyncio.Lock()
self._speco_initial_draft_weights_lock = lock

async with lock:
if bool(getattr(self, "_speco_initial_draft_weights_ready", False)):
return
collective_rpc = getattr(self, "collective_rpc", None)
if not callable(collective_rpc):
raise RuntimeError(
"vLLM HTTP server does not expose collective_rpc for "
"initial drafter weight loading"
)
await collective_rpc("speco_ensure_draft_initialized")
self._speco_initial_draft_weights_ready = True

async def generate(self, *args, **kwargs):
await self._speco_ensure_initial_draft_weights()
output = await super().generate(*args, **kwargs)
extra_fields = getattr(output, "extra_fields", None)
if isinstance(extra_fields, dict):
Expand Down Expand Up @@ -2438,6 +2472,7 @@ class SpecoVLLMColocateWorkerExtension(_VLLMWorkerExtensionBase):
"""vLLM worker extension that can update only the speculative draft model."""

_speco_draft_level2_snapshot: dict[str, Any] | None = None
_speco_draft_level2_snapshot_source: str | None = None

def __new__(cls, **kwargs):
try:
Expand Down Expand Up @@ -2486,6 +2521,11 @@ def _speco_wake_up_hook(*args, **kwargs):
)
return result

if getattr(instance, "_speco_draft_weight_source", None) == "online":
raise RuntimeError(
"Cannot restore the online drafter after level-2 wake-up: "
"the online weight snapshot is missing"
)
reloaded = instance._speco_reload_draft_from_checkpoint()
if reloaded > 0:
logger.warning(
Expand Down Expand Up @@ -2560,6 +2600,15 @@ def _speco_snapshot_draft_for_level2(self) -> int:
target model from the actor. Snapshotting the draft preserves the
latest online-published state for every speculative method.
"""
source = getattr(self, "_speco_draft_weight_source", None)
if source not in {"checkpoint", "online"}:
if self._speco_is_dflash_draft():
return 0
# Other speculative methods use vLLM's native checkpoint loading,
# which is already valid before this extension observes the model.
source = "checkpoint"
self._speco_draft_weight_source = source

draft_model, _ = self._speco_resolve_draft_model()
if draft_model is None:
return 0
Expand All @@ -2576,12 +2625,18 @@ def _speco_snapshot_draft_for_level2(self) -> int:
return 0

self._speco_draft_level2_snapshot = snapshot
self._speco_draft_level2_snapshot_source = source
return len(snapshot)

def _speco_restore_draft_after_level2(self) -> int:
snapshot = getattr(self, "_speco_draft_level2_snapshot", None)
if snapshot is None:
return 0
snapshot_source = getattr(self, "_speco_draft_level2_snapshot_source", None)
if snapshot_source not in {"checkpoint", "online"}:
raise RuntimeError(
"Cannot restore the draft level-2 snapshot: weight source is unknown"
)

draft_model, _ = self._speco_resolve_draft_model()
if draft_model is None:
Expand Down Expand Up @@ -2619,8 +2674,10 @@ def _speco_restore_draft_after_level2(self) -> int:
)

self._speco_rebuild_draft_metadata_buffers(draft_model)
self._speco_draft_weight_source = snapshot_source
restored = len(snapshot)
self._speco_draft_level2_snapshot = None
self._speco_draft_level2_snapshot_source = None
return restored

@staticmethod
Expand Down Expand Up @@ -2821,6 +2878,7 @@ def on_bucket_received(bucket_weights):
"[speco draft update] _build_fused_kv_buffers failed: %s", exc
)

self._speco_draft_weight_source = "online"
self._speco_diag_draft_state("after_draft_ipc_update")
# One-time diagnostic: check whether probabilistic sampling is active
proposer = self._speco_resolve_draft_proposer()
Expand Down Expand Up @@ -2852,7 +2910,7 @@ def on_bucket_received(bucket_weights):
)

# ----------------------------------------------------------------
# Fix: reload DFlash drafter weights from checkpoint after wake_up
# Initial DFlash/DSpark checkpoint load and level-2 wake-up fallback
# ----------------------------------------------------------------

def _speco_get_draft_checkpoint_path(self) -> str | None:
Expand All @@ -2872,10 +2930,10 @@ def _speco_get_draft_checkpoint_path(self) -> str | None:
return getattr(draft_model_cfg, "model", None)

def _speco_reload_draft_from_checkpoint(self) -> int:
"""Reload DFlash drafter weights from its checkpoint (safetensors).
"""Reload DFlash/DSpark drafter weights from checkpoint (safetensors).

Called after target model wake_up to restore drafter weights that were
lost during sleep(level=2). Returns the number of weight tensors loaded.
Used before the first serving request and as the fallback after
sleep(level=2). Returns the number of weight tensors loaded.
"""
import glob as _glob

Expand Down Expand Up @@ -2919,25 +2977,52 @@ def _speco_reload_draft_from_checkpoint(self) -> int:

try:
draft_model.load_weights(iter(weights_iter))
self._speco_rebuild_draft_metadata_buffers(draft_model)
loaded_count = len(weights_iter)
except Exception as exc:
logger.warning("[speco draft reload] load_weights failed: %s", exc)
return 0

self._speco_draft_weight_source = "checkpoint"
return loaded_count

def speco_ensure_draft_initialized(self) -> dict[str, Any]:
"""Load the base drafter once, before its first serving request."""
source = getattr(self, "_speco_draft_weight_source", None)
if source in {"checkpoint", "online"}:
return {"initialized": True, "source": source, "loaded_params": 0}
if not self._speco_is_dflash_draft():
return {
"initialized": False,
"source": "not_applicable",
"loaded_params": 0,
}

loaded_params = self._speco_reload_draft_from_checkpoint()
if loaded_params <= 0:
raise RuntimeError(
"Failed to initialize the serving drafter from its configured "
"checkpoint before the first rollout request"
)
self._speco_draft_weight_source = "checkpoint"
return {
"initialized": True,
"source": "checkpoint",
"loaded_params": loaded_params,
}

def update_weights_from_ipc(
self,
peft_config: dict | None = None,
base_sync_done=False,
use_shm: bool = False,
):
"""Override target weight sync to also reload drafter from checkpoint."""
"""Sync target weights without replacing the current speculative drafter."""
patch_verl_bucketed_weight_transfer_rebuild_ipc()
patch_verl_bucketed_weight_transfer_shm_reuse()
patch_verl_bucketed_weight_transfer_npu_staging()
is_npu = _speco_is_npu_vllm_worker(self)
# Diagnostic: check draft state BEFORE target sync
# Diagnostic: check draft state BEFORE target sync.
self._speco_diag_draft_state("before_target_sync")
try:
with _speco_npu_target_staging(
Expand All @@ -2950,14 +3035,10 @@ def update_weights_from_ipc(
)
# Diagnostic: check draft state AFTER target sync (may be zeroed by wake_up)
self._speco_diag_draft_state("after_target_sync")
reloaded = self._speco_reload_draft_from_checkpoint()
if reloaded > 0:
logger.warning(
"[speco draft reload] drafter weights restored after target sync (%d tensors)",
reloaded,
)
# Diagnostic: check draft state AFTER reload
self._speco_diag_draft_state("after_draft_reload")
# Do not reload the configured checkpoint here. The wake-up hook
# above already restores a level-2 snapshot (or falls back to the
# checkpoint when no snapshot exists). Reloading again after target
# sync would overwrite the most recently published online drafter.
return result
finally:
if is_npu:
Expand Down
Loading