diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 25a80696c..033da16d0 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -1308,15 +1308,9 @@ class Engine(EngineScoreMixin, EngineBase): ): """Update weights from distributed source. If there are going to be more updates, set `flush_cache` to be false to avoid duplicated cache cleaning operation.""" - if load_format == "flattened_bucket": - serialized_named_tensors = normalize_serialized_named_tensor_payloads( - cast(List[SerializedTensorPayload], named_tensors) - ) - else: - serialized_named_tensors = [ - MultiprocessingSerializer.serialize(named_tensors) - for _ in range(self.server_args.tp_size) - ] + serialized_named_tensors = self._serialize_tensors_per_rank( + named_tensors, load_format + ) obj = UpdateWeightsFromTensorReqInput( serialized_named_tensors=serialized_named_tensors, load_format=load_format, @@ -1367,23 +1361,38 @@ class Engine(EngineScoreMixin, EngineBase): self.tokenizer_manager.get_weights_by_name(obj, None) ) + def _serialize_tensors_per_rank( + self, + tensors, + load_format: Optional[str], + ) -> List[bytes]: + """One serialized payload per TP rank: each rank deserializes only its + own copy, so producer-side CUDA-IPC refcounts drop cleanly after every + load. flattened_bucket callers pass pre-serialized per-rank payloads.""" + if load_format == "flattened_bucket": + return normalize_serialized_named_tensor_payloads( + cast(List[SerializedTensorPayload], tensors) + ) + else: + return [ + MultiprocessingSerializer.serialize(tensors) + for _ in range(self.server_args.tp_size) + ] + def load_lora_adapter_from_tensors( self, lora_name: str, - tensors, + tensors: Union[Dict[str, torch.Tensor], List[SerializedTensorPayload]], config_dict: Dict, load_format: Optional[str] = None, ): - if load_format == "flattened_bucket": - serialized_tensors = tensors - else: - serialized_tensors = MultiprocessingSerializer.serialize( - tensors, output_str=True - ) + serialized_named_tensors = self._serialize_tensors_per_rank( + tensors, load_format + ) lora_req = LoadLoRAAdapterFromTensorsReqInput( lora_name=lora_name, config_dict=config_dict, - serialized_tensors=serialized_tensors, + serialized_named_tensors=serialized_named_tensors, load_format=load_format, ) return self.loop.run_until_complete( diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 3ce179755..ef1617c8e 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -2049,7 +2049,10 @@ class LoadLoRAAdapterFromTensorsReqInput(BaseReq, kw_only=True): # The PEFT adapter_config.json, already JSON — a tighter type would only add # decode strictness with no benefit. config_dict: Dict[str, Any] - serialized_tensors: str + # One serialized copy of the adapter tensors per TP rank; each rank + # deserializes only its own copy. Same normalization conventions as + # UpdateWeightsFromTensorReqInput.serialized_named_tensors. + serialized_named_tensors: Annotated[List[bytes], Base64Bytes()] pinned: bool = False added_tokens_config: Optional[Dict[str, int]] = None lora_id: Optional[str] = None diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index 86b7b378f..b4fa41ef6 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -653,13 +653,17 @@ class TokenizerControlMixin: ) assert ( - self.server_args.dp_size == 1 - ), "dp_size must be 1 for dynamic lora loading" + self.server_args.dp_size == 1 or self.server_args.enable_dp_attention + ), "dp_size must be 1 or dp attention must be enabled for dynamic lora loading" logger.info( "Start load Lora adapter from tensors. Lora name=%s", obj.lora_name, ) + obj.serialized_named_tensors = normalize_serialized_named_tensor_payloads( + obj.serialized_named_tensors + ) + async with self.lora_update_lock: new_adapter = LoRARef( lora_name=obj.lora_name, diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 1ad3711fc..19ff16881 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -173,13 +173,18 @@ class BaseTpWorker(ABC): ) return success, message - def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput): - + def _deserialize_own_rank(self, serialized_named_tensors): + """Each rank deserializes only its own payload (index ps.tp_rank); + deserializing another rank's copy would break producer-side CUDA-IPC + refcounting.""" monkey_patch_torch_reductions() + return MultiprocessingSerializer.deserialize( + serialized_named_tensors[self.ps.tp_rank] + ) + + def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput): success, message = self.model_runner.weight_updater.update_weights_from_tensor( - named_tensors=MultiprocessingSerializer.deserialize( - recv_req.serialized_named_tensors[self.ps.tp_rank] - ), + named_tensors=self._deserialize_own_rank(recv_req.serialized_named_tensors), load_format=recv_req.load_format, ) return success, message @@ -209,18 +214,16 @@ class BaseTpWorker(ABC): self, recv_req: LoadLoRAAdapterFromTensorsReqInput ): # The LoRA code handles TP sharding internally using slice_lora_a_weights - # and slice_lora_b_weights methods (see lora/layers.py:46-49, mem_pool.py:437-440). + # and slice_lora_b_weights methods (see lora/layers.py and mem_pool.py). + data = self._deserialize_own_rank(recv_req.serialized_named_tensors) if recv_req.load_format == "flattened_bucket": - flattened_data = MultiprocessingSerializer.deserialize( - recv_req.serialized_tensors - ) bucket = FlattenedTensorBucket( - flattened_tensor=flattened_data["flattened_tensor"], - metadata=flattened_data["metadata"], + flattened_tensor=data["flattened_tensor"], + metadata=data["metadata"], ) tensors = dict(bucket.reconstruct_tensors()) else: - tensors = MultiprocessingSerializer.deserialize(recv_req.serialized_tensors) + tensors = data if recv_req.expected_checksums is not None: import hashlib @@ -245,12 +248,12 @@ class BaseTpWorker(ABC): extra = [n for n in tensors if n not in exp] if mismatch or missing or extra: raise RuntimeError( - f"[LORA-CHECK] rank{self.tp_rank} adapter sync MISMATCH of {len(exp)} expected: " + f"[LORA-CHECK] rank{self.ps.tp_rank} adapter sync MISMATCH of {len(exp)} expected: " f"{len(mismatch)} value-diff {mismatch[:5]}, {len(missing)} missing {missing[:5]}, " f"{len(extra)} extra {extra[:5]}" ) logger.info( - f"[LORA-CHECK] rank{self.tp_rank} adapter sync OK: {len(exp)}/{len(exp)} tensors match (sha256)" + f"[LORA-CHECK] rank{self.ps.tp_rank} adapter sync OK: {len(exp)}/{len(exp)} tensors match (sha256)" ) result = self.model_runner.load_lora_adapter_from_tensors( recv_req.to_ref(), diff --git a/test/registered/rl/test_lora_load_from_tensor.py b/test/registered/rl/test_lora_load_from_tensor.py index d2505215a..c37aab68b 100644 --- a/test/registered/rl/test_lora_load_from_tensor.py +++ b/test/registered/rl/test_lora_load_from_tensor.py @@ -342,9 +342,11 @@ class TestLoRALoadFromTensor(CustomTestCase): } serialized = MultiprocessingSerializer.serialize(bucket_dict, output_str=True) + # flattened_bucket callers pass one serialized copy per TP rank, same + # as Engine.update_weights_from_tensor. result = self.engine.load_lora_adapter_from_tensors( lora_name="self_cognition_Alice_flattened", - tensors=serialized, + tensors=[serialized], config_dict=self.lora_config_dict, load_format="flattened_bucket", ) diff --git a/test/registered/unit/managers/test_msgpack_ipc_roundtrip.py b/test/registered/unit/managers/test_msgpack_ipc_roundtrip.py index e3d3e95d0..0dabbe330 100644 --- a/test/registered/unit/managers/test_msgpack_ipc_roundtrip.py +++ b/test/registered/unit/managers/test_msgpack_ipc_roundtrip.py @@ -126,7 +126,7 @@ REGISTRY_TYPE_INSTANCES = { "LoadLoRAAdapterFromTensorsReqInput": LoadLoRAAdapterFromTensorsReqInput( lora_name="adapter", config_dict={"r": 8, "lora_alpha": 16, "target_modules": ["q_proj", "v_proj"]}, - serialized_tensors="", + serialized_named_tensors=[b"tp0-bytes", b"tp1-bytes"], added_tokens_config={"": 32000}, ), "DumperControlReqInput": DumperControlReqInput(method="start", body={"k": "v"}),