[diffusion] perf: merge LTX-2 stage-1 distilled LoRA into the base in original mode (#28594)
This commit is contained in:
@@ -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
|
||||
Reference in New Issue
Block a user