[diffusion] perf: merge LTX-2 stage-1 distilled LoRA into the base in original mode (#28594)

This commit is contained in:
Mick
2026-06-19 15:41:47 +08:00
committed by GitHub
parent 59eb142ec2
commit af2ec2a0dd
4 changed files with 264 additions and 3 deletions
@@ -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):
"""
@@ -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):
@@ -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
@@ -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