[sgl] add support for weight update function in spedec (#22088)
This commit is contained in:
@@ -48,11 +48,13 @@ class SchedulerUpdateWeightsMixin:
|
|||||||
):
|
):
|
||||||
"""In-place update of the weights from disk."""
|
"""In-place update of the weights from disk."""
|
||||||
success, message = self.tp_worker.update_weights_from_disk(recv_req)
|
success, message = self.tp_worker.update_weights_from_disk(recv_req)
|
||||||
if success:
|
tp_success = success
|
||||||
if recv_req.flush_cache:
|
if success and self.draft_worker is not None:
|
||||||
flush_cache_success = self.flush_cache()
|
success, message = self.draft_worker.update_weights_from_disk(recv_req)
|
||||||
assert flush_cache_success, "Cache flush failed after updating weights"
|
if tp_success and recv_req.flush_cache:
|
||||||
else:
|
flush_cache_success = self.flush_cache()
|
||||||
|
assert flush_cache_success, "Cache flush failed after updating weights"
|
||||||
|
if not success:
|
||||||
logger.error(message)
|
logger.error(message)
|
||||||
return UpdateWeightFromDiskReqOutput(success, message, 0)
|
return UpdateWeightFromDiskReqOutput(success, message, 0)
|
||||||
|
|
||||||
@@ -108,11 +110,13 @@ class SchedulerUpdateWeightsMixin:
|
|||||||
):
|
):
|
||||||
"""Update the online model parameter from IPC for checkpoint-engine integration."""
|
"""Update the online model parameter from IPC for checkpoint-engine integration."""
|
||||||
success, message = self.tp_worker.update_weights_from_ipc(recv_req)
|
success, message = self.tp_worker.update_weights_from_ipc(recv_req)
|
||||||
if success:
|
tp_success = success
|
||||||
if recv_req.flush_cache:
|
if success and self.draft_worker is not None:
|
||||||
flush_cache_success = self.flush_cache()
|
success, message = self.draft_worker.update_weights_from_ipc(recv_req)
|
||||||
assert flush_cache_success, "Cache flush failed after updating weights"
|
if tp_success and recv_req.flush_cache:
|
||||||
else:
|
flush_cache_success = self.flush_cache()
|
||||||
|
assert flush_cache_success, "Cache flush failed after updating weights"
|
||||||
|
if not success:
|
||||||
logger.error(message)
|
logger.error(message)
|
||||||
torch.distributed.barrier(group=self.tp_cpu_group)
|
torch.distributed.barrier(group=self.tp_cpu_group)
|
||||||
return UpdateWeightsFromIPCReqOutput(success, message)
|
return UpdateWeightsFromIPCReqOutput(success, message)
|
||||||
|
|||||||
@@ -22,7 +22,11 @@ from sglang.srt.layers.moe.utils import (
|
|||||||
speculative_moe_backend_context,
|
speculative_moe_backend_context,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs
|
from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs
|
||||||
from sglang.srt.managers.io_struct import UpdateWeightsFromTensorReqInput
|
from sglang.srt.managers.io_struct import (
|
||||||
|
UpdateWeightFromDiskReqInput,
|
||||||
|
UpdateWeightsFromIPCReqInput,
|
||||||
|
UpdateWeightsFromTensorReqInput,
|
||||||
|
)
|
||||||
from sglang.srt.managers.schedule_batch import ModelWorkerBatch
|
from sglang.srt.managers.schedule_batch import ModelWorkerBatch
|
||||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
from sglang.srt.managers.scheduler import GenerationBatchResult
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||||
@@ -978,6 +982,24 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
tgt_cache_loc, accepted_out_cache_loc
|
tgt_cache_loc, accepted_out_cache_loc
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput):
|
||||||
|
success, message = self._draft_worker.draft_runner.update_weights_from_disk(
|
||||||
|
recv_req.model_path,
|
||||||
|
recv_req.load_format,
|
||||||
|
recapture_cuda_graph=recv_req.recapture_cuda_graph,
|
||||||
|
)
|
||||||
|
if not success:
|
||||||
|
return success, message
|
||||||
|
return True, "Succeeded to update model weights."
|
||||||
|
|
||||||
|
def update_weights_from_ipc(self, recv_req: UpdateWeightsFromIPCReqInput):
|
||||||
|
success, message = self._draft_worker.draft_runner.update_weights_from_ipc(
|
||||||
|
recv_req
|
||||||
|
)
|
||||||
|
if not success:
|
||||||
|
return success, message
|
||||||
|
return True, "Succeeded to update model weights."
|
||||||
|
|
||||||
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
|
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
|
||||||
monkey_patch_torch_reductions()
|
monkey_patch_torch_reductions()
|
||||||
named_tensors = MultiprocessingSerializer.deserialize(
|
named_tensors = MultiprocessingSerializer.deserialize(
|
||||||
|
|||||||
@@ -21,6 +21,10 @@ import torch
|
|||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.moe.utils import speculative_moe_backend_context
|
from sglang.srt.layers.moe.utils import speculative_moe_backend_context
|
||||||
from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs
|
from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs
|
||||||
|
from sglang.srt.managers.io_struct import (
|
||||||
|
UpdateWeightFromDiskReqInput,
|
||||||
|
UpdateWeightsFromIPCReqInput,
|
||||||
|
)
|
||||||
from sglang.srt.managers.schedule_batch import ModelWorkerBatch
|
from sglang.srt.managers.schedule_batch import ModelWorkerBatch
|
||||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
from sglang.srt.managers.scheduler import GenerationBatchResult
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||||
@@ -770,3 +774,25 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
next_draft_input=next_draft_input,
|
next_draft_input=next_draft_input,
|
||||||
accept_lens=accept_length,
|
accept_lens=accept_length,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput):
|
||||||
|
for i in range(self.speculative_num_steps):
|
||||||
|
success, message = self._draft_worker.draft_runner_list[
|
||||||
|
i
|
||||||
|
].update_weights_from_disk(
|
||||||
|
recv_req.model_path,
|
||||||
|
recv_req.load_format,
|
||||||
|
recapture_cuda_graph=recv_req.recapture_cuda_graph,
|
||||||
|
)
|
||||||
|
if not success:
|
||||||
|
return success, message
|
||||||
|
return True, "Succeeded to update model weights."
|
||||||
|
|
||||||
|
def update_weights_from_ipc(self, recv_req: UpdateWeightsFromIPCReqInput):
|
||||||
|
for i in range(self.speculative_num_steps):
|
||||||
|
success, message = self._draft_worker.draft_runner_list[
|
||||||
|
i
|
||||||
|
].update_weights_from_ipc(recv_req)
|
||||||
|
if not success:
|
||||||
|
return success, message
|
||||||
|
return True, "Succeeded to update model weights."
|
||||||
|
|||||||
Reference in New Issue
Block a user