diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py index e736a65b9..45b3676c6 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/weights_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/weights_api.py index 8f4fd1270..554feea4e 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/weights_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/weights_api.py @@ -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: diff --git a/python/sglang/multimodal_gen/runtime/pipelines/stable_diffusion_3.py b/python/sglang/multimodal_gen/runtime/pipelines/stable_diffusion_3.py index f0b8c0b9f..039838040 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/stable_diffusion_3.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/stable_diffusion_3.py @@ -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" diff --git a/python/sglang/multimodal_gen/runtime/post_training/gpu_worker_post_training_mixin.py b/python/sglang/multimodal_gen/runtime/post_training/gpu_worker_post_training_mixin.py index ee0045d87..f6520591f 100644 --- a/python/sglang/multimodal_gen/runtime/post_training/gpu_worker_post_training_mixin.py +++ b/python/sglang/multimodal_gen/runtime/post_training/gpu_worker_post_training_mixin.py @@ -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( diff --git a/python/sglang/multimodal_gen/runtime/post_training/weights_updater.py b/python/sglang/multimodal_gen/runtime/post_training/weights_updater.py index 0b69ba2ac..b5d35136b 100644 --- a/python/sglang/multimodal_gen/runtime/post_training/weights_updater.py +++ b/python/sglang/multimodal_gen/runtime/post_training/weights_updater.py @@ -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,