diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 665c551da..251922a2e 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -7,6 +7,7 @@ import torch from sglang.srt.environ import envs from sglang.srt.layers.moe.utils import speculative_moe_backend_context +from sglang.srt.managers.io_struct import UpdateWeightsFromTensorReqInput from sglang.srt.managers.schedule_batch import ModelWorkerBatch from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.tp_worker import TpModelWorker @@ -40,12 +41,14 @@ from sglang.srt.speculative.spec_utils import ( load_token_map, ) from sglang.srt.utils.common import ( + MultiprocessingSerializer, empty_context, fast_topk, get_available_gpu_memory, is_npu, next_power_of_2, ) +from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions _is_npu = is_npu() @@ -553,6 +556,7 @@ class EAGLEWorkerV2(BaseSpecWorker): self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens self.enable_nan_detection = server_args.enable_nan_detection + self.tp_rank = tp_rank self.gpu_id = gpu_id self.device = server_args.device self._target_worker = target_worker @@ -787,3 +791,21 @@ class EAGLEWorkerV2(BaseSpecWorker): self.token_to_kv_pool_allocator.get_kvcache().move_kv_cache( tgt_cache_loc, accepted_out_cache_loc ) + + 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.draft_worker.draft_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