[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]
|
serialized_named_tensors: list[str | bytes]
|
||||||
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
|
||||||
|
lora_alpha: int | None = None
|
||||||
|
lora_rank: int | None = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -72,6 +72,9 @@ async def update_weights_from_tensor(request: Request):
|
|||||||
serialized_named_tensors=serialized_named_tensors,
|
serialized_named_tensors=serialized_named_tensors,
|
||||||
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"),
|
||||||
|
lora_alpha=body.get("lora_alpha"),
|
||||||
|
lora_rank=body.get("lora_rank"),
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import torch
|
|||||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||||
ComposedPipelineBase,
|
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.schedule_batch import Req
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
||||||
InputValidationStage,
|
InputValidationStage,
|
||||||
@@ -64,7 +65,7 @@ class SD3ConditioningStage(PipelineStage):
|
|||||||
return merged_embeds, merged_pooled
|
return merged_embeds, merged_pooled
|
||||||
|
|
||||||
|
|
||||||
class StableDiffusion3Pipeline(ComposedPipelineBase):
|
class StableDiffusion3Pipeline(LoRAPipeline, ComposedPipelineBase):
|
||||||
"""StableDiffusion3 pipeline implementation."""
|
"""StableDiffusion3 pipeline implementation."""
|
||||||
|
|
||||||
pipeline_name = "StableDiffusion3Pipeline"
|
pipeline_name = "StableDiffusion3Pipeline"
|
||||||
|
|||||||
@@ -70,6 +70,9 @@ class GPUWorkerPostTrainingMixin:
|
|||||||
named_tensors=named_tensors,
|
named_tensors=named_tensors,
|
||||||
load_format=req.load_format,
|
load_format=req.load_format,
|
||||||
target_modules=req.target_modules,
|
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(
|
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,
|
is_layerwise_offloaded_module,
|
||||||
)
|
)
|
||||||
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.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.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 (
|
from sglang.srt.weight_sync.tensor_bucket import (
|
||||||
@@ -67,6 +68,54 @@ from sglang.srt.weight_sync.tensor_bucket import (
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
_DEFAULT_TENSOR_TARGET_MODULE = "transformer"
|
_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]:
|
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
|
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(
|
def _iter_module_weight_updates(
|
||||||
module: torch.nn.Module,
|
module: torch.nn.Module,
|
||||||
weights_iter,
|
weights_iter,
|
||||||
@@ -385,7 +465,19 @@ class WeightsUpdater:
|
|||||||
named_tensors: Any,
|
named_tensors: Any,
|
||||||
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,
|
||||||
|
lora_alpha: int | None = None,
|
||||||
|
lora_rank: int | None = None,
|
||||||
) -> tuple[bool, str]:
|
) -> 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:
|
if target_modules is None:
|
||||||
target_modules = [_DEFAULT_TENSOR_TARGET_MODULE]
|
target_modules = [_DEFAULT_TENSOR_TARGET_MODULE]
|
||||||
try:
|
try:
|
||||||
@@ -435,6 +527,141 @@ class WeightsUpdater:
|
|||||||
logger.info(message)
|
logger.info(message)
|
||||||
return True, 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(
|
def _resolve_module_payloads(
|
||||||
self,
|
self,
|
||||||
named_tensors: Any,
|
named_tensors: Any,
|
||||||
|
|||||||
Reference in New Issue
Block a user