From 5155016b568f8ef123e3ef7a7b64371a56a04885 Mon Sep 17 00:00:00 2001 From: ShawnY112358 <61113840+ShawnY112358@users.noreply.github.com> Date: Thu, 27 Nov 2025 07:30:45 +0800 Subject: [PATCH] [feat] update bucketed weights from distributed (#13824) Co-authored-by: Stefan He --- python/sglang/srt/entrypoints/engine.py | 2 + python/sglang/srt/managers/io_struct.py | 2 + python/sglang/srt/managers/tp_worker.py | 6 +- .../sglang/srt/model_executor/model_runner.py | 44 +++++++++- .../test_update_weights_from_distributed.py | 85 ++++++++++++++++--- 5 files changed, 124 insertions(+), 15 deletions(-) diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 51a809222..7ca884e41 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -466,6 +466,7 @@ class Engine(EngineBase): shapes: list[list[int]], group_name: str = "weight_update_group", flush_cache: bool = True, + load_format: Optional[str] = None, ): """Update weights from distributed source.""" obj = UpdateWeightsFromDistributedReqInput( @@ -474,6 +475,7 @@ class Engine(EngineBase): shapes=shapes, group_name=group_name, flush_cache=flush_cache, + load_format=load_format, ) return self.loop.run_until_complete( self.tokenizer_manager.update_weights_from_distributed(obj, None) diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index f4acf6907..46647d01f 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -1144,6 +1144,8 @@ class UpdateWeightsFromDistributedReqInput(BaseReq): abort_all_requests: bool = False # Optional: Update weight version along with weights weight_version: Optional[str] = None + # Optional format specification for loading + load_format: Optional[str] = None @dataclass diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 5e8bd9241..c5f20b5f9 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -153,7 +153,11 @@ class BaseTpWorker(ABC): self, recv_req: UpdateWeightsFromDistributedReqInput ): success, message = self.model_runner.update_weights_from_distributed( - recv_req.names, recv_req.dtypes, recv_req.shapes, recv_req.group_name + recv_req.names, + recv_req.dtypes, + recv_req.shapes, + recv_req.group_name, + recv_req.load_format, ) return success, message diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index eeb4beade..90112ef1c 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -1153,7 +1153,14 @@ class ModelRunner: logger.error(message) return False, message - def update_weights_from_distributed(self, names, dtypes, shapes, group_name): + def update_weights_from_distributed( + self, + names, + dtypes, + shapes, + group_name, + load_format: Optional[str] = None, + ): """ Update specific parameter in the model weights online through `_model_update_group` process group. @@ -1169,6 +1176,10 @@ class ModelRunner: "Please call `init_weights_update_group` first." ) + if load_format == "flattened_bucket": + return self._update_bucketed_weights_from_distributed( + names, dtypes, shapes, group_name + ) try: weights = [] handles = [] @@ -1201,6 +1212,37 @@ class ModelRunner: logger.error(error_msg) return False, error_msg + def _update_bucketed_weights_from_distributed( + self, names, dtypes, shapes, group_name + ): + try: + named_tensors = [] + for name, dtype, shape in zip(names, dtypes, shapes): + target_dtype = ( + dtype if isinstance(dtype, torch.dtype) else getattr(torch, dtype) + ) + named_tensors.append( + (name, torch.empty(shape, dtype=target_dtype, device=self.device)) + ) + bucket = FlattenedTensorBucket(named_tensors=named_tensors) + flattened_tensor = bucket.get_flattened_tensor() + torch.distributed.broadcast( + flattened_tensor, + src=0, + group=self._model_update_group[group_name], + ) + reconstructed_tensors = bucket.reconstruct_tensors() + self.model.load_weights(reconstructed_tensors) + return True, f"Succeeded to update parameter online." + except Exception as e: + error_msg = ( + f"Failed to update parameter online: {e}. " + f"The full weights of the ModelRunner are partially updated. " + f"Please discard the whole weights." + ) + logger.error(error_msg) + return False, error_msg + def update_weights_from_tensor( self, named_tensors: List[Tuple[str, Union[torch.Tensor, "LocalSerializedTensor"]]], diff --git a/test/srt/rl/test_update_weights_from_distributed.py b/test/srt/rl/test_update_weights_from_distributed.py index fd45a4e5c..e48342660 100644 --- a/test/srt/rl/test_update_weights_from_distributed.py +++ b/test/srt/rl/test_update_weights_from_distributed.py @@ -28,6 +28,7 @@ from transformers import AutoModelForCausalLM import sglang as sgl from sglang.srt.utils import init_custom_process_group +from sglang.srt.weight_sync.tensor_bucket import FlattenedTensorBucket from sglang.test.test_utils import ( DEFAULT_MODEL_NAME_FOR_TEST, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, @@ -69,6 +70,7 @@ def init_process( backend, checking_parameters, tie_word_embeddings, + load_format, barrier, pause_generation_mode, ): @@ -84,6 +86,7 @@ def init_process( checking_parameters, tie_word_embeddings, state_dict_key_to_shape, + load_format, barrier, ) elif rank in [1, 2]: @@ -98,6 +101,7 @@ def init_process( state_dict_key_to_shape, backend, tp_size, + load_format, barrier, pause_generation_mode, ) @@ -112,6 +116,7 @@ def init_process_hf( checking_parameters, tie_word_embeddings, state_dict_key_to_shape, + load_format, barrier, ): # These two environment variables are very important @@ -179,14 +184,23 @@ def init_process_hf( if tie_word_embeddings: broadcast_parameters.remove("lm_head.weight") - # Broadcast all the weights from the training - # engine to other ranks (inference engine). - for parameter_name in broadcast_parameters: - torch.distributed.broadcast( - hf_base_model.get_parameter(parameter_name), - src=0, - group=group, - ) + if load_format == "flattened_bucket": + named_tensors = [ + (parameter_name, hf_base_model.get_parameter(parameter_name)) + for parameter_name in broadcast_parameters + ] + bucket = FlattenedTensorBucket(named_tensors=named_tensors) + flattened_tensor = bucket.get_flattened_tensor() + torch.distributed.broadcast(flattened_tensor, src=0, group=group) + else: + # Broadcast all the weights from the training + # engine to other ranks (inference engine). + for parameter_name in broadcast_parameters: + torch.distributed.broadcast( + hf_base_model.get_parameter(parameter_name), + src=0, + group=group, + ) torch.cuda.synchronize() time_end_broadcast = time.perf_counter() @@ -216,6 +230,7 @@ def init_process_sgl( state_dict_key_to_shape, backend, tp_size, + load_format, barrier, pause_generation_mode, ): @@ -341,6 +356,7 @@ def init_process_sgl( dtypes=dtypes, shapes=shapes, group_name="test_parameter_update_group", + load_format=load_format, ) else: requests.post( @@ -350,6 +366,7 @@ def init_process_sgl( "dtypes": dtypes, "shapes": shapes, "group_name": "test_parameter_update_group", + "load_format": load_format, "flush_cache": not (pause_generation_mode == "in_place"), }, ) @@ -431,6 +448,7 @@ def test_update_weights_from_distributed( state_dict_key_to_shape, truncate_size, checking_parameters, + load_format=None, pause_generation_mode=None, ): tie_word_embeddings = ( @@ -456,6 +474,7 @@ def test_update_weights_from_distributed( backend, checking_parameters, tie_word_embeddings, + load_format, barrier, pause_generation_mode, ), @@ -614,31 +633,55 @@ class TestUpdateWeightsFromDistributed(CustomTestCase): pause_generation_mode = random.choice(["in_place", "retract"]) else: pause_generation_mode = None + load_format = random.choice(["flattened_bucket", None]) test_suits = [ - (1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, mode, pause_generation_mode), + ( + 1, + 1, + DEFAULT_SMALL_MODEL_NAME_FOR_TEST, + mode, + pause_generation_mode, + load_format, + ), ] else: test_suits = [ - (1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, "Engine", None), + ( + 1, + 1, + DEFAULT_SMALL_MODEL_NAME_FOR_TEST, + "Engine", + None, + random.choice(["flattened_bucket", None]), + ), ( 1, 1, DEFAULT_MODEL_NAME_FOR_TEST, "Sever", random.choice(["in_place", "retract"]), + random.choice(["flattened_bucket", None]), ), ] if torch.cuda.device_count() >= 4: test_suits.extend( [ - (2, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, "Engine", None), + ( + 2, + 1, + DEFAULT_SMALL_MODEL_NAME_FOR_TEST, + "Engine", + None, + random.choice(["flattened_bucket", None]), + ), ( 1, 2, DEFAULT_MODEL_NAME_FOR_TEST, "Server", random.choice(["in_place", "retract"]), + random.choice(["flattened_bucket", None]), ), ] ) @@ -646,13 +689,21 @@ class TestUpdateWeightsFromDistributed(CustomTestCase): if torch.cuda.device_count() >= 5: test_suits.extend( [ - (2, 2, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, "Engine", None), + ( + 2, + 2, + DEFAULT_SMALL_MODEL_NAME_FOR_TEST, + "Engine", + None, + random.choice(["flattened_bucket", None]), + ), ( 2, 2, DEFAULT_MODEL_NAME_FOR_TEST, "Server", random.choice(["in_place", "retract"]), + random.choice(["flattened_bucket", None]), ), ] ) @@ -689,7 +740,14 @@ class TestUpdateWeightsFromDistributed(CustomTestCase): "lm_head.weight", ] - for tp_size, dp_size, model_name, backend, pause_generation_mode in test_suits: + for ( + tp_size, + dp_size, + model_name, + backend, + pause_generation_mode, + load_format, + ) in test_suits: test_update_weights_from_distributed( tp_size, dp_size, @@ -698,6 +756,7 @@ class TestUpdateWeightsFromDistributed(CustomTestCase): model_state_dict_shapes[model_name], truncate_size, checking_parameters, + load_format, pause_generation_mode, )