[diffusion] rl: extract post-training weight apis into mixins and add tensor update/checker paths (#22817)
Co-authored-by: dreamyang-liu <nikolaliu@icloud.com> Co-authored-by: Xiaole Guo <vera0315@connect.hku.hk> Co-authored-by: zhaochenyang20 <zhaochen20@outlook.com> Co-authored-by: Xiaole Guo <119391905+gxlvera@users.noreply.github.com> Co-authored-by: zhihengy <zhihengy@andrew.cmu.edu> Co-authored-by: MikukuOvO <MikukuOvO@users.noreply.github.com>
This commit is contained in:
co-authored by
dreamyang-liu
Xiaole Guo
zhaochenyang20
Xiaole Guo
zhihengy
MikukuOvO
parent
cae022aa88
commit
5c0b2859e8
@@ -17,6 +17,23 @@ class UpdateWeightFromDiskReqInput:
|
|||||||
target_modules: list[str] | None = None
|
target_modules: list[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class UpdateWeightFromTensorReqInput:
|
||||||
|
"""Request to update model weights from tensor payloads for diffusion models."""
|
||||||
|
|
||||||
|
serialized_named_tensors: list[str | bytes]
|
||||||
|
load_format: str | None = None
|
||||||
|
target_modules: list[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class UpdateWeightFromTensorCheckerReqInput:
|
||||||
|
"""Request to verify live module weights against expected SHA-256 values."""
|
||||||
|
|
||||||
|
target_module: str
|
||||||
|
expected_named_tensors_sha256: dict[str, str]
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class GetWeightsChecksumReqInput:
|
class GetWeightsChecksumReqInput:
|
||||||
"""Compute SHA-256 checksum of loaded module weights for verification."""
|
"""Compute SHA-256 checksum of loaded module weights for verification."""
|
||||||
|
|||||||
@@ -5,6 +5,8 @@ from fastapi import APIRouter, Request
|
|||||||
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
|
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
|
||||||
GetWeightsChecksumReqInput,
|
GetWeightsChecksumReqInput,
|
||||||
UpdateWeightFromDiskReqInput,
|
UpdateWeightFromDiskReqInput,
|
||||||
|
UpdateWeightFromTensorCheckerReqInput,
|
||||||
|
UpdateWeightFromTensorReqInput,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
|
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
|
||||||
from sglang.srt.utils.json_response import orjson_response
|
from sglang.srt.utils.json_response import orjson_response
|
||||||
@@ -29,6 +31,82 @@ async def update_weights_from_disk(request: Request):
|
|||||||
target_modules=body.get("target_modules"),
|
target_modules=body.get("target_modules"),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
response = await async_scheduler_client.forward(req)
|
||||||
|
except Exception as e:
|
||||||
|
return orjson_response(
|
||||||
|
{"success": False, "message": str(e)},
|
||||||
|
status_code=500,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = response.output
|
||||||
|
return orjson_response(
|
||||||
|
result,
|
||||||
|
status_code=200 if result["success"] else 400,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/update_weights_from_tensor")
|
||||||
|
async def update_weights_from_tensor(request: Request):
|
||||||
|
"""Update model weights from serialized tensor payloads."""
|
||||||
|
body = await request.json()
|
||||||
|
serialized_named_tensors = body.get("serialized_named_tensors")
|
||||||
|
if not serialized_named_tensors:
|
||||||
|
return orjson_response(
|
||||||
|
{"success": False, "message": "serialized_named_tensors is required"},
|
||||||
|
status_code=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
req = UpdateWeightFromTensorReqInput(
|
||||||
|
serialized_named_tensors=serialized_named_tensors,
|
||||||
|
load_format=body.get("load_format"),
|
||||||
|
target_modules=body.get("target_modules"),
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
response = await async_scheduler_client.forward(req)
|
||||||
|
except Exception as e:
|
||||||
|
return orjson_response(
|
||||||
|
{"success": False, "message": str(e)},
|
||||||
|
status_code=500,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = response.output
|
||||||
|
return orjson_response(
|
||||||
|
result,
|
||||||
|
status_code=200 if result["success"] else 400,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/update_weights_from_tensor_checker")
|
||||||
|
async def update_weights_from_tensor_checker(request: Request):
|
||||||
|
"""Verify live module weights against expected SHA-256 values."""
|
||||||
|
body = await request.json()
|
||||||
|
target_module = body.get("target_module")
|
||||||
|
if not target_module:
|
||||||
|
return orjson_response(
|
||||||
|
{"success": False, "message": "target_module is required"},
|
||||||
|
status_code=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
expected_named_tensors_sha256 = body.get("expected_named_tensors_sha256")
|
||||||
|
if (
|
||||||
|
not isinstance(expected_named_tensors_sha256, dict)
|
||||||
|
or not expected_named_tensors_sha256
|
||||||
|
):
|
||||||
|
return orjson_response(
|
||||||
|
{
|
||||||
|
"success": False,
|
||||||
|
"message": "expected_named_tensors_sha256 is required",
|
||||||
|
},
|
||||||
|
status_code=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
req = UpdateWeightFromTensorCheckerReqInput(
|
||||||
|
target_module=target_module,
|
||||||
|
expected_named_tensors_sha256=expected_named_tensors_sha256,
|
||||||
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
response = await async_scheduler_client.forward(req)
|
response = await async_scheduler_client.forward(req)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -37,14 +37,8 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
|||||||
post_process_sample,
|
post_process_sample,
|
||||||
save_outputs,
|
save_outputs,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.weight_utils import compute_weights_checksum
|
|
||||||
from sglang.multimodal_gen.runtime.loader.weights_updater import (
|
|
||||||
WeightsUpdater,
|
|
||||||
get_updatable_modules,
|
|
||||||
)
|
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
||||||
configure_layerwise_offload_modules,
|
configure_layerwise_offload_modules,
|
||||||
iter_materialized_weights,
|
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core import (
|
from sglang.multimodal_gen.runtime.pipelines_core import (
|
||||||
ComposedPipelineBase,
|
ComposedPipelineBase,
|
||||||
@@ -54,6 +48,9 @@ from sglang.multimodal_gen.runtime.pipelines_core import (
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
|
||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
|
from sglang.multimodal_gen.runtime.post_training.gpu_worker_post_training_mixin import (
|
||||||
|
GPUWorkerPostTrainingMixin,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.realtime.session import (
|
from sglang.multimodal_gen.runtime.realtime.session import (
|
||||||
RealtimeSessionCache,
|
RealtimeSessionCache,
|
||||||
)
|
)
|
||||||
@@ -102,7 +99,7 @@ class _ExpandedOutputParts:
|
|||||||
trajectory_decoded_parts: list[list[torch.Tensor]] | None = None
|
trajectory_decoded_parts: list[list[torch.Tensor]] | None = None
|
||||||
|
|
||||||
|
|
||||||
class GPUWorker:
|
class GPUWorker(GPUWorkerPostTrainingMixin):
|
||||||
"""
|
"""
|
||||||
A worker that executes the model on a single GPU.
|
A worker that executes the model on a single GPU.
|
||||||
"""
|
"""
|
||||||
@@ -894,48 +891,6 @@ class GPUWorker:
|
|||||||
status = self.pipeline.get_lora_status()
|
status = self.pipeline.get_lora_status()
|
||||||
return OutputBatch(output=status)
|
return OutputBatch(output=status)
|
||||||
|
|
||||||
def update_weights_from_disk(
|
|
||||||
self,
|
|
||||||
model_path: str,
|
|
||||||
flush_cache: bool = True,
|
|
||||||
target_modules: list[str] | None = None,
|
|
||||||
) -> tuple[bool, str]:
|
|
||||||
"""Update model weights from disk inplace without restarting the server."""
|
|
||||||
if not self.pipeline:
|
|
||||||
return False, "Pipeline is not initialized"
|
|
||||||
|
|
||||||
updater = WeightsUpdater(self.pipeline)
|
|
||||||
success, message = updater.update_weights_from_disk(
|
|
||||||
model_path,
|
|
||||||
flush_cache=flush_cache,
|
|
||||||
target_modules=target_modules,
|
|
||||||
)
|
|
||||||
if success:
|
|
||||||
self.server_args.model_path = model_path
|
|
||||||
self.pipeline.model_path = model_path
|
|
||||||
return success, message
|
|
||||||
|
|
||||||
def get_weights_checksum(
|
|
||||||
self, module_names: list[str] | None = None
|
|
||||||
) -> dict[str, str]:
|
|
||||||
"""Compute SHA-256 checksum of each module's weights."""
|
|
||||||
if not self.pipeline:
|
|
||||||
return {"error": "Pipeline is not initialized"}
|
|
||||||
|
|
||||||
all_modules = get_updatable_modules(self.pipeline)
|
|
||||||
names = module_names if module_names is not None else list(all_modules.keys())
|
|
||||||
|
|
||||||
checksums: dict[str, str] = {}
|
|
||||||
for name in names:
|
|
||||||
module = all_modules.get(name)
|
|
||||||
if module is None:
|
|
||||||
checksums[name] = "not_found"
|
|
||||||
continue
|
|
||||||
checksums[name] = compute_weights_checksum(
|
|
||||||
iter_materialized_weights(module)
|
|
||||||
)
|
|
||||||
return checksums
|
|
||||||
|
|
||||||
|
|
||||||
OOM_MSG = """
|
OOM_MSG = """
|
||||||
OOM detected. Possible solutions:
|
OOM detected. Possible solutions:
|
||||||
|
|||||||
@@ -20,6 +20,8 @@ from sglang.multimodal_gen.runtime.distributed import get_world_group
|
|||||||
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
|
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
|
||||||
GetWeightsChecksumReqInput,
|
GetWeightsChecksumReqInput,
|
||||||
UpdateWeightFromDiskReqInput,
|
UpdateWeightFromDiskReqInput,
|
||||||
|
UpdateWeightFromTensorCheckerReqInput,
|
||||||
|
UpdateWeightFromTensorReqInput,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
||||||
GetDisaggStatsReq,
|
GetDisaggStatsReq,
|
||||||
@@ -44,6 +46,9 @@ from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import (
|
|||||||
BatchMetricsWindow,
|
BatchMetricsWindow,
|
||||||
OutputBatch,
|
OutputBatch,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.post_training.scheduler_post_training_mixin import (
|
||||||
|
SchedulerPostTrainingMixin,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.server_args import (
|
from sglang.multimodal_gen.runtime.server_args import (
|
||||||
PortArgs,
|
PortArgs,
|
||||||
ServerArgs,
|
ServerArgs,
|
||||||
@@ -69,7 +74,7 @@ _MAX_RECV_REQS_PER_POLL = 1024
|
|||||||
_BATCH_METRICS_LOG_INTERVAL = 5
|
_BATCH_METRICS_LOG_INTERVAL = 5
|
||||||
|
|
||||||
|
|
||||||
class Scheduler(SchedulerDisaggMixin):
|
class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin):
|
||||||
"""
|
"""
|
||||||
Runs the main event loop for the rank 0 worker.
|
Runs the main event loop for the rank 0 worker.
|
||||||
It listens for external requests via ZMQ and coordinates with other workers.
|
It listens for external requests via ZMQ and coordinates with other workers.
|
||||||
@@ -132,6 +137,10 @@ class Scheduler(SchedulerDisaggMixin):
|
|||||||
ReleaseRealtimeSessionReq: self._handle_release_realtime_session,
|
ReleaseRealtimeSessionReq: self._handle_release_realtime_session,
|
||||||
GetDisaggStatsReq: self._handle_get_disagg_stats,
|
GetDisaggStatsReq: self._handle_get_disagg_stats,
|
||||||
UpdateWeightFromDiskReqInput: self._handle_update_weights_from_disk,
|
UpdateWeightFromDiskReqInput: self._handle_update_weights_from_disk,
|
||||||
|
UpdateWeightFromTensorReqInput: self._handle_update_weights_from_tensor,
|
||||||
|
UpdateWeightFromTensorCheckerReqInput: (
|
||||||
|
self._handle_update_weights_from_tensor_checker
|
||||||
|
),
|
||||||
GetWeightsChecksumReqInput: self._handle_get_weights_checksum,
|
GetWeightsChecksumReqInput: self._handle_get_weights_checksum,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -211,25 +220,6 @@ class Scheduler(SchedulerDisaggMixin):
|
|||||||
req = reqs[0]
|
req = reqs[0]
|
||||||
return self.worker.release_realtime_session(req.session_id)
|
return self.worker.release_realtime_session(req.session_id)
|
||||||
|
|
||||||
def _handle_update_weights_from_disk(self, reqs: List[Any]) -> OutputBatch:
|
|
||||||
"""Handle update_weights_from_disk request for RL workflows."""
|
|
||||||
req = reqs[0]
|
|
||||||
success, message = self.worker.update_weights_from_disk(
|
|
||||||
model_path=req.model_path,
|
|
||||||
flush_cache=req.flush_cache,
|
|
||||||
target_modules=req.target_modules,
|
|
||||||
)
|
|
||||||
return OutputBatch(
|
|
||||||
output={"success": success, "message": message},
|
|
||||||
error=None if success else message,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _handle_get_weights_checksum(self, reqs: List[Any]) -> OutputBatch:
|
|
||||||
"""Handle get_weights_checksum request."""
|
|
||||||
req = reqs[0]
|
|
||||||
checksums = self.worker.get_weights_checksum(module_names=req.module_names)
|
|
||||||
return OutputBatch(output=checksums)
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _normalize_generation_reqs(reqs: list[Any]) -> list[Req]:
|
def _normalize_generation_reqs(reqs: list[Any]) -> list[Req]:
|
||||||
if len(reqs) == 1 and isinstance(reqs[0], list):
|
if len(reqs) == 1 and isinstance(reqs[0], list):
|
||||||
|
|||||||
@@ -0,0 +1,172 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.distributed import get_tp_rank, get_tp_world_size
|
||||||
|
from sglang.multimodal_gen.runtime.loader.weight_utils import compute_weights_checksum
|
||||||
|
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
||||||
|
iter_materialized_weights,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.post_training.tensor_update_checker import (
|
||||||
|
TensorUpdateChecker,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.post_training.weights_updater import (
|
||||||
|
WeightsUpdater,
|
||||||
|
get_updatable_modules,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils import MultiprocessingSerializer
|
||||||
|
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
|
||||||
|
UpdateWeightFromTensorCheckerReqInput,
|
||||||
|
UpdateWeightFromTensorReqInput,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class GPUWorkerPostTrainingMixin:
|
||||||
|
def update_weights_from_disk(
|
||||||
|
self,
|
||||||
|
model_path: str,
|
||||||
|
flush_cache: bool = True,
|
||||||
|
target_modules: list[str] | None = None,
|
||||||
|
) -> tuple[bool, str]:
|
||||||
|
if not self.pipeline:
|
||||||
|
return False, "Pipeline is not initialized"
|
||||||
|
|
||||||
|
updater = WeightsUpdater(self.pipeline)
|
||||||
|
success, message = updater.update_weights_from_disk(
|
||||||
|
model_path,
|
||||||
|
flush_cache=flush_cache,
|
||||||
|
target_modules=target_modules,
|
||||||
|
)
|
||||||
|
if success:
|
||||||
|
self.server_args.model_path = model_path
|
||||||
|
self.pipeline.model_path = model_path
|
||||||
|
return success, message
|
||||||
|
|
||||||
|
def update_weights_from_tensor(
|
||||||
|
self,
|
||||||
|
req: UpdateWeightFromTensorReqInput,
|
||||||
|
) -> tuple[bool, str]:
|
||||||
|
if not self.pipeline:
|
||||||
|
return False, "Pipeline is not initialized"
|
||||||
|
|
||||||
|
payload, error = self._select_rank_scoped_payload(
|
||||||
|
payloads=req.serialized_named_tensors,
|
||||||
|
field_name="serialized_named_tensors",
|
||||||
|
)
|
||||||
|
if error is not None:
|
||||||
|
return False, error
|
||||||
|
|
||||||
|
monkey_patch_torch_reductions()
|
||||||
|
try:
|
||||||
|
named_tensors = MultiprocessingSerializer.deserialize(payload)
|
||||||
|
except Exception as e:
|
||||||
|
return False, f"Failed to deserialize serialized_named_tensors: {e}"
|
||||||
|
|
||||||
|
updater = WeightsUpdater(self.pipeline)
|
||||||
|
return updater.update_weights_from_tensor(
|
||||||
|
named_tensors=named_tensors,
|
||||||
|
load_format=req.load_format,
|
||||||
|
target_modules=req.target_modules,
|
||||||
|
)
|
||||||
|
|
||||||
|
def update_weights_from_tensor_checker(
|
||||||
|
self,
|
||||||
|
req: UpdateWeightFromTensorCheckerReqInput,
|
||||||
|
) -> tuple[bool, str]:
|
||||||
|
if not self.pipeline:
|
||||||
|
return False, "Pipeline is not initialized"
|
||||||
|
|
||||||
|
checker = TensorUpdateChecker(self.pipeline)
|
||||||
|
result = checker.verify_across_tp(
|
||||||
|
target_module=req.target_module,
|
||||||
|
expected_named_tensors_sha256=req.expected_named_tensors_sha256,
|
||||||
|
tp_rank=get_tp_rank(),
|
||||||
|
tp_world_size=get_tp_world_size(),
|
||||||
|
tp_cpu_group=self.tp_cpu_group,
|
||||||
|
tp_root_rank=self.tp_group.first_rank,
|
||||||
|
)
|
||||||
|
if self.sp_group.world_size == 1:
|
||||||
|
return result
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
is_sp_root = self.sp_group.rank_in_group == 0
|
||||||
|
gathered_results = [None] * self.sp_group.world_size if is_sp_root else None
|
||||||
|
torch.distributed.gather_object(
|
||||||
|
result,
|
||||||
|
gathered_results,
|
||||||
|
dst=self.sp_group.first_rank,
|
||||||
|
group=self.sp_cpu_group,
|
||||||
|
)
|
||||||
|
|
||||||
|
final_result = None
|
||||||
|
if is_sp_root:
|
||||||
|
failures = [
|
||||||
|
(rank, message)
|
||||||
|
for rank, (success, message) in enumerate(gathered_results)
|
||||||
|
if not success
|
||||||
|
]
|
||||||
|
if failures:
|
||||||
|
rank, message = failures[0]
|
||||||
|
if len(failures) == 1:
|
||||||
|
final_result = (False, f"SP rank {rank}: {message}")
|
||||||
|
else:
|
||||||
|
final_result = (
|
||||||
|
False,
|
||||||
|
f"{len(failures)} SP ranks failed update_weight_from_tensor_checker; "
|
||||||
|
f"first failure on rank {rank}: {message}",
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
final_result = result
|
||||||
|
|
||||||
|
final_result_holder = [final_result]
|
||||||
|
torch.distributed.broadcast_object_list(
|
||||||
|
final_result_holder,
|
||||||
|
src=self.sp_group.first_rank,
|
||||||
|
group=self.sp_cpu_group,
|
||||||
|
)
|
||||||
|
return final_result_holder[0]
|
||||||
|
|
||||||
|
def get_weights_checksum(
|
||||||
|
self, module_names: list[str] | None = None
|
||||||
|
) -> dict[str, str]:
|
||||||
|
if not self.pipeline:
|
||||||
|
return {"error": "Pipeline is not initialized"}
|
||||||
|
|
||||||
|
all_modules = get_updatable_modules(self.pipeline)
|
||||||
|
names = module_names if module_names is not None else list(all_modules.keys())
|
||||||
|
|
||||||
|
checksums: dict[str, str] = {}
|
||||||
|
for name in names:
|
||||||
|
module = all_modules.get(name)
|
||||||
|
if module is None:
|
||||||
|
checksums[name] = "not_found"
|
||||||
|
continue
|
||||||
|
checksums[name] = compute_weights_checksum(
|
||||||
|
iter_materialized_weights(module)
|
||||||
|
)
|
||||||
|
return checksums
|
||||||
|
|
||||||
|
def _select_rank_scoped_payload(
|
||||||
|
self,
|
||||||
|
payloads: list,
|
||||||
|
field_name: str,
|
||||||
|
) -> tuple[object | None, str | None]:
|
||||||
|
if not isinstance(payloads, list):
|
||||||
|
return None, f"{field_name} must be a list"
|
||||||
|
if not payloads:
|
||||||
|
return None, f"{field_name} is required"
|
||||||
|
|
||||||
|
tp_world_size = get_tp_world_size()
|
||||||
|
if len(payloads) not in (1, tp_world_size):
|
||||||
|
return (
|
||||||
|
None,
|
||||||
|
f"{field_name} size must be 1 or tp_size ({tp_world_size}), "
|
||||||
|
f"got {len(payloads)}",
|
||||||
|
)
|
||||||
|
|
||||||
|
payload_idx = get_tp_rank() if len(payloads) == tp_world_size else 0
|
||||||
|
return payloads[payload_idx], None
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any, List
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
|
||||||
|
|
||||||
|
|
||||||
|
class SchedulerPostTrainingMixin:
|
||||||
|
def _handle_update_weights_from_disk(self, reqs: List[Any]) -> OutputBatch:
|
||||||
|
req = reqs[0]
|
||||||
|
success, message = self.worker.update_weights_from_disk(
|
||||||
|
model_path=req.model_path,
|
||||||
|
flush_cache=req.flush_cache,
|
||||||
|
target_modules=req.target_modules,
|
||||||
|
)
|
||||||
|
return OutputBatch(
|
||||||
|
output={"success": success, "message": message},
|
||||||
|
error=None if success else message,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _handle_update_weights_from_tensor(self, reqs: List[Any]) -> OutputBatch:
|
||||||
|
req = reqs[0]
|
||||||
|
success, message = self.worker.update_weights_from_tensor(req)
|
||||||
|
if self.server_args.tp_size > 1:
|
||||||
|
import torch
|
||||||
|
|
||||||
|
torch.distributed.barrier(group=self.worker.tp_cpu_group)
|
||||||
|
return OutputBatch(
|
||||||
|
output={"success": success, "message": message},
|
||||||
|
error=None if success else message,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _handle_update_weights_from_tensor_checker(
|
||||||
|
self, reqs: List[Any]
|
||||||
|
) -> OutputBatch:
|
||||||
|
req = reqs[0]
|
||||||
|
success, message = self.worker.update_weights_from_tensor_checker(req)
|
||||||
|
return OutputBatch(
|
||||||
|
output={"success": success, "message": message},
|
||||||
|
error=None if success else message,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _handle_get_weights_checksum(self, reqs: List[Any]) -> OutputBatch:
|
||||||
|
req = reqs[0]
|
||||||
|
checksums = self.worker.get_weights_checksum(module_names=req.module_names)
|
||||||
|
return OutputBatch(output=checksums)
|
||||||
@@ -0,0 +1,263 @@
|
|||||||
|
"""Verification helpers for diffusion update_weights_from_tensor workflows."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
from collections.abc import Iterable
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch.distributed.tensor import DTensor
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
||||||
|
iter_materialized_weights,
|
||||||
|
)
|
||||||
|
|
||||||
|
_MAX_DISPLAY_TENSORS = 5
|
||||||
|
|
||||||
|
|
||||||
|
def _materialize_local_tensor(tensor: torch.Tensor) -> torch.Tensor:
|
||||||
|
if isinstance(tensor, DTensor):
|
||||||
|
tensor = tensor._local_tensor
|
||||||
|
return tensor.detach().cpu().contiguous()
|
||||||
|
|
||||||
|
|
||||||
|
def compute_tensor_sha256(tensor: torch.Tensor) -> str:
|
||||||
|
tensor = _materialize_local_tensor(tensor)
|
||||||
|
hasher = hashlib.sha256()
|
||||||
|
hasher.update(str(tensor.dtype).encode("utf-8"))
|
||||||
|
hasher.update(repr(tuple(tensor.shape)).encode("utf-8"))
|
||||||
|
hasher.update(tensor.view(torch.uint8).numpy().tobytes())
|
||||||
|
return hasher.hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
def build_named_tensor_sha256(
|
||||||
|
named_tensors: Iterable[tuple[str, torch.Tensor]],
|
||||||
|
) -> dict[str, str]:
|
||||||
|
sha256_by_name: dict[str, str] = {}
|
||||||
|
for name, tensor in named_tensors:
|
||||||
|
sha256_by_name[name] = compute_tensor_sha256(tensor)
|
||||||
|
return sha256_by_name
|
||||||
|
|
||||||
|
|
||||||
|
class TensorUpdateChecker:
|
||||||
|
def __init__(self, pipeline):
|
||||||
|
self.pipeline = pipeline
|
||||||
|
|
||||||
|
def verify_across_tp(
|
||||||
|
self,
|
||||||
|
target_module: str,
|
||||||
|
expected_named_tensors_sha256: dict[str, str],
|
||||||
|
tp_rank: int,
|
||||||
|
tp_world_size: int,
|
||||||
|
tp_cpu_group,
|
||||||
|
tp_root_rank: int,
|
||||||
|
) -> tuple[bool, str]:
|
||||||
|
if tp_world_size == 1:
|
||||||
|
return self.verify(
|
||||||
|
target_module=target_module,
|
||||||
|
expected_named_tensors_sha256=expected_named_tensors_sha256,
|
||||||
|
)
|
||||||
|
|
||||||
|
module = self.pipeline.get_module(target_module)
|
||||||
|
if module is None:
|
||||||
|
return False, f"Module '{target_module}' is not initialized"
|
||||||
|
|
||||||
|
local_named_tensors = dict(
|
||||||
|
self._iter_module_named_tensors(
|
||||||
|
module, expected_named_tensors_sha256.keys()
|
||||||
|
)
|
||||||
|
)
|
||||||
|
reference_tensors = dict(module.named_parameters())
|
||||||
|
reference_tensors.update(dict(module.named_buffers()))
|
||||||
|
actual_named_tensors_sha256: dict[str, str] | None = (
|
||||||
|
{} if tp_rank == 0 else None
|
||||||
|
)
|
||||||
|
|
||||||
|
for name, expected_sha256 in expected_named_tensors_sha256.items():
|
||||||
|
gathered_tensors: list[torch.Tensor | None] | None = (
|
||||||
|
[None] * tp_world_size if tp_rank == 0 else None
|
||||||
|
)
|
||||||
|
torch.distributed.gather_object(
|
||||||
|
(
|
||||||
|
_materialize_local_tensor(local_named_tensors[name])
|
||||||
|
if name in local_named_tensors
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
gathered_tensors,
|
||||||
|
dst=tp_root_rank,
|
||||||
|
group=tp_cpu_group,
|
||||||
|
)
|
||||||
|
if tp_rank != 0:
|
||||||
|
continue
|
||||||
|
|
||||||
|
valid_tensors = [
|
||||||
|
tensor for tensor in gathered_tensors if tensor is not None
|
||||||
|
]
|
||||||
|
if len(valid_tensors) != len(gathered_tensors):
|
||||||
|
continue
|
||||||
|
|
||||||
|
local_sha256s = [compute_tensor_sha256(tensor) for tensor in valid_tensors]
|
||||||
|
if all(local_sha256 == expected_sha256 for local_sha256 in local_sha256s):
|
||||||
|
actual_named_tensors_sha256[name] = expected_sha256
|
||||||
|
continue
|
||||||
|
|
||||||
|
reference_tensor = reference_tensors.get(name)
|
||||||
|
candidate_dims: list[int] = []
|
||||||
|
if isinstance(reference_tensor, DTensor):
|
||||||
|
for placement in reference_tensor.placements:
|
||||||
|
shard_dim = getattr(placement, "dim", None)
|
||||||
|
if isinstance(shard_dim, int) and shard_dim not in candidate_dims:
|
||||||
|
candidate_dims.append(shard_dim)
|
||||||
|
for attr in ("input_dim", "output_dim"):
|
||||||
|
shard_dim = getattr(reference_tensor, attr, None)
|
||||||
|
if isinstance(shard_dim, int) and shard_dim not in candidate_dims:
|
||||||
|
candidate_dims.append(shard_dim)
|
||||||
|
|
||||||
|
reconstructed_sha256 = None
|
||||||
|
first_tensor = valid_tensors[0]
|
||||||
|
for shard_dim in candidate_dims:
|
||||||
|
if first_tensor.ndim == 0:
|
||||||
|
break
|
||||||
|
|
||||||
|
shard_dim %= first_tensor.ndim
|
||||||
|
compatible = True
|
||||||
|
for tensor in valid_tensors[1:]:
|
||||||
|
if (
|
||||||
|
tensor.ndim != first_tensor.ndim
|
||||||
|
or tensor.dtype != first_tensor.dtype
|
||||||
|
):
|
||||||
|
compatible = False
|
||||||
|
break
|
||||||
|
if any(
|
||||||
|
lhs != rhs
|
||||||
|
for dim, (lhs, rhs) in enumerate(
|
||||||
|
zip(first_tensor.shape, tensor.shape)
|
||||||
|
)
|
||||||
|
if dim != shard_dim
|
||||||
|
):
|
||||||
|
compatible = False
|
||||||
|
break
|
||||||
|
if not compatible:
|
||||||
|
continue
|
||||||
|
|
||||||
|
reconstructed = torch.cat(valid_tensors, dim=shard_dim).contiguous()
|
||||||
|
if compute_tensor_sha256(reconstructed) == expected_sha256:
|
||||||
|
reconstructed_sha256 = expected_sha256
|
||||||
|
break
|
||||||
|
|
||||||
|
actual_named_tensors_sha256[name] = reconstructed_sha256 or local_sha256s[0]
|
||||||
|
|
||||||
|
final_result: tuple[bool, str] | None = None
|
||||||
|
if tp_rank == 0:
|
||||||
|
final_result = self._compare_manifests(
|
||||||
|
target_module=target_module,
|
||||||
|
expected_named_tensors_sha256=expected_named_tensors_sha256,
|
||||||
|
actual_named_tensors_sha256=actual_named_tensors_sha256,
|
||||||
|
)
|
||||||
|
if final_result[0]:
|
||||||
|
final_result = (
|
||||||
|
True,
|
||||||
|
f"Verified module '{target_module}' update across {tp_world_size} TP ranks.",
|
||||||
|
)
|
||||||
|
|
||||||
|
final_result_holder = [final_result]
|
||||||
|
torch.distributed.broadcast_object_list(
|
||||||
|
final_result_holder,
|
||||||
|
src=tp_root_rank,
|
||||||
|
group=tp_cpu_group,
|
||||||
|
)
|
||||||
|
final_result = final_result_holder[0]
|
||||||
|
assert final_result is not None
|
||||||
|
return final_result
|
||||||
|
|
||||||
|
def verify(
|
||||||
|
self,
|
||||||
|
target_module: str,
|
||||||
|
expected_named_tensors_sha256: dict[str, str],
|
||||||
|
) -> tuple[bool, str]:
|
||||||
|
module = self.pipeline.get_module(target_module)
|
||||||
|
if module is None:
|
||||||
|
return False, f"Module '{target_module}' is not initialized"
|
||||||
|
|
||||||
|
actual_named_tensors_sha256 = build_named_tensor_sha256(
|
||||||
|
self._iter_module_named_tensors(
|
||||||
|
module, expected_named_tensors_sha256.keys()
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return self._compare_manifests(
|
||||||
|
target_module=target_module,
|
||||||
|
expected_named_tensors_sha256=expected_named_tensors_sha256,
|
||||||
|
actual_named_tensors_sha256=actual_named_tensors_sha256,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _iter_module_named_tensors(
|
||||||
|
self,
|
||||||
|
module: torch.nn.Module,
|
||||||
|
expected_names: Iterable[str],
|
||||||
|
):
|
||||||
|
expected_name_set = set(expected_names)
|
||||||
|
seen_names: set[str] = set()
|
||||||
|
|
||||||
|
for name, tensor in iter_materialized_weights(module):
|
||||||
|
if name not in expected_name_set:
|
||||||
|
continue
|
||||||
|
seen_names.add(name)
|
||||||
|
yield name, tensor
|
||||||
|
|
||||||
|
for name, tensor in module.named_buffers():
|
||||||
|
if name in seen_names or name not in expected_name_set:
|
||||||
|
continue
|
||||||
|
seen_names.add(name)
|
||||||
|
yield name, tensor
|
||||||
|
|
||||||
|
def _compare_manifests(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
target_module: str,
|
||||||
|
expected_named_tensors_sha256: dict[str, str],
|
||||||
|
actual_named_tensors_sha256: dict[str, str],
|
||||||
|
) -> tuple[bool, str]:
|
||||||
|
missing_names = sorted(
|
||||||
|
name
|
||||||
|
for name in expected_named_tensors_sha256
|
||||||
|
if name not in actual_named_tensors_sha256
|
||||||
|
)
|
||||||
|
mismatched_names = sorted(
|
||||||
|
name
|
||||||
|
for name, expected_sha256 in expected_named_tensors_sha256.items()
|
||||||
|
if name in actual_named_tensors_sha256
|
||||||
|
and actual_named_tensors_sha256[name] != expected_sha256
|
||||||
|
)
|
||||||
|
|
||||||
|
if missing_names or mismatched_names:
|
||||||
|
parts: list[str] = []
|
||||||
|
if missing_names:
|
||||||
|
parts.append(
|
||||||
|
"missing "
|
||||||
|
f"{len(missing_names)} tensor(s): "
|
||||||
|
f"{self._format_tensor_names(missing_names)}"
|
||||||
|
)
|
||||||
|
if mismatched_names:
|
||||||
|
parts.append(
|
||||||
|
"checksum mismatch for "
|
||||||
|
f"{len(mismatched_names)} tensor(s): "
|
||||||
|
f"{self._format_tensor_names(mismatched_names)}"
|
||||||
|
)
|
||||||
|
return (
|
||||||
|
False,
|
||||||
|
f"Module '{target_module}' update weight check failed: "
|
||||||
|
+ "; ".join(parts),
|
||||||
|
)
|
||||||
|
|
||||||
|
return (
|
||||||
|
True,
|
||||||
|
f"Verified module '{target_module}' update for "
|
||||||
|
f"{len(expected_named_tensors_sha256)} tensor(s).",
|
||||||
|
)
|
||||||
|
|
||||||
|
def _format_tensor_names(self, names: list[str]) -> str:
|
||||||
|
displayed = names[:_MAX_DISPLAY_TENSORS]
|
||||||
|
formatted = ", ".join(displayed)
|
||||||
|
if len(names) > _MAX_DISPLAY_TENSORS:
|
||||||
|
formatted += f", ... (+{len(names) - _MAX_DISPLAY_TENSORS} more)"
|
||||||
|
return formatted
|
||||||
+222
-11
@@ -42,6 +42,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import gc
|
import gc
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch.distributed.tensor import DTensor, distribute_tensor
|
from torch.distributed.tensor import DTensor, distribute_tensor
|
||||||
@@ -49,6 +50,7 @@ from torch.distributed.tensor import DTensor, distribute_tensor
|
|||||||
from sglang.multimodal_gen.runtime.cache.teacache import TeaCacheMixin
|
from sglang.multimodal_gen.runtime.cache.teacache import TeaCacheMixin
|
||||||
from sglang.multimodal_gen.runtime.loader.utils import (
|
from sglang.multimodal_gen.runtime.loader.utils import (
|
||||||
_list_safetensors_files,
|
_list_safetensors_files,
|
||||||
|
get_param_names_mapping,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.weight_utils import (
|
from sglang.multimodal_gen.runtime.loader.weight_utils import (
|
||||||
safetensors_weights_iterator,
|
safetensors_weights_iterator,
|
||||||
@@ -59,8 +61,13 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
|
|||||||
from sglang.multimodal_gen.runtime.pipelines.diffusers_pipeline import DiffusersPipeline
|
from sglang.multimodal_gen.runtime.pipelines.diffusers_pipeline import DiffusersPipeline
|
||||||
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model
|
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
from sglang.srt.weight_sync.tensor_bucket import (
|
||||||
|
FlattenedTensorBucket,
|
||||||
|
FlattenedTensorMetadata,
|
||||||
|
)
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
_DEFAULT_TENSOR_TARGET_MODULE = "transformer"
|
||||||
|
|
||||||
|
|
||||||
def get_updatable_modules(pipeline) -> dict[str, torch.nn.Module]:
|
def get_updatable_modules(pipeline) -> dict[str, torch.nn.Module]:
|
||||||
@@ -115,6 +122,9 @@ def _load_weights_into_module(module: torch.nn.Module, weights_iter) -> None:
|
|||||||
For offloaded modules, updates CPU buffers directly via
|
For offloaded modules, updates CPU buffers directly via
|
||||||
update_cpu_weights(); non-offloaded parameters use in-place copy.
|
update_cpu_weights(); non-offloaded parameters use in-place copy.
|
||||||
"""
|
"""
|
||||||
|
model_params = dict(module.named_parameters())
|
||||||
|
weights_iter = _iter_module_weight_updates(module, weights_iter, model_params)
|
||||||
|
|
||||||
offload_managers: list = []
|
offload_managers: list = []
|
||||||
if is_layerwise_offloaded_module(module):
|
if is_layerwise_offloaded_module(module):
|
||||||
offload_managers = [m for m in module.layerwise_offload_managers if m.enabled]
|
offload_managers = [m for m in module.layerwise_offload_managers if m.enabled]
|
||||||
@@ -125,29 +135,91 @@ def _load_weights_into_module(module: torch.nn.Module, weights_iter) -> None:
|
|||||||
for manager in offload_managers:
|
for manager in offload_managers:
|
||||||
offloaded_names.update(manager.update_cpu_weights(weight_dict))
|
offloaded_names.update(manager.update_cpu_weights(weight_dict))
|
||||||
remaining = ((n, w) for n, w in weight_dict.items() if n not in offloaded_names)
|
remaining = ((n, w) for n, w in weight_dict.items() if n not in offloaded_names)
|
||||||
load_weights_into_model(remaining, dict(module.named_parameters()))
|
load_weights_into_model(remaining, model_params)
|
||||||
else:
|
else:
|
||||||
load_weights_into_model(weights_iter, dict(module.named_parameters()))
|
load_weights_into_model(weights_iter, model_params)
|
||||||
|
|
||||||
|
|
||||||
def load_weights_into_model(weights_iter, model_params: dict) -> None:
|
def _build_module_weight_name_mapper(module: torch.nn.Module):
|
||||||
|
"""Build a chained regex mapper from mapping dicts exposed by the module."""
|
||||||
|
mapping_fns = []
|
||||||
|
for attr in ("lora_param_names_mapping", "param_names_mapping"):
|
||||||
|
mapping = getattr(module, attr, None)
|
||||||
|
if not mapping:
|
||||||
|
continue
|
||||||
|
mapping_fns.append(get_param_names_mapping(mapping))
|
||||||
|
|
||||||
|
if not mapping_fns:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def map_name(name: str) -> str:
|
||||||
|
mapped_name = name
|
||||||
|
for mapping_fn in mapping_fns:
|
||||||
|
mapped_name = mapping_fn(mapped_name)[0]
|
||||||
|
return mapped_name
|
||||||
|
|
||||||
|
return map_name
|
||||||
|
|
||||||
|
|
||||||
|
def _iter_module_weight_updates(
|
||||||
|
module: torch.nn.Module,
|
||||||
|
weights_iter,
|
||||||
|
model_params: dict,
|
||||||
|
):
|
||||||
|
map_name = _build_module_weight_name_mapper(module)
|
||||||
|
module_name = type(module).__name__
|
||||||
|
|
||||||
|
for name, loaded_weight in weights_iter:
|
||||||
|
if name in model_params:
|
||||||
|
yield name, loaded_weight
|
||||||
|
continue
|
||||||
|
|
||||||
|
mapped_name = map_name(name) if map_name is not None else name
|
||||||
|
if mapped_name in model_params:
|
||||||
|
yield mapped_name, loaded_weight
|
||||||
|
continue
|
||||||
|
|
||||||
|
logger.warning(
|
||||||
|
"Skipping weight update for %s: parameter %r not found after mapping to %r",
|
||||||
|
module_name,
|
||||||
|
name,
|
||||||
|
mapped_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def load_weights_into_model(
|
||||||
|
weights_iter, model_params: dict, module_name: str | None = None
|
||||||
|
) -> None:
|
||||||
"""Copy weights from weights_iter into model_params in-place."""
|
"""Copy weights from weights_iter into model_params in-place."""
|
||||||
for name, loaded_weight in weights_iter:
|
for name, loaded_weight in weights_iter:
|
||||||
if name not in model_params:
|
if name not in model_params:
|
||||||
|
logger.warning("Skipping weight update: parameter %r not found", name)
|
||||||
continue
|
continue
|
||||||
param = model_params[name]
|
param = model_params[name]
|
||||||
if param.shape != loaded_weight.shape:
|
weight_loader = getattr(param, "weight_loader", None)
|
||||||
raise ValueError(
|
if callable(weight_loader):
|
||||||
f"Shape mismatch for {name}: model={param.shape}, loaded={loaded_weight.shape}"
|
weight_loader(param, loaded_weight.to(param.dtype))
|
||||||
)
|
else:
|
||||||
if isinstance(param, DTensor):
|
dtensor_param = param if isinstance(param, DTensor) else None
|
||||||
|
if dtensor_param is None and isinstance(
|
||||||
|
getattr(param, "data", None), DTensor
|
||||||
|
):
|
||||||
|
dtensor_param = param.data
|
||||||
|
|
||||||
|
if dtensor_param is not None:
|
||||||
distributed_weight = distribute_tensor(
|
distributed_weight = distribute_tensor(
|
||||||
loaded_weight.to(param.dtype),
|
loaded_weight.to(param.dtype),
|
||||||
param.device_mesh,
|
dtensor_param.device_mesh,
|
||||||
param.placements,
|
dtensor_param.placements,
|
||||||
)
|
)
|
||||||
param._local_tensor.copy_(distributed_weight._local_tensor)
|
dtensor_param._local_tensor.copy_(distributed_weight._local_tensor)
|
||||||
else:
|
else:
|
||||||
|
if param.shape != loaded_weight.shape:
|
||||||
|
module_prefix = f"{module_name}." if module_name else ""
|
||||||
|
raise ValueError(
|
||||||
|
f"Shape mismatch for {module_prefix}{name}: "
|
||||||
|
f"model={param.shape}, loaded={loaded_weight.shape}"
|
||||||
|
)
|
||||||
param.data.copy_(loaded_weight.to(param.dtype))
|
param.data.copy_(loaded_weight.to(param.dtype))
|
||||||
|
|
||||||
|
|
||||||
@@ -293,3 +365,142 @@ class WeightsUpdater:
|
|||||||
continue
|
continue
|
||||||
weights_iter = _get_weights_iter(str(weights_dir))
|
weights_iter = _get_weights_iter(str(weights_dir))
|
||||||
_load_weights_into_module(module, weights_iter)
|
_load_weights_into_module(module, weights_iter)
|
||||||
|
|
||||||
|
def update_weights_from_tensor(
|
||||||
|
self,
|
||||||
|
named_tensors: Any,
|
||||||
|
load_format: str | None = None,
|
||||||
|
target_modules: list[str] | None = None,
|
||||||
|
) -> tuple[bool, str]:
|
||||||
|
if target_modules is None:
|
||||||
|
target_modules = [_DEFAULT_TENSOR_TARGET_MODULE]
|
||||||
|
try:
|
||||||
|
modules_to_update = self._collect_modules(target_modules)
|
||||||
|
except ValueError as e:
|
||||||
|
logger.error(str(e))
|
||||||
|
return False, str(e)
|
||||||
|
|
||||||
|
if not modules_to_update:
|
||||||
|
error_msg = (
|
||||||
|
f"No matching modules found for update. "
|
||||||
|
f"Requested: {target_modules}. "
|
||||||
|
f"Available nn.Module(s): {list(get_updatable_modules(self.pipeline).keys())}"
|
||||||
|
)
|
||||||
|
logger.error(error_msg)
|
||||||
|
return False, error_msg
|
||||||
|
|
||||||
|
try:
|
||||||
|
module_payloads = self._resolve_module_payloads(
|
||||||
|
named_tensors=named_tensors,
|
||||||
|
modules_to_update=modules_to_update,
|
||||||
|
)
|
||||||
|
except ValueError as e:
|
||||||
|
logger.error(str(e))
|
||||||
|
return False, str(e)
|
||||||
|
|
||||||
|
updated_modules: list[str] = []
|
||||||
|
for module_name, module in modules_to_update:
|
||||||
|
try:
|
||||||
|
payload = module_payloads[module_name]
|
||||||
|
weights_iter = self._materialize_weights_iter(payload, load_format)
|
||||||
|
_load_weights_into_module(module, weights_iter)
|
||||||
|
updated_modules.append(module_name)
|
||||||
|
except Exception as e:
|
||||||
|
error_msg = (
|
||||||
|
f"Failed to update module '{module_name}' from tensor: {e}. "
|
||||||
|
f"The pipeline may be partially updated. "
|
||||||
|
f"Please discard the whole weights and reload from a known-good checkpoint."
|
||||||
|
)
|
||||||
|
logger.error(error_msg, exc_info=True)
|
||||||
|
return False, error_msg
|
||||||
|
|
||||||
|
gc.collect()
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
names = ", ".join(updated_modules)
|
||||||
|
message = f"Updated {len(updated_modules)} modules from tensor ({names})."
|
||||||
|
logger.info(message)
|
||||||
|
return True, message
|
||||||
|
|
||||||
|
def _resolve_module_payloads(
|
||||||
|
self,
|
||||||
|
named_tensors: Any,
|
||||||
|
modules_to_update: list[tuple[str, torch.nn.Module]],
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
module_names = [name for name, _ in modules_to_update]
|
||||||
|
if isinstance(named_tensors, dict):
|
||||||
|
missing = [name for name in module_names if name not in named_tensors]
|
||||||
|
if missing:
|
||||||
|
raise ValueError(
|
||||||
|
f"Missing tensor payload for module(s): {missing}. "
|
||||||
|
f"Provided modules: {list(named_tensors.keys())}"
|
||||||
|
)
|
||||||
|
return {name: named_tensors[name] for name in module_names}
|
||||||
|
|
||||||
|
if len(module_names) == 1:
|
||||||
|
return {module_names[0]: named_tensors}
|
||||||
|
|
||||||
|
raise ValueError(
|
||||||
|
"Ambiguous tensor payload for multi-module update. "
|
||||||
|
"Provide a dict mapping module_name -> module payload, "
|
||||||
|
f"requested modules: {module_names}."
|
||||||
|
)
|
||||||
|
|
||||||
|
def _materialize_weights_iter(self, module_payload: Any, load_format: str | None):
|
||||||
|
if load_format == "flattened_bucket":
|
||||||
|
if not isinstance(module_payload, dict):
|
||||||
|
raise ValueError(
|
||||||
|
"flattened_bucket payload must be a dict with "
|
||||||
|
"'flattened_tensor' and 'metadata'."
|
||||||
|
)
|
||||||
|
flattened_tensor = module_payload.get("flattened_tensor")
|
||||||
|
metadata = module_payload.get("metadata")
|
||||||
|
if flattened_tensor is None or metadata is None:
|
||||||
|
raise ValueError(
|
||||||
|
"flattened_bucket payload missing 'flattened_tensor' or 'metadata'."
|
||||||
|
)
|
||||||
|
return self._reconstruct_from_flattened_bucket(flattened_tensor, metadata)
|
||||||
|
|
||||||
|
if isinstance(module_payload, (list, tuple)):
|
||||||
|
return iter(module_payload)
|
||||||
|
|
||||||
|
raise ValueError(
|
||||||
|
f"Unsupported module payload type for load_format={load_format}: "
|
||||||
|
f"{type(module_payload).__name__}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def _reconstruct_from_flattened_bucket(self, flattened_tensor: Any, metadata: Any):
|
||||||
|
if not isinstance(flattened_tensor, torch.Tensor):
|
||||||
|
raise ValueError(
|
||||||
|
"flattened_bucket 'flattened_tensor' must be a torch.Tensor."
|
||||||
|
)
|
||||||
|
if not isinstance(metadata, list):
|
||||||
|
raise ValueError("flattened_bucket 'metadata' must be a list.")
|
||||||
|
|
||||||
|
converted_metadata: list[FlattenedTensorMetadata] = []
|
||||||
|
for meta in metadata:
|
||||||
|
converted_metadata.append(
|
||||||
|
FlattenedTensorMetadata(
|
||||||
|
name=meta.name,
|
||||||
|
shape=torch.Size(meta.shape),
|
||||||
|
dtype=self._normalize_torch_dtype(meta.dtype),
|
||||||
|
start_idx=int(meta.start_idx),
|
||||||
|
end_idx=int(meta.end_idx),
|
||||||
|
numel=int(meta.numel),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
bucket = FlattenedTensorBucket(
|
||||||
|
flattened_tensor=flattened_tensor,
|
||||||
|
metadata=converted_metadata,
|
||||||
|
)
|
||||||
|
return bucket.reconstruct_tensors()
|
||||||
|
|
||||||
|
def _normalize_torch_dtype(self, dtype: Any) -> torch.dtype:
|
||||||
|
if isinstance(dtype, torch.dtype):
|
||||||
|
return dtype
|
||||||
|
if isinstance(dtype, str):
|
||||||
|
name = dtype.split(".")[-1]
|
||||||
|
normalized = getattr(torch, name, None)
|
||||||
|
if isinstance(normalized, torch.dtype):
|
||||||
|
return normalized
|
||||||
|
raise ValueError(f"Unsupported dtype in flattened_bucket metadata: {dtype!r}")
|
||||||
Reference in New Issue
Block a user