diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py index d7979cff3..5c902c406 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/lora_pipeline.py @@ -4,7 +4,7 @@ import os from collections import defaultdict from collections.abc import Hashable -from contextlib import contextmanager +from contextlib import contextmanager, nullcontext from typing import Any import torch @@ -19,6 +19,9 @@ from sglang.multimodal_gen.runtime.layers.lora.linear import ( wrap_with_lora_layer, ) from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping +from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( + is_layerwise_offloaded_module, +) from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import ( ComposedPipelineBase, ) @@ -170,10 +173,6 @@ class LoRAPipeline(ComposedPipelineBase): Yields: List of modules that had offload disabled. """ - from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( - is_layerwise_offloaded_module, - ) - module_names = [] if target_modules is not None: # Extract module names from target_modules @@ -215,6 +214,22 @@ class LoRAPipeline(ComposedPipelineBase): for module in offload_disabled_modules: module.enable_offload() + def _needs_lora_weight_update_context( + self, + target_modules: list[tuple[str, dict[str, BaseLayerWithLoRA]]], + merge_weights_by_module: dict[str, bool], + ) -> bool: + + for module_name, lora_layers_dict in target_modules: + if merge_weights_by_module[module_name]: + return True + if any(layer.merged for layer in lora_layers_dict.values()): + return True + module = self.modules.get(module_name) + if module is not None and is_layerwise_offloaded_module(module): + return True + return False + def convert_module_lora_layers( self, module: torch.nn.Module, @@ -578,6 +593,58 @@ class LoRAPipeline(ComposedPipelineBase): ) return adapted_count + def _reactivate_cached_dynamic_lora_layers( + self, + lora_layers: dict[str, BaseLayerWithLoRA], + lora_nicknames: list[str], + lora_paths: list[str | None], + strengths: list[float], + ) -> int | None: + """ + Re-enable a previously applied dynamic LoRA without rebuilding per-layer state. + + Dynamic LoRA keeps adapter tensors on the wrapped layers. When a later stage only + disables them and the next stage asks for the same single adapter again, toggling + `disable_lora` is enough; the stored A/B tensors, rank, alpha, and strength still + describe the requested adapter. + """ + if len(lora_nicknames) != 1: + return None + + nickname = lora_nicknames[0] + strength = strengths[0] + adapter = self.lora_adapters.get(nickname) + if adapter is None: + return None + path = lora_paths[0] or self.loaded_adapter_paths.get(nickname) + if path is None: + return None + + active_count = 0 + for name, layer in lora_layers.items(): + if layer.merged or len(layer.lora_weights_list) != 1: + return None + has_adapter = name + ".lora_A" in adapter and name + ".lora_B" in adapter + if not has_adapter: + continue + if ( + layer.lora_A is None + or layer.lora_B is None + or layer.lora_path != path + or layer.strength != strength + ): + return None + active_count += 1 + + if active_count == 0: + return None + + for name, layer in lora_layers.items(): + has_adapter = name + ".lora_A" in adapter and name + ".lora_B" in adapter + layer.disable_lora = not has_adapter + + return active_count + def is_lora_effective(self, target: str = "all") -> bool: """ Check if LoRA is currently effective for the specified target. @@ -756,58 +823,79 @@ class LoRAPipeline(ComposedPipelineBase): if not target_modules: continue - # Disable layerwise offload if enabled: load all layers to GPU - # the LoRA weights merging process requires weights being on device - with self._temporarily_disable_offload(target_modules=target_modules): - tgt_nicknames = [lora_nicknames[i] for i in idx_list] - tgt_paths = [lora_paths[i] for i in idx_list] - tgt_strengths = [strengths[i] for i in idx_list] + tgt_nicknames = [lora_nicknames[i] for i in idx_list] + tgt_paths = [lora_paths[i] for i in idx_list] + tgt_strengths = [strengths[i] for i in idx_list] - merged_name = ( - ",".join(tgt_nicknames) - if len(tgt_nicknames) > 1 - else tgt_nicknames[0] + merged_name = ( + ",".join(tgt_nicknames) if len(tgt_nicknames) > 1 else tgt_nicknames[0] + ) + + # Skip if LoRA configuration matches exactly (including order and strength) + # Since all modules for the same target apply the same config, checking one is sufficient + first_module_name, first_lora_layers_dict = target_modules[0] + first_effective_merge_weights = self._should_merge_lora_for_layers( + first_module_name, first_lora_layers_dict, merge_mode + ) + if not first_effective_merge_weights and len(tgt_nicknames) > 1: + raise ValueError( + "Dynamic LoRA currently supports only one adapter per target. " + "Use merge_mode='merge' for multiple adapters." ) - # Skip if LoRA configuration matches exactly (including order and strength) - # Since all modules for the same target apply the same config, checking one is sufficient - first_module_name, first_lora_layers_dict = target_modules[0] - first_effective_merge_weights = self._should_merge_lora_for_layers( - first_module_name, first_lora_layers_dict, merge_mode - ) - if not first_effective_merge_weights and len(tgt_nicknames) > 1: - raise ValueError( - "Dynamic LoRA currently supports only one adapter per target. " - "Use merge_mode='merge' for multiple adapters." + merge_weights_by_module = {} + for module_name, lora_layers_dict in target_modules: + merge_weights_by_module[module_name] = ( + first_effective_merge_weights + if module_name == first_module_name + else self._should_merge_lora_for_layers( + module_name, lora_layers_dict, merge_mode ) - if self._check_lora_config_matches( - first_module_name, - tgt_nicknames, - tgt_strengths, - first_effective_merge_weights, - adapter_updated, - ): - logger.info("LoRA configuration matches exactly, skipping") - continue + ) + if self._check_lora_config_matches( + first_module_name, + tgt_nicknames, + tgt_strengths, + first_effective_merge_weights, + adapter_updated, + ): + logger.info("LoRA configuration matches exactly, skipping") + continue + + # merged LoRA and offloaded modules update backing weights; dynamic + # reactivation only toggles wrapper metadata when cached tensors match + if self._needs_lora_weight_update_context( + target_modules, merge_weights_by_module + ): + weight_update_context = self._temporarily_disable_offload( + target_modules=target_modules + ) + else: + weight_update_context = nullcontext() + + with weight_update_context: # Apply LoRA to modules for this target for module_name, lora_layers_dict in target_modules: - effective_merge_weights = ( - first_effective_merge_weights - if module_name == first_module_name - else self._should_merge_lora_for_layers( - module_name, lora_layers_dict, merge_mode + effective_merge_weights = merge_weights_by_module[module_name] + count = None + if not effective_merge_weights and not adapter_updated: + count = self._reactivate_cached_dynamic_lora_layers( + lora_layers_dict, + tgt_nicknames, + tgt_paths, + tgt_strengths, + ) + if count is None: + count = self._apply_lora_to_layers( + lora_layers_dict, + tgt_nicknames, + tgt_paths, + rank, + tgt_strengths, + clear_existing=True, + merge_weights=effective_merge_weights, ) - ) - count = self._apply_lora_to_layers( - lora_layers_dict, - tgt_nicknames, - tgt_paths, - rank, - tgt_strengths, - clear_existing=True, - merge_weights=effective_merge_weights, - ) adapted_count += count self.cur_adapter_name[module_name] = merged_name self.cur_adapter_path[module_name] = ",".join( diff --git a/python/sglang/multimodal_gen/test/server/perf_baselines.json b/python/sglang/multimodal_gen/test/server/perf_baselines.json index aec8a7ffe..9aaf37698 100644 --- a/python/sglang/multimodal_gen/test/server/perf_baselines.json +++ b/python/sglang/multimodal_gen/test/server/perf_baselines.json @@ -2442,7 +2442,7 @@ "per_frame_generation": null }, "denoise_step_ms": { - "0": 179.26, + "0": 225.43, "1": 283.15, "2": 193.3, "3": 166.82, @@ -2534,7 +2534,7 @@ "32": 209.80 }, "expected_e2e_ms": 12000.0, - "expected_avg_denoise_ms": 238.8, + "expected_avg_denoise_ms": 299.01, "expected_median_denoise_ms": 246.85, "estimated_full_test_time_s": 170.0 }, diff --git a/python/sglang/multimodal_gen/test/unit/test_lora_pipeline.py b/python/sglang/multimodal_gen/test/unit/test_lora_pipeline.py new file mode 100644 index 000000000..e4e14d402 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_lora_pipeline.py @@ -0,0 +1,121 @@ +from collections import defaultdict +from contextlib import contextmanager, nullcontext +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from sglang.multimodal_gen.runtime.layers.lora.linear import BaseLayerWithLoRA +from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline + +_RANK_PATCH = "sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline.dist.get_rank" + + +class _TestLoRAPipeline(LoRAPipeline): + def create_pipeline_stages(self, server_args): + return None + + +def _make_layer() -> BaseLayerWithLoRA: + return BaseLayerWithLoRA(torch.nn.Linear(2, 2, bias=False)) + + +def _make_pipeline(layer: BaseLayerWithLoRA) -> _TestLoRAPipeline: + pipeline = object.__new__(_TestLoRAPipeline) + pipeline.modules = {"transformer": torch.nn.Module()} + pipeline.server_args = SimpleNamespace(lora_merge_mode="dynamic") + pipeline.lora_initialized = True + pipeline.lora_adapters = defaultdict(dict) + pipeline.loaded_adapter_paths = {"adapter": "/adapter"} + pipeline.cur_adapter_name = {} + pipeline.cur_adapter_path = {} + pipeline.cur_adapter_strength = {} + pipeline.cur_adapter_config = {} + pipeline.lora_layers = {"linear": layer} + pipeline.lora_layers_transformer_2 = {} + pipeline.lora_layers_critic = {} + pipeline.is_lora_merged = {} + + pipeline.lora_adapters["adapter"]["linear.lora_A"] = torch.ones(1, 2) + pipeline.lora_adapters["adapter"]["linear.lora_B"] = torch.ones(2, 1) + return pipeline + + +def test_dynamic_lora_reactivates_cached_layers_without_weight_update_context(): + layer = _make_layer() + pipeline = _make_pipeline(layer) + context_calls = 0 + + @contextmanager + def counted_context(*args, **kwargs): + nonlocal context_calls + context_calls += 1 + yield [] + + pipeline._temporarily_disable_offload = counted_context + + with patch(_RANK_PATCH, return_value=0): + pipeline.set_lora( + "adapter", + "/adapter", + target="transformer", + strength=0.75, + merge_mode="dynamic", + ) + + first_lora_a = layer.lora_A + first_lora_b = layer.lora_B + assert context_calls == 0 + assert not layer.disable_lora + + pipeline._temporarily_disable_offload = lambda *args, **kwargs: nullcontext([]) + pipeline.deactivate_lora_weights("transformer") + assert layer.disable_lora + + def fail_apply(*args, **kwargs): + raise AssertionError("cached dynamic LoRA should not rebuild weights") + + context_calls = 0 + pipeline._temporarily_disable_offload = counted_context + pipeline._apply_lora_to_layers = fail_apply + + with patch(_RANK_PATCH, return_value=0): + pipeline.set_lora( + "adapter", + None, + target="transformer", + strength=0.75, + merge_mode="dynamic", + ) + + assert context_calls == 0 + assert not layer.disable_lora + assert layer.lora_A is first_lora_a + assert layer.lora_B is first_lora_b + + +def test_merged_lora_still_uses_weight_update_context(): + layer = _make_layer() + pipeline = _make_pipeline(layer) + context_calls = 0 + + @contextmanager + def counted_context(*args, **kwargs): + nonlocal context_calls + context_calls += 1 + yield [] + + pipeline._temporarily_disable_offload = counted_context + + with patch(_RANK_PATCH, return_value=0): + pipeline.set_lora( + "adapter", + "/adapter", + target="transformer", + strength=1.0, + merge_mode="merge", + ) + + assert context_calls == 1 + assert layer.merged + assert pipeline.is_lora_merged["transformer"]