[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
|
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):
|
class VocabParallelEmbeddingWithLoRA(BaseLayerWithLoRA):
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -494,6 +494,11 @@ class LTX2TwoStageResidencyController:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def initialize(self) -> None:
|
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:
|
if not self.should_use_premerged:
|
||||||
return
|
return
|
||||||
self.pipeline._initialize_premerged_stage2_transformer(self.server_args)
|
self.pipeline._initialize_premerged_stage2_transformer(self.server_args)
|
||||||
@@ -587,6 +592,10 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
|
|||||||
self._active_lora_phase = None
|
self._active_lora_phase = None
|
||||||
self._active_lora_signature = None
|
self._active_lora_signature = None
|
||||||
self._use_premerged_stage2_transformer = False
|
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:
|
def _initialize_premerged_stage2_transformer(self, server_args: ServerArgs) -> None:
|
||||||
transformer_path = self._resolve_component_path(
|
transformer_path = self._resolve_component_path(
|
||||||
@@ -611,6 +620,123 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
|
|||||||
merge_weights=True,
|
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:
|
def should_skip_ltx2_lora_switch_stage(self) -> bool:
|
||||||
return (
|
return (
|
||||||
self._use_premerged_stage2_transformer
|
self._use_premerged_stage2_transformer
|
||||||
@@ -697,6 +823,14 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
|
|||||||
if phase_signature == self._active_lora_signature:
|
if phase_signature == self._active_lora_signature:
|
||||||
return
|
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(
|
if self._ltx2_residency.enter_phase(
|
||||||
phase
|
phase
|
||||||
) and self._can_short_circuit_lora_switch(phase, batch):
|
) and self._can_short_circuit_lora_switch(phase, batch):
|
||||||
|
|||||||
@@ -2685,14 +2685,14 @@
|
|||||||
"TextEncodingStage": 401.75,
|
"TextEncodingStage": 401.75,
|
||||||
"LTX2TextConnectorStage": 27.48,
|
"LTX2TextConnectorStage": 27.48,
|
||||||
"LTX2HalveResolutionStage": 0.04,
|
"LTX2HalveResolutionStage": 0.04,
|
||||||
"LTX2LoRASwitchStage": 180.0,
|
"LTX2LoRASwitchStage": 6.12,
|
||||||
"LTX2SigmaPreparationStage": 0.26,
|
"LTX2SigmaPreparationStage": 0.26,
|
||||||
"TimestepPreparationStage": 14.45,
|
"TimestepPreparationStage": 14.45,
|
||||||
"LTX2AVLatentPreparationStage": 0.13,
|
"LTX2AVLatentPreparationStage": 0.13,
|
||||||
"LTX2ImageEncodingStage": 57.62,
|
"LTX2ImageEncodingStage": 57.62,
|
||||||
"LTX2AVDenoisingStage": 12162.51,
|
"LTX2AVDenoisingStage": 12162.51,
|
||||||
"LTX2UpsampleStage": 11.04,
|
"LTX2UpsampleStage": 11.04,
|
||||||
"ltx2_lora_switch_stage2": 9155.58,
|
"ltx2_lora_switch_stage2": 5.68,
|
||||||
"ltx2_image_encoding_stage2": 64.54,
|
"ltx2_image_encoding_stage2": 64.54,
|
||||||
"LTX2RefinementStage": 3484.86,
|
"LTX2RefinementStage": 3484.86,
|
||||||
"LTX2AVDecodingStage": 1054.91,
|
"LTX2AVDecodingStage": 1054.91,
|
||||||
@@ -2718,7 +2718,7 @@
|
|||||||
"16": 1144.87,
|
"16": 1144.87,
|
||||||
"17": 1143.19
|
"17": 1143.19
|
||||||
},
|
},
|
||||||
"expected_e2e_ms": 26673.22,
|
"expected_e2e_ms": 15981.39,
|
||||||
"expected_avg_denoise_ms": 868.06,
|
"expected_avg_denoise_ms": 868.06,
|
||||||
"expected_median_denoise_ms": 747.62,
|
"expected_median_denoise_ms": 747.62,
|
||||||
"estimated_full_test_time_s": 363.2
|
"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