diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 18ee1f6e0..c17db0bd1 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -718,11 +718,6 @@ async def update_weights_from_disk(obj: UpdateWeightFromDiskReqInput, request: R await _global_state.tokenizer_manager.update_weights_from_disk(obj, request) ) - # Update weight version if provided and weights update was successful - if success and obj.weight_version is not None: - _update_weight_version_if_provided(obj.weight_version) - message += f" Weight version updated to {obj.weight_version}." - content = { "success": success, "message": message, @@ -816,11 +811,6 @@ async def update_weights_from_tensor( obj, request ) - # Update weight version if provided and weights update was successful - if success and obj.weight_version is not None: - _update_weight_version_if_provided(obj.weight_version) - message += f" Weight version updated to {obj.weight_version}." - content = {"success": success, "message": message} return ORJSONResponse( content, status_code=200 if success else HTTPStatus.BAD_REQUEST @@ -838,11 +828,6 @@ async def update_weights_from_distributed( ) ) - # Update weight version if provided and weights update was successful - if success and obj.weight_version is not None: - _update_weight_version_if_provided(obj.weight_version) - message += f" Weight version updated to {obj.weight_version}." - content = {"success": success, "message": message} if success: return ORJSONResponse(content, status_code=200) @@ -857,11 +842,6 @@ async def update_weights_from_ipc(obj: UpdateWeightsFromIPCReqInput, request: Re obj, request ) - # Update weight version if provided and weights update was successful - if success and obj.weight_version is not None: - _update_weight_version_if_provided(obj.weight_version) - message += f" Weight version updated to {obj.weight_version}." - content = {"success": success, "message": message} if success: if _global_state.tokenizer_manager.initial_weights_loaded is False: @@ -1325,12 +1305,6 @@ async def vertex_generate(vertex_req: VertexGenerateReqInput, raw_request: Reque return ORJSONResponse({"predictions": ret}) -def _update_weight_version_if_provided(weight_version: Optional[str]) -> None: - """Update weight version if provided.""" - if weight_version is not None: - _global_state.tokenizer_manager.server_args.weight_version = weight_version - - def _create_error_response(e): return ORJSONResponse( {"error": {"message": str(e)}}, status_code=HTTPStatus.BAD_REQUEST diff --git a/python/sglang/srt/managers/tokenizer_communicator_mixin.py b/python/sglang/srt/managers/tokenizer_communicator_mixin.py index 7f49e48c4..9205db0a9 100644 --- a/python/sglang/srt/managers/tokenizer_communicator_mixin.py +++ b/python/sglang/srt/managers/tokenizer_communicator_mixin.py @@ -406,7 +406,13 @@ class TokenizerCommunicatorMixin: # cannot run while requests are in progress. async with self.model_update_lock.writer_lock: results = await self.update_weights_from_distributed_communicator(obj) - return _Communicator.merge_results(results) + success, message = _Communicator.merge_results(results) + + if success and obj.weight_version is not None: + self._update_weight_version_if_provided(obj.weight_version) + message += f" Weight version updated to {obj.weight_version}." + + return success, message async def init_weights_send_group_for_remote_instance( self, @@ -453,7 +459,13 @@ class TokenizerCommunicatorMixin: # cannot run while requests are in progress. async with self.model_update_lock.writer_lock: result = (await self.update_weights_from_tensor_communicator(obj))[0] - return result.success, result.message + success, message = result.success, result.message + + if success and obj.weight_version is not None: + self._update_weight_version_if_provided(obj.weight_version) + message += f" Weight version updated to {obj.weight_version}." + + return success, message async def update_weights_from_ipc( self, @@ -471,11 +483,17 @@ class TokenizerCommunicatorMixin: # This means that weight sync cannot run while requests are in progress. async with self.model_update_lock.writer_lock: result = (await self.update_weights_from_ipc_communicator(obj))[0] - return result.success, result.message + success, message = result.success, result.message except Exception as e: error_msg = f"IPC weight update failed: {str(e)}" logger.error(error_msg) - return False, error_msg + success, message = False, error_msg + + if success and obj.weight_version is not None: + self._update_weight_version_if_provided(obj.weight_version) + message += f" Weight version updated to {obj.weight_version}." + + return success, message async def load_lora_adapter( self: TokenizerManager, @@ -693,3 +711,8 @@ class TokenizerCommunicatorMixin: f"Invalid --log-requests-level: {self.log_requests_level=}" ) return max_length, skip_names, out_skip_names + + def _update_weight_version_if_provided(self, weight_version: Optional[str]) -> None: + """Update weight version if provided.""" + if weight_version is not None: + self.server_args.weight_version = weight_version diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 4dbf078c8..3e51939e0 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -1216,7 +1216,15 @@ class TokenizerManager(TokenizerCommunicatorMixin): # Hold the lock if it is not async. This means that weight sync # cannot run while requests are in progress. async with self.model_update_lock.writer_lock: - return await self._wait_for_model_update_from_disk(obj) + success, message, num_paused_requests = ( + await self._wait_for_model_update_from_disk(obj) + ) + + if success and obj.weight_version is not None: + self._update_weight_version_if_provided(obj.weight_version) + message += f" Weight version updated to {obj.weight_version}." + + return success, message, num_paused_requests async def _wait_for_model_update_from_disk( self, obj: UpdateWeightFromDiskReqInput