diff --git a/python/sglang/srt/managers/scheduler_update_weights_mixin.py b/python/sglang/srt/managers/scheduler_update_weights_mixin.py index fafdd48f4..fa0d612e2 100644 --- a/python/sglang/srt/managers/scheduler_update_weights_mixin.py +++ b/python/sglang/srt/managers/scheduler_update_weights_mixin.py @@ -76,7 +76,8 @@ class SchedulerUpdateWeightsMixin: def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput): """Update the online model parameter from tensors.""" - success, message = self.tp_worker.update_weights_from_tensor(recv_req) + worker = self.draft_worker or self.tp_worker + success, message = worker.update_weights_from_tensor(recv_req) # TODO extract common code b/t update_weights_from_distributed and update_weights_from_tensor later if success: if recv_req.flush_cache: diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index 08e6516bc..60cb81125 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -9,6 +9,7 @@ from sglang.srt.layers.dp_attention import get_attention_tp_group from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.moe.utils import speculative_moe_backend_context from sglang.srt.layers.sampler import get_token_ids_logprobs, get_top_logprobs +from sglang.srt.managers.io_struct import UpdateWeightsFromTensorReqInput from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.tp_worker import TpModelWorker @@ -50,6 +51,7 @@ from sglang.srt.speculative.spec_utils import ( select_top_k_tokens, ) from sglang.srt.utils import ( + MultiprocessingSerializer, empty_context, get_available_gpu_memory, get_bool_env_var, @@ -57,6 +59,7 @@ from sglang.srt.utils import ( is_npu, next_power_of_2, ) +from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions _is_npu = is_npu() @@ -984,6 +987,24 @@ class EAGLEWorker(TpModelWorker): draft_input.topk_p, draft_input.topk_index = fast_topk(probs, self.topk, dim=-1) draft_input.hidden_states = logits_output.hidden_states + def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput): + monkey_patch_torch_reductions() + named_tensors = MultiprocessingSerializer.deserialize( + recv_req.serialized_named_tensors[self.tp_rank] + ) + success, message = self.model_runner.update_weights_from_tensor( + named_tensors=named_tensors, + load_format=recv_req.load_format, + ) + if not success: + return success, message + + success, message = self.target_worker.model_runner.update_weights_from_tensor( + named_tensors=named_tensors, + load_format=recv_req.load_format, + ) + return success, message + @torch.compile(dynamic=True, disable=_is_npu) def get_last_loc_large_page_size_top_k_1(