[LoRA] 1/n Per-rank tensor serialization for load_lora_adapter_from_tensors under dp_size > 1 (#32580)

This commit is contained in:
Ethan (Yusheng) Su
2026-07-28 14:28:56 -07:00
committed by GitHub
parent d943636a48
commit 0a49226d19
6 changed files with 57 additions and 36 deletions
+26 -17
View File
@@ -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(
+4 -1
View File
@@ -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
@@ -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,
+17 -14
View File
@@ -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(),