From af2ec2a0ddbd0249ba8e8a755765fa685a79a426 Mon Sep 17 00:00:00 2001 From: Mick Date: Fri, 19 Jun 2026 15:41:47 +0800 Subject: [PATCH] [diffusion] perf: merge LTX-2 stage-1 distilled LoRA into the base in original mode (#28594) --- .../runtime/layers/lora/linear.py | 24 ++++ .../runtime/pipelines/ltx_2_pipeline.py | 134 ++++++++++++++++++ .../test/server/perf_baselines.json | 6 +- .../test/unit/test_lora_commit_as_base.py | 103 ++++++++++++++ 4 files changed, 264 insertions(+), 3 deletions(-) create mode 100644 python/sglang/multimodal_gen/test/unit/test_lora_commit_as_base.py diff --git a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py index 094a62a6f..ba3ca1571 100644 --- a/python/sglang/multimodal_gen/runtime/layers/lora/linear.py +++ b/python/sglang/multimodal_gen/runtime/layers/lora/linear.py @@ -403,6 +403,30 @@ class BaseLayerWithLoRA(nn.Module): self.merged = False + @torch.no_grad() + def commit_merged_as_base(self) -> None: + """Promote the currently merged weights to the permanent base. + + Re-snapshots ``cpu_weight`` so the merged weights become the restore + target and resets adapter bookkeeping (``merged=False``). A later dynamic + ``set_lora_weights`` then adds its delta on top of the merged base instead + of unmerging it. + """ + if not self.merged: + return + weight = self.base_layer.weight + if isinstance(weight, DTensor): + weight = weight.to_local() + # clone(): to("cpu") may alias storage; we must not mutate this backup. + self.cpu_weight = weight.detach().to("cpu").clone() + self.merged = False + self.disable_lora = True + self.lora_weights_list = [] + self.lora_A = None + self.lora_B = None + self.lora_path = None + self.strength = 1.0 + class VocabParallelEmbeddingWithLoRA(BaseLayerWithLoRA): """ diff --git a/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py index 5acd524bb..ca5ca703a 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py @@ -494,6 +494,11 @@ class LTX2TwoStageResidencyController: ) def initialize(self) -> None: + if self.mode == "original": + # maybe merge the fixed stage-1 distilled LoRA into the base once so phase switches skip per-request + # merge/unmerge. + self.pipeline._maybe_merge_stage1_distilled_into_base(self.server_args) + return if not self.should_use_premerged: return self.pipeline._initialize_premerged_stage2_transformer(self.server_args) @@ -587,6 +592,10 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline): self._active_lora_phase = None self._active_lora_signature = None self._use_premerged_stage2_transformer = False + # set when original mode merges stage-1 distilled LoRA into the DiT base + # once at init (see _merge_stage1_distilled_into_base). + self._stage1_distilled_in_base = False + self._stage1_distilled_base_strength: float | None = None def _initialize_premerged_stage2_transformer(self, server_args: ServerArgs) -> None: transformer_path = self._resolve_component_path( @@ -611,6 +620,123 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline): merge_weights=True, ) + def _can_merge_stage1_distilled_into_base(self, server_args: ServerArgs) -> bool: + """Whether original mode can merge stage-1 distilled LoRA into the base once. + + For a fixed non-zero stage-1 strength (HQ only), we merge it into the base once and run + stage 2 as a dynamic delta. Requires native LTX-2.3, no user stage-1 + LoRA, plain (non-FSDP/DTensor, unquantized) weights. + """ + return ( + self._ltx2_residency.mode == "original" + and self._should_merge_stage2_distilled_lora(server_args) + and self._stage1_lora_path is None + and float(self.STAGE_1_DISTILLED_LORA_STRENGTH) != 0.0 + and not bool(getattr(server_args, "use_fsdp_inference", False)) + and getattr(server_args, "quantization", None) is None + ) + + def _maybe_merge_stage1_distilled_into_base(self, server_args: ServerArgs) -> None: + """Merge stage-1 distilled LoRA into the single DiT base once at init. + + Stage 1 then runs on the base; stage 2 adds a dynamic delta of + ``stage2 - stage1`` strength on top. No per-request merge/unmerge. + """ + self._stage1_distilled_in_base = False + self._stage1_distilled_base_strength = None + if not self._can_merge_stage1_distilled_into_base(server_args): + return + + strength = float(self.STAGE_1_DISTILLED_LORA_STRENGTH) + # Canonical merge path (handles offload/TP), then commit it as the base. + self.set_lora( + lora_nickname="ltx2_stage1_distilled", + lora_path=self._distilled_lora_path, + target="transformer", + strength=strength, + merge_weights=True, + ) + if self._uses_dtensor_weights(self.lora_layers): + # Unsupported layout; undo and fall back to per-request merge. + self.deactivate_lora_weights(target="transformer") + return + + for layer in self.lora_layers.values(): + layer.commit_merged_as_base() + # Keep the adapter loaded for the stage-2 delta; clear merged bookkeeping. + self.is_lora_merged["transformer"] = False + self.cur_adapter_strength.pop("transformer", None) + self.cur_adapter_config.pop("transformer", None) + + self._stage1_distilled_in_base = True + self._stage1_distilled_base_strength = strength + self._active_lora_phase = "stage1" + self._active_lora_signature = None + logger.info( + "Merged LTX-2 stage-1 distilled LoRA (strength=%.4f) into the DiT base; " + "stage-2 uses a dynamic delta to avoid per-request merge/unmerge.", + strength, + ) + + def _unmerge_stage1_distilled_from_base(self) -> None: + """Restore the base weights and revert to per-request merging. + + Used when a request overrides the stage-1 strength away from the merged + value. Subtracts the merged delta, then disables the optimization. + """ + if not self._stage1_distilled_in_base: + return + self.set_lora( + lora_nickname="ltx2_stage1_distilled", + lora_path=self._distilled_lora_path, + target="transformer", + strength=-float(self._stage1_distilled_base_strength), + merge_weights=True, + ) + for layer in self.lora_layers.values(): + layer.commit_merged_as_base() + self.is_lora_merged["transformer"] = False + self.cur_adapter_strength.pop("transformer", None) + self.cur_adapter_config.pop("transformer", None) + self._stage1_distilled_in_base = False + self._stage1_distilled_base_strength = None + self._active_lora_signature = None + logger.info("Restored LTX-2 base; reverting to per-request stage-1 merge.") + + def _switch_lora_phase_base_merged( + self, phase: str, distilled_lora_strength: float + ) -> bool: + """Phase switch when stage-1 distilled is merged into the base, unmerge or apply dynamic lora + + Returns True if handled, False to fall back to the per-request path + (after restoring the base). + """ + if phase == "stage1": + if distilled_lora_strength != self._stage1_distilled_base_strength: + self._unmerge_stage1_distilled_from_base() + return False + # Base already holds stage-1 distilled; just drop the stage-2 delta. + self.deactivate_lora_weights(target="transformer") + return True + if phase == "stage2": + delta = distilled_lora_strength - float( + self._stage1_distilled_base_strength + ) + if delta == 0.0: + self.deactivate_lora_weights(target="transformer") + return True + # Dynamic delta on the merged base (base + delta == stage-2 strength); + # reuse the loaded adapter, so no reload/merge/unmerge. + self.set_lora( + lora_nickname="ltx2_stage1_distilled", + lora_path=self._distilled_lora_path, + target="transformer", + strength=delta, + merge_weights=False, + ) + return True + return False + def should_skip_ltx2_lora_switch_stage(self) -> bool: return ( self._use_premerged_stage2_transformer @@ -697,6 +823,14 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline): if phase_signature == self._active_lora_signature: return + if self._stage1_distilled_in_base: + if self._switch_lora_phase_base_merged(phase, distilled_lora_strength): + self._active_lora_phase = phase + self._active_lora_signature = phase_signature + return + # Base was restored (stage-1 strength override); fall through to the + # legacy per-request merge path below. + if self._ltx2_residency.enter_phase( phase ) and self._can_short_circuit_lora_switch(phase, batch): diff --git a/python/sglang/multimodal_gen/test/server/perf_baselines.json b/python/sglang/multimodal_gen/test/server/perf_baselines.json index 7d49fc138..a4b030dd7 100644 --- a/python/sglang/multimodal_gen/test/server/perf_baselines.json +++ b/python/sglang/multimodal_gen/test/server/perf_baselines.json @@ -2685,14 +2685,14 @@ "TextEncodingStage": 401.75, "LTX2TextConnectorStage": 27.48, "LTX2HalveResolutionStage": 0.04, - "LTX2LoRASwitchStage": 180.0, + "LTX2LoRASwitchStage": 6.12, "LTX2SigmaPreparationStage": 0.26, "TimestepPreparationStage": 14.45, "LTX2AVLatentPreparationStage": 0.13, "LTX2ImageEncodingStage": 57.62, "LTX2AVDenoisingStage": 12162.51, "LTX2UpsampleStage": 11.04, - "ltx2_lora_switch_stage2": 9155.58, + "ltx2_lora_switch_stage2": 5.68, "ltx2_image_encoding_stage2": 64.54, "LTX2RefinementStage": 3484.86, "LTX2AVDecodingStage": 1054.91, @@ -2718,7 +2718,7 @@ "16": 1144.87, "17": 1143.19 }, - "expected_e2e_ms": 26673.22, + "expected_e2e_ms": 15981.39, "expected_avg_denoise_ms": 868.06, "expected_median_denoise_ms": 747.62, "estimated_full_test_time_s": 363.2 diff --git a/python/sglang/multimodal_gen/test/unit/test_lora_commit_as_base.py b/python/sglang/multimodal_gen/test/unit/test_lora_commit_as_base.py new file mode 100644 index 000000000..0eac3bf91 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_lora_commit_as_base.py @@ -0,0 +1,103 @@ +"""Unit tests for BaseLayerWithLoRA.commit_merged_as_base. + +Validates the "merge a fixed-strength LoRA into the base once, then apply the +rest as a dynamic delta" primitive used by LTX-2 original mode. Weight-level +invariants are checked on CPU with a plain nn.Linear base (the parallel forward +path needs a distributed runtime and is covered by the diffusion server tests). +""" + +import torch +from torch import nn + +from sglang.multimodal_gen.runtime.layers.lora.linear import LinearWithLoRA + + +def _make_layer(base_weight: torch.Tensor, rank: int = 2, alpha: int = 2): + base = nn.Linear(base_weight.shape[1], base_weight.shape[0], bias=False) + with torch.no_grad(): + base.weight.copy_(base_weight) + return LinearWithLoRA(base, lora_rank=rank, lora_alpha=alpha) + + +def test_commit_merged_as_base_promotes_weights_and_resets_state(): + torch.manual_seed(0) + out_f, in_f, rank = 6, 8, 2 + base_w = torch.randn(out_f, in_f) + A = torch.randn(rank, in_f) + B = torch.randn(out_f, rank) + s1 = 0.25 # alpha == rank below, so scale == strength + + layer = _make_layer(base_w, rank=rank, alpha=rank) + layer.set_lora_weights( + A.clone(), B.clone(), strength=s1, clear_existing=True, merge_weights=True + ) + assert layer.merged + + layer.commit_merged_as_base() + + merged_w = base_w + s1 * (B @ A) + # Merged weights become the permanent base + restore target. + assert torch.allclose(layer.base_layer.weight, merged_w, atol=1e-5) + assert torch.allclose(layer.cpu_weight, merged_w, atol=1e-5) + # Bookkeeping reset so a later dynamic set_lora adds a delta on top, and + # deactivate (which only unmerges when merged=True) leaves the base intact. + assert layer.merged is False + assert layer.disable_lora is True + assert layer.lora_weights_list == [] + assert layer.lora_A is None and layer.lora_B is None + + +def test_dynamic_delta_after_commit_does_not_unmerge_base(): + torch.manual_seed(1) + out_f, in_f, rank = 4, 5, 2 + base_w = torch.randn(out_f, in_f) + A = torch.randn(rank, in_f) + B = torch.randn(out_f, rank) + s1, delta = 0.25, 0.25 + + layer = _make_layer(base_w, rank=rank, alpha=rank) + layer.set_lora_weights( + A.clone(), B.clone(), strength=s1, clear_existing=True, merge_weights=True + ) + layer.commit_merged_as_base() + merged_w = layer.base_layer.weight.detach().clone() + + layer.set_lora_weights( + A.clone(), B.clone(), strength=delta, clear_existing=True, merge_weights=False + ) + # Dynamic mode must NOT touch the merged base weights (no unmerge happened). + assert layer.merged is False + assert layer.disable_lora is False + assert torch.allclose(layer.base_layer.weight, merged_w, atol=1e-6) + assert len(layer.lora_weights_list) == 1 + assert layer.strength == delta + # Effective transform (base + dynamic delta) equals a single merge at s1+delta. + effective = layer.base_layer.weight + layer.strength * (B @ A) + assert torch.allclose(effective, base_w + (s1 + delta) * (B @ A), atol=1e-5) + + +def test_negative_merge_after_commit_restores_original_base(): + """Restore path: re-merging at the negative strength recovers the base.""" + torch.manual_seed(2) + out_f, in_f, rank = 5, 7, 3 + base_w = torch.randn(out_f, in_f) + A = torch.randn(rank, in_f) + B = torch.randn(out_f, rank) + s1 = 0.25 + + layer = _make_layer(base_w, rank=rank, alpha=rank) + layer.set_lora_weights( + A.clone(), B.clone(), strength=s1, clear_existing=True, merge_weights=True + ) + layer.commit_merged_as_base() + assert not torch.allclose(layer.base_layer.weight, base_w, atol=1e-5) + + # Subtract the merged delta (as _unmerge_stage1_distilled_from_base does). + layer.set_lora_weights( + A.clone(), B.clone(), strength=-s1, clear_existing=True, merge_weights=True + ) + layer.commit_merged_as_base() + + assert torch.allclose(layer.base_layer.weight, base_w, atol=1e-5) + assert torch.allclose(layer.cpu_weight, base_w, atol=1e-5) + assert layer.merged is False