[diffusion] post_training: Add LoRA IPC weight sync via lora_merge mode (#31029)
This commit is contained in:
@@ -24,6 +24,9 @@ class UpdateWeightFromTensorReqInput:
|
||||
serialized_named_tensors: list[str | bytes]
|
||||
load_format: str | None = None
|
||||
target_modules: list[str] | None = None
|
||||
weight_update_mode: str | None = None
|
||||
lora_alpha: int | None = None
|
||||
lora_rank: int | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -72,6 +72,9 @@ async def update_weights_from_tensor(request: Request):
|
||||
serialized_named_tensors=serialized_named_tensors,
|
||||
load_format=body.get("load_format"),
|
||||
target_modules=body.get("target_modules"),
|
||||
weight_update_mode=body.get("weight_update_mode"),
|
||||
lora_alpha=body.get("lora_alpha"),
|
||||
lora_rank=body.get("lora_rank"),
|
||||
)
|
||||
|
||||
try:
|
||||
|
||||
@@ -8,6 +8,7 @@ import torch
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
||||
InputValidationStage,
|
||||
@@ -64,7 +65,7 @@ class SD3ConditioningStage(PipelineStage):
|
||||
return merged_embeds, merged_pooled
|
||||
|
||||
|
||||
class StableDiffusion3Pipeline(ComposedPipelineBase):
|
||||
class StableDiffusion3Pipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""StableDiffusion3 pipeline implementation."""
|
||||
|
||||
pipeline_name = "StableDiffusion3Pipeline"
|
||||
|
||||
@@ -70,6 +70,9 @@ class GPUWorkerPostTrainingMixin:
|
||||
named_tensors=named_tensors,
|
||||
load_format=req.load_format,
|
||||
target_modules=req.target_modules,
|
||||
weight_update_mode=req.weight_update_mode,
|
||||
lora_alpha=req.lora_alpha,
|
||||
lora_rank=req.lora_rank,
|
||||
)
|
||||
|
||||
def update_weights_from_tensor_checker(
|
||||
|
||||
@@ -58,6 +58,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
|
||||
is_layerwise_offloaded_module,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines.diffusers_pipeline import DiffusersPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
|
||||
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.srt.weight_sync.tensor_bucket import (
|
||||
@@ -67,6 +68,54 @@ from sglang.srt.weight_sync.tensor_bucket import (
|
||||
|
||||
logger = init_logger(__name__)
|
||||
_DEFAULT_TENSOR_TARGET_MODULE = "transformer"
|
||||
LORA_MERGE_WEIGHT_UPDATE_MODE = "lora_merge"
|
||||
_LORA_IPC_TARGET_MODULES = frozenset({"transformer", "transformer_2"})
|
||||
|
||||
|
||||
def _get_lora_layer_dict(
|
||||
lora_pipeline: LoRAPipeline, target_module: str
|
||||
) -> dict[str, object]:
|
||||
if target_module == "transformer":
|
||||
return lora_pipeline.lora_layers
|
||||
if target_module == "transformer_2":
|
||||
if not lora_pipeline.lora_layers_transformer_2:
|
||||
raise ValueError(
|
||||
"transformer_2 is not present or has no LoRA layers in this pipeline"
|
||||
)
|
||||
return lora_pipeline.lora_layers_transformer_2
|
||||
raise ValueError(
|
||||
f"Unsupported LoRA IPC target_module={target_module!r}; "
|
||||
f"expected one of {sorted(_LORA_IPC_TARGET_MODULES)}"
|
||||
)
|
||||
|
||||
|
||||
def _group_lora_ab_tensors(
|
||||
named_tensors: list[tuple[str, torch.Tensor]],
|
||||
) -> dict[str, tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""Group flattened IPC tensors into {layer_name: (lora_A, lora_B)} pairs."""
|
||||
partial: dict[str, dict[str, torch.Tensor]] = {}
|
||||
for name, tensor in named_tensors:
|
||||
if ".lora_A" in name:
|
||||
layer_name = name.split(".lora_A", 1)[0]
|
||||
partial.setdefault(layer_name, {})["A"] = tensor
|
||||
elif ".lora_B" in name:
|
||||
layer_name = name.split(".lora_B", 1)[0]
|
||||
partial.setdefault(layer_name, {})["B"] = tensor
|
||||
|
||||
pairs: dict[str, tuple[torch.Tensor, torch.Tensor]] = {}
|
||||
for layer_name, ab in partial.items():
|
||||
lora_a = ab.get("A")
|
||||
lora_b = ab.get("B")
|
||||
if lora_a is None or lora_b is None:
|
||||
logger.warning(
|
||||
"Incomplete LoRA pair for layer %s (has_A=%s has_B=%s); skipping",
|
||||
layer_name,
|
||||
lora_a is not None,
|
||||
lora_b is not None,
|
||||
)
|
||||
continue
|
||||
pairs[layer_name] = (lora_a, lora_b)
|
||||
return pairs
|
||||
|
||||
|
||||
def get_updatable_modules(pipeline) -> dict[str, torch.nn.Module]:
|
||||
@@ -160,6 +209,37 @@ def _build_module_weight_name_mapper(module: torch.nn.Module):
|
||||
return map_name
|
||||
|
||||
|
||||
def _strip_param_weight_suffix(param_name: str) -> str:
|
||||
if param_name.endswith(".weight"):
|
||||
return param_name[: -len(".weight")]
|
||||
if param_name.endswith(".bias"):
|
||||
return param_name[: -len(".bias")]
|
||||
return param_name
|
||||
|
||||
|
||||
def _resolve_lora_ipc_layer_dict_key(
|
||||
layer_prefix: str,
|
||||
layer_dict: dict,
|
||||
module: torch.nn.Module,
|
||||
) -> tuple[Any | None, str]:
|
||||
"""Map training-side LoRA layer prefix to lora_layers key (Layer 2)."""
|
||||
layer = layer_dict.get(layer_prefix)
|
||||
if layer is not None:
|
||||
return layer, layer_prefix
|
||||
|
||||
map_name = _build_module_weight_name_mapper(module)
|
||||
if map_name is None:
|
||||
return None, layer_prefix
|
||||
|
||||
mapped = _strip_param_weight_suffix(map_name(f"{layer_prefix}.weight"))
|
||||
if mapped != layer_prefix:
|
||||
layer = layer_dict.get(mapped)
|
||||
if layer is not None:
|
||||
return layer, mapped
|
||||
|
||||
return None, layer_prefix
|
||||
|
||||
|
||||
def _iter_module_weight_updates(
|
||||
module: torch.nn.Module,
|
||||
weights_iter,
|
||||
@@ -385,7 +465,19 @@ class WeightsUpdater:
|
||||
named_tensors: Any,
|
||||
load_format: str | None = None,
|
||||
target_modules: list[str] | None = None,
|
||||
weight_update_mode: str | None = None,
|
||||
lora_alpha: int | None = None,
|
||||
lora_rank: int | None = None,
|
||||
) -> tuple[bool, str]:
|
||||
if weight_update_mode == LORA_MERGE_WEIGHT_UPDATE_MODE:
|
||||
return self._update_lora_from_tensor(
|
||||
named_tensors=named_tensors,
|
||||
load_format=load_format,
|
||||
target_modules=target_modules,
|
||||
lora_alpha=lora_alpha,
|
||||
lora_rank=lora_rank,
|
||||
)
|
||||
|
||||
if target_modules is None:
|
||||
target_modules = [_DEFAULT_TENSOR_TARGET_MODULE]
|
||||
try:
|
||||
@@ -435,6 +527,141 @@ class WeightsUpdater:
|
||||
logger.info(message)
|
||||
return True, message
|
||||
|
||||
def _update_lora_from_tensor(
|
||||
self,
|
||||
named_tensors: Any,
|
||||
load_format: str | None,
|
||||
target_modules: list[str] | None,
|
||||
lora_alpha: int | None,
|
||||
lora_rank: int | None,
|
||||
) -> tuple[bool, str]:
|
||||
if not isinstance(self.pipeline, LoRAPipeline):
|
||||
return (
|
||||
False,
|
||||
"LoRA merge weight update requires a LoRAPipeline-backed model",
|
||||
)
|
||||
|
||||
if target_modules is None:
|
||||
target_modules = [_DEFAULT_TENSOR_TARGET_MODULE]
|
||||
if len(target_modules) != 1:
|
||||
return (
|
||||
False,
|
||||
"LoRA IPC weight update requires exactly one target module per request",
|
||||
)
|
||||
target_module = target_modules[0]
|
||||
if target_module not in _LORA_IPC_TARGET_MODULES:
|
||||
return (
|
||||
False,
|
||||
f"LoRA IPC weight update supports target_modules in "
|
||||
f"{sorted(_LORA_IPC_TARGET_MODULES)}, got {target_module!r}",
|
||||
)
|
||||
|
||||
try:
|
||||
modules_to_update = self._collect_modules([target_module])
|
||||
except ValueError as e:
|
||||
logger.error(str(e))
|
||||
return False, str(e)
|
||||
|
||||
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)
|
||||
|
||||
materialized: list[tuple[str, torch.Tensor]] = []
|
||||
for module_name, _module in modules_to_update:
|
||||
payload = module_payloads[module_name]
|
||||
weights_iter = self._materialize_weights_iter(payload, load_format)
|
||||
materialized.extend(list(weights_iter))
|
||||
|
||||
pairs = _group_lora_ab_tensors(materialized)
|
||||
if not pairs:
|
||||
return False, "No LoRA A/B tensor pairs found in payload"
|
||||
|
||||
lora_pipeline: LoRAPipeline = self.pipeline
|
||||
if not lora_pipeline.lora_initialized:
|
||||
convert_target = (
|
||||
"all"
|
||||
if "transformer_2" in get_updatable_modules(lora_pipeline)
|
||||
else "transformer"
|
||||
)
|
||||
# Match disk LoRA loading: wrap all supported Linear layers regardless
|
||||
# of lora_target_modules. Training-side HF keys are resolved at write time.
|
||||
saved_lora_target_modules = lora_pipeline.lora_target_modules
|
||||
lora_pipeline.lora_target_modules = None
|
||||
try:
|
||||
with lora_pipeline._temporarily_disable_offload(
|
||||
target=convert_target, use_module_names_only=True
|
||||
):
|
||||
lora_pipeline.convert_to_lora_layers()
|
||||
finally:
|
||||
lora_pipeline.lora_target_modules = saved_lora_target_modules
|
||||
|
||||
try:
|
||||
layer_dict = _get_lora_layer_dict(lora_pipeline, target_module)
|
||||
except ValueError as e:
|
||||
logger.error(str(e))
|
||||
return False, str(e)
|
||||
|
||||
dit_module = dict(modules_to_update).get(target_module)
|
||||
if dit_module is None:
|
||||
return False, f"No DiT module found for LoRA IPC target {target_module!r}"
|
||||
|
||||
updated = 0
|
||||
skipped = 0
|
||||
unknown_layers: list[str] = []
|
||||
with lora_pipeline._temporarily_disable_offload(target=target_module):
|
||||
for layer_name, (lora_a, lora_b) in pairs.items():
|
||||
layer, _resolved_key = _resolve_lora_ipc_layer_dict_key(
|
||||
layer_name, layer_dict, dit_module
|
||||
)
|
||||
if layer is None:
|
||||
logger.warning(
|
||||
"Unknown LoRA layer name %s for target %s; skipping",
|
||||
layer_name,
|
||||
target_module,
|
||||
)
|
||||
unknown_layers.append(layer_name)
|
||||
skipped += 1
|
||||
continue
|
||||
inferred_rank = int(lora_a.shape[0])
|
||||
alpha = lora_alpha if lora_alpha is not None else inferred_rank
|
||||
if lora_rank is not None and lora_rank != inferred_rank:
|
||||
logger.warning(
|
||||
"LoRA rank mismatch for %s: payload=%d request=%d; using payload rank",
|
||||
layer_name,
|
||||
inferred_rank,
|
||||
lora_rank,
|
||||
)
|
||||
layer.lora_rank = inferred_rank
|
||||
layer.lora_alpha = alpha
|
||||
layer.set_lora_weights(
|
||||
lora_a, lora_b, merge_weights=True, clear_existing=True
|
||||
)
|
||||
updated += 1
|
||||
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
if updated == 0:
|
||||
sample = unknown_layers[:5]
|
||||
return (
|
||||
False,
|
||||
f"No LoRA layers updated for {target_module} ({skipped} unknown layer names"
|
||||
f"{f', e.g. {sample}' if sample else ''}); "
|
||||
"check training-side layer name mapping",
|
||||
)
|
||||
|
||||
message = (
|
||||
f"Updated {updated} LoRA layers in {target_module} from IPC tensors "
|
||||
f"(skipped {skipped} unknown layers)."
|
||||
)
|
||||
logger.info(message)
|
||||
return True, message
|
||||
|
||||
def _resolve_module_payloads(
|
||||
self,
|
||||
named_tensors: Any,
|
||||
|
||||
Reference in New Issue
Block a user