[diffusion] fix: update weight from tensor detects device by uuid (#32685)
This commit is contained in:
@@ -22,6 +22,8 @@ class UpdateWeightFromTensorReqInput:
|
|||||||
"""Request to update model weights from tensor payloads for diffusion models."""
|
"""Request to update model weights from tensor payloads for diffusion models."""
|
||||||
|
|
||||||
serialized_named_tensors: list[str | bytes]
|
serialized_named_tensors: list[str | bytes]
|
||||||
|
# Physical GPU UUID each payload was exported from, one per payload.
|
||||||
|
payload_gpu_uuids: list[str] | None = None
|
||||||
load_format: str | None = None
|
load_format: str | None = None
|
||||||
target_modules: list[str] | None = None
|
target_modules: list[str] | None = None
|
||||||
weight_update_mode: str | None = None
|
weight_update_mode: str | None = None
|
||||||
|
|||||||
@@ -70,6 +70,7 @@ async def update_weights_from_tensor(request: Request):
|
|||||||
|
|
||||||
req = UpdateWeightFromTensorReqInput(
|
req = UpdateWeightFromTensorReqInput(
|
||||||
serialized_named_tensors=serialized_named_tensors,
|
serialized_named_tensors=serialized_named_tensors,
|
||||||
|
payload_gpu_uuids=body.get("payload_gpu_uuids"),
|
||||||
load_format=body.get("load_format"),
|
load_format=body.get("load_format"),
|
||||||
target_modules=body.get("target_modules"),
|
target_modules=body.get("target_modules"),
|
||||||
weight_update_mode=body.get("weight_update_mode"),
|
weight_update_mode=body.get("weight_update_mode"),
|
||||||
|
|||||||
+43
-15
@@ -2,7 +2,11 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.distributed import get_tp_rank, get_tp_world_size
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
|
get_tp_rank,
|
||||||
|
get_tp_world_size,
|
||||||
|
get_world_group,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.weight_utils import compute_weights_checksum
|
from sglang.multimodal_gen.runtime.loader.weight_utils import compute_weights_checksum
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
||||||
iter_materialized_weights,
|
iter_materialized_weights,
|
||||||
@@ -14,6 +18,7 @@ from sglang.multimodal_gen.runtime.post_training.weights_updater import (
|
|||||||
WeightsUpdater,
|
WeightsUpdater,
|
||||||
get_updatable_modules,
|
get_updatable_modules,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.utils import MultiprocessingSerializer
|
from sglang.srt.utils import MultiprocessingSerializer
|
||||||
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
|
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
|
||||||
|
|
||||||
@@ -24,6 +29,11 @@ if TYPE_CHECKING:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_gpu_uuid(uuid: str) -> str:
|
||||||
|
# NVML prefixes uuids with "GPU-"/"MIG-"; torch device properties do not.
|
||||||
|
return uuid.removeprefix("MIG-").removeprefix("GPU-").lower()
|
||||||
|
|
||||||
|
|
||||||
class GPUWorkerPostTrainingMixin:
|
class GPUWorkerPostTrainingMixin:
|
||||||
def update_weights_from_disk(
|
def update_weights_from_disk(
|
||||||
self,
|
self,
|
||||||
@@ -52,9 +62,9 @@ class GPUWorkerPostTrainingMixin:
|
|||||||
if not self.pipeline:
|
if not self.pipeline:
|
||||||
return False, "Pipeline is not initialized"
|
return False, "Pipeline is not initialized"
|
||||||
|
|
||||||
payload, error = self._select_rank_scoped_payload(
|
payload, error = self._select_own_gpu_payload(
|
||||||
payloads=req.serialized_named_tensors,
|
payloads=req.serialized_named_tensors,
|
||||||
field_name="serialized_named_tensors",
|
payload_gpu_uuids=req.payload_gpu_uuids,
|
||||||
)
|
)
|
||||||
if error is not None:
|
if error is not None:
|
||||||
return False, error
|
return False, error
|
||||||
@@ -153,23 +163,41 @@ class GPUWorkerPostTrainingMixin:
|
|||||||
)
|
)
|
||||||
return checksums
|
return checksums
|
||||||
|
|
||||||
def _select_rank_scoped_payload(
|
def _select_own_gpu_payload(
|
||||||
self,
|
self,
|
||||||
payloads: list,
|
payloads: list,
|
||||||
field_name: str,
|
payload_gpu_uuids: list[str] | None,
|
||||||
) -> tuple[object | None, str | None]:
|
) -> tuple[object | None, str | None]:
|
||||||
if not isinstance(payloads, list):
|
if not isinstance(payloads, list):
|
||||||
return None, f"{field_name} must be a list"
|
return None, "serialized_named_tensors must be a list"
|
||||||
if not payloads:
|
if not payloads:
|
||||||
return None, f"{field_name} is required"
|
return None, "serialized_named_tensors is required"
|
||||||
|
|
||||||
tp_world_size = get_tp_world_size()
|
if payload_gpu_uuids is None:
|
||||||
if len(payloads) not in (1, tp_world_size):
|
# Unlabeled fallback: world rank equals tp rank for tp-only runs.
|
||||||
return (
|
if len(payloads) == 1:
|
||||||
None,
|
return payloads[0], None
|
||||||
f"{field_name} size must be 1 or tp_size ({tp_world_size}), "
|
world_group = get_world_group()
|
||||||
f"got {len(payloads)}",
|
if len(payloads) != world_group.world_size:
|
||||||
|
return None, (
|
||||||
|
f"serialized_named_tensors size must be 1 or world_size "
|
||||||
|
f"({world_group.world_size}), got {len(payloads)}"
|
||||||
|
)
|
||||||
|
return payloads[world_group.rank_in_group], None
|
||||||
|
|
||||||
|
if len(payload_gpu_uuids) != len(payloads):
|
||||||
|
return None, (
|
||||||
|
f"payload_gpu_uuids needs one entry per payload, "
|
||||||
|
f"got {len(payload_gpu_uuids)} for {len(payloads)} payloads"
|
||||||
)
|
)
|
||||||
|
|
||||||
payload_idx = get_tp_rank() if len(payloads) == tp_world_size else 0
|
own_gpu_uuid = _normalize_gpu_uuid(
|
||||||
return payloads[payload_idx], None
|
current_platform.get_device_uuid(self.local_rank)
|
||||||
|
)
|
||||||
|
normalized_uuids = [_normalize_gpu_uuid(uuid) for uuid in payload_gpu_uuids]
|
||||||
|
if own_gpu_uuid not in normalized_uuids:
|
||||||
|
return None, (
|
||||||
|
f"no payload was exported from this worker's GPU {own_gpu_uuid}, "
|
||||||
|
f"got {payload_gpu_uuids}"
|
||||||
|
)
|
||||||
|
return payloads[normalized_uuids.index(own_gpu_uuid)], None
|
||||||
|
|||||||
Reference in New Issue
Block a user