[LoRA] 1/n Per-rank tensor serialization for load_lora_adapter_from_tensors under dp_size > 1 (#32580)
This commit is contained in:
@@ -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
|
"""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."""
|
to avoid duplicated cache cleaning operation."""
|
||||||
if load_format == "flattened_bucket":
|
serialized_named_tensors = self._serialize_tensors_per_rank(
|
||||||
serialized_named_tensors = normalize_serialized_named_tensor_payloads(
|
named_tensors, load_format
|
||||||
cast(List[SerializedTensorPayload], named_tensors)
|
|
||||||
)
|
)
|
||||||
else:
|
|
||||||
serialized_named_tensors = [
|
|
||||||
MultiprocessingSerializer.serialize(named_tensors)
|
|
||||||
for _ in range(self.server_args.tp_size)
|
|
||||||
]
|
|
||||||
obj = UpdateWeightsFromTensorReqInput(
|
obj = UpdateWeightsFromTensorReqInput(
|
||||||
serialized_named_tensors=serialized_named_tensors,
|
serialized_named_tensors=serialized_named_tensors,
|
||||||
load_format=load_format,
|
load_format=load_format,
|
||||||
@@ -1367,23 +1361,38 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
self.tokenizer_manager.get_weights_by_name(obj, None)
|
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(
|
def load_lora_adapter_from_tensors(
|
||||||
self,
|
self,
|
||||||
lora_name: str,
|
lora_name: str,
|
||||||
tensors,
|
tensors: Union[Dict[str, torch.Tensor], List[SerializedTensorPayload]],
|
||||||
config_dict: Dict,
|
config_dict: Dict,
|
||||||
load_format: Optional[str] = None,
|
load_format: Optional[str] = None,
|
||||||
):
|
):
|
||||||
if load_format == "flattened_bucket":
|
serialized_named_tensors = self._serialize_tensors_per_rank(
|
||||||
serialized_tensors = tensors
|
tensors, load_format
|
||||||
else:
|
|
||||||
serialized_tensors = MultiprocessingSerializer.serialize(
|
|
||||||
tensors, output_str=True
|
|
||||||
)
|
)
|
||||||
lora_req = LoadLoRAAdapterFromTensorsReqInput(
|
lora_req = LoadLoRAAdapterFromTensorsReqInput(
|
||||||
lora_name=lora_name,
|
lora_name=lora_name,
|
||||||
config_dict=config_dict,
|
config_dict=config_dict,
|
||||||
serialized_tensors=serialized_tensors,
|
serialized_named_tensors=serialized_named_tensors,
|
||||||
load_format=load_format,
|
load_format=load_format,
|
||||||
)
|
)
|
||||||
return self.loop.run_until_complete(
|
return self.loop.run_until_complete(
|
||||||
|
|||||||
@@ -2049,7 +2049,10 @@ class LoadLoRAAdapterFromTensorsReqInput(BaseReq, kw_only=True):
|
|||||||
# The PEFT adapter_config.json, already JSON — a tighter type would only add
|
# The PEFT adapter_config.json, already JSON — a tighter type would only add
|
||||||
# decode strictness with no benefit.
|
# decode strictness with no benefit.
|
||||||
config_dict: Dict[str, Any]
|
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
|
pinned: bool = False
|
||||||
added_tokens_config: Optional[Dict[str, int]] = None
|
added_tokens_config: Optional[Dict[str, int]] = None
|
||||||
lora_id: Optional[str] = None
|
lora_id: Optional[str] = None
|
||||||
|
|||||||
@@ -653,13 +653,17 @@ class TokenizerControlMixin:
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert (
|
assert (
|
||||||
self.server_args.dp_size == 1
|
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
|
||||||
), "dp_size must be 1 for dynamic lora loading"
|
), "dp_size must be 1 or dp attention must be enabled for dynamic lora loading"
|
||||||
logger.info(
|
logger.info(
|
||||||
"Start load Lora adapter from tensors. Lora name=%s",
|
"Start load Lora adapter from tensors. Lora name=%s",
|
||||||
obj.lora_name,
|
obj.lora_name,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
obj.serialized_named_tensors = normalize_serialized_named_tensor_payloads(
|
||||||
|
obj.serialized_named_tensors
|
||||||
|
)
|
||||||
|
|
||||||
async with self.lora_update_lock:
|
async with self.lora_update_lock:
|
||||||
new_adapter = LoRARef(
|
new_adapter = LoRARef(
|
||||||
lora_name=obj.lora_name,
|
lora_name=obj.lora_name,
|
||||||
|
|||||||
@@ -173,13 +173,18 @@ class BaseTpWorker(ABC):
|
|||||||
)
|
)
|
||||||
return success, message
|
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()
|
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(
|
success, message = self.model_runner.weight_updater.update_weights_from_tensor(
|
||||||
named_tensors=MultiprocessingSerializer.deserialize(
|
named_tensors=self._deserialize_own_rank(recv_req.serialized_named_tensors),
|
||||||
recv_req.serialized_named_tensors[self.ps.tp_rank]
|
|
||||||
),
|
|
||||||
load_format=recv_req.load_format,
|
load_format=recv_req.load_format,
|
||||||
)
|
)
|
||||||
return success, message
|
return success, message
|
||||||
@@ -209,18 +214,16 @@ class BaseTpWorker(ABC):
|
|||||||
self, recv_req: LoadLoRAAdapterFromTensorsReqInput
|
self, recv_req: LoadLoRAAdapterFromTensorsReqInput
|
||||||
):
|
):
|
||||||
# The LoRA code handles TP sharding internally using slice_lora_a_weights
|
# 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":
|
if recv_req.load_format == "flattened_bucket":
|
||||||
flattened_data = MultiprocessingSerializer.deserialize(
|
|
||||||
recv_req.serialized_tensors
|
|
||||||
)
|
|
||||||
bucket = FlattenedTensorBucket(
|
bucket = FlattenedTensorBucket(
|
||||||
flattened_tensor=flattened_data["flattened_tensor"],
|
flattened_tensor=data["flattened_tensor"],
|
||||||
metadata=flattened_data["metadata"],
|
metadata=data["metadata"],
|
||||||
)
|
)
|
||||||
tensors = dict(bucket.reconstruct_tensors())
|
tensors = dict(bucket.reconstruct_tensors())
|
||||||
else:
|
else:
|
||||||
tensors = MultiprocessingSerializer.deserialize(recv_req.serialized_tensors)
|
tensors = data
|
||||||
if recv_req.expected_checksums is not None:
|
if recv_req.expected_checksums is not None:
|
||||||
import hashlib
|
import hashlib
|
||||||
|
|
||||||
@@ -245,12 +248,12 @@ class BaseTpWorker(ABC):
|
|||||||
extra = [n for n in tensors if n not in exp]
|
extra = [n for n in tensors if n not in exp]
|
||||||
if mismatch or missing or extra:
|
if mismatch or missing or extra:
|
||||||
raise RuntimeError(
|
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(mismatch)} value-diff {mismatch[:5]}, {len(missing)} missing {missing[:5]}, "
|
||||||
f"{len(extra)} extra {extra[:5]}"
|
f"{len(extra)} extra {extra[:5]}"
|
||||||
)
|
)
|
||||||
logger.info(
|
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(
|
result = self.model_runner.load_lora_adapter_from_tensors(
|
||||||
recv_req.to_ref(),
|
recv_req.to_ref(),
|
||||||
|
|||||||
@@ -342,9 +342,11 @@ class TestLoRALoadFromTensor(CustomTestCase):
|
|||||||
}
|
}
|
||||||
serialized = MultiprocessingSerializer.serialize(bucket_dict, output_str=True)
|
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(
|
result = self.engine.load_lora_adapter_from_tensors(
|
||||||
lora_name="self_cognition_Alice_flattened",
|
lora_name="self_cognition_Alice_flattened",
|
||||||
tensors=serialized,
|
tensors=[serialized],
|
||||||
config_dict=self.lora_config_dict,
|
config_dict=self.lora_config_dict,
|
||||||
load_format="flattened_bucket",
|
load_format="flattened_bucket",
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -126,7 +126,7 @@ REGISTRY_TYPE_INSTANCES = {
|
|||||||
"LoadLoRAAdapterFromTensorsReqInput": LoadLoRAAdapterFromTensorsReqInput(
|
"LoadLoRAAdapterFromTensorsReqInput": LoadLoRAAdapterFromTensorsReqInput(
|
||||||
lora_name="adapter",
|
lora_name="adapter",
|
||||||
config_dict={"r": 8, "lora_alpha": 16, "target_modules": ["q_proj", "v_proj"]},
|
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={"<extra>": 32000},
|
added_tokens_config={"<extra>": 32000},
|
||||||
),
|
),
|
||||||
"DumperControlReqInput": DumperControlReqInput(method="start", body={"k": "v"}),
|
"DumperControlReqInput": DumperControlReqInput(method="start", body={"k": "v"}),
|
||||||
|
|||||||
Reference in New Issue
Block a user