[diffusion] fix: update weight from tensor detects device by uuid (#32685)

This commit is contained in:
Kangrui Du
2026-08-09 16:11:28 +08:00
committed by GitHub
parent bc285b2064
commit d0aa37b49b
3 changed files with 46 additions and 15 deletions
@@ -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"),
@@ -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