[diffusion] feat: introduce ltx-2-two-stage device manager (#22869)

This commit is contained in:
Mick
2026-04-18 11:04:33 +08:00
committed by GitHub
parent 615d6c93b2
commit 0d94c3366a
8 changed files with 611 additions and 7 deletions
@@ -9,9 +9,11 @@ from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
is_ltx23_native_variant,
sync_ltx23_runtime_vae_markers,
)
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
PipelineComponentLoader,
)
from sglang.multimodal_gen.runtime.loader.utils import BYTES_PER_GB
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
@@ -31,7 +33,12 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages import (
TextEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.server_args import (
LTX2_RESIDENT_AUTO_ENABLE_MEM_GB,
ServerArgs,
)
from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
@@ -270,10 +277,404 @@ class LTX2Pipeline(_BaseLTX2Pipeline):
_add_ltx2_decoding_stage(self)
class LTX2TwoStageDeviceManager:
"""
Device residency manager for LTX-2.3 two-stage DiT switching.
Modes:
- resident: keep both DiTs on GPU; phase switch is pointer rebinding only.
- snapshot: keep CPU snapshots and prefetch the target DiT (DiT2 with pre-merged LoRA) with async H2D (similar to dit layerwise offload).
The DiT_1 will always be kept a replica in CPU.
- default snapshot behavior: allow stage1/stage2 overlap by prefetching
stage2 while stage1 is still running.
- snapshot low-VRAM behavior (`_snapshot_low_vram_mode=True`): evict
stage1 before stage2 prefetch and disable early overlap prefetch to
reduce peak VRAM, at the cost of higher phase-switch latency.
- default toggle: low-VRAM auto-enables on H100-like (<130 GiB) CUDA
GPUs, and stays disabled by default on higher-memory GPUs. It can be
overridden with `SGLANG_LTX2_SNAPSHOT_LOW_VRAM_MODE`.
- original: official two-stage semantics without premerged stage-2.
"""
VALID_MODES = ("original", "snapshot", "resident")
def __init__(self, pipeline: "LTX2TwoStagePipeline", server_args: ServerArgs):
self.pipeline = pipeline
self.server_args = server_args
self.mode = self._resolve_mode(server_args)
self._cpu_param_snapshots: dict[str, dict[str, torch.Tensor]] = {}
self._cpu_buffer_snapshots: dict[str, dict[str, torch.Tensor]] = {}
self._active_phase: str | None = None
self._prefetch_stream: object | None = None
self._phase_ready_events: dict[str, object] = {}
self._snapshot_low_vram_mode = self._resolve_snapshot_low_vram_mode()
self._snapshot_release_empty_cache = get_bool_env_var(
"SGLANG_LTX2_SNAPSHOT_RELEASE_EMPTY_CACHE",
default="false",
)
def _resolve_snapshot_low_vram_mode(self) -> bool:
if self.mode != "snapshot" or not current_platform.is_cuda():
return False
device_name = str(current_platform.get_device_name(0)).upper()
device_total_memory_gb = (
current_platform.get_device_total_memory() / BYTES_PER_GB
)
# H100-class (<130 GiB) cards are sensitive to stage1/stage2 overlap windows.
h100_like_memory_class = (
"H100" in device_name
or device_total_memory_gb < LTX2_RESIDENT_AUTO_ENABLE_MEM_GB
)
default = "true" if h100_like_memory_class else "false"
enabled = get_bool_env_var(
"SGLANG_LTX2_SNAPSHOT_LOW_VRAM_MODE",
default=default,
)
if enabled:
logger.info(
"Enabled LTX2 snapshot low-VRAM mode "
"(SGLANG_LTX2_SNAPSHOT_LOW_VRAM_MODE=%s, device=%s, %.2f GiB total)",
os.getenv("SGLANG_LTX2_SNAPSHOT_LOW_VRAM_MODE", default),
device_name,
device_total_memory_gb,
)
return enabled
@classmethod
def _resolve_mode(cls, server_args: ServerArgs) -> str:
mode = getattr(server_args, "ltx2_two_stage_device_mode", None)
if mode is None:
env_mode = os.getenv("SGLANG_LTX2_TWO_STAGE_DEVICE_MODE")
mode = env_mode.lower() if env_mode else "snapshot"
if mode not in cls.VALID_MODES:
raise ValueError(
f"Invalid ltx2_two_stage_device_mode={mode!r}. "
f"Expected one of {cls.VALID_MODES}."
)
return mode
@property
def should_use_premerged(self) -> bool:
"""Whether to keep a pre-merged stage-2 DiT for LTX-2.3 two-stage.
We only enable this optimization for native LTX-2.3 two-stage and when
users did not explicitly provide a stage-1 LoRA path
"""
return (
self.mode != "original"
and self.pipeline._should_merge_stage2_distilled_lora(self.server_args)
and getattr(self.pipeline, "_stage1_lora_path", None) is None
)
def initialize(self) -> None:
if not self.should_use_premerged:
return
self.pipeline._initialize_premerged_stage2_transformer(self.server_args)
if self.mode == "snapshot":
# Snapshot mode keeps both DiT CPU snapshots for cheap GPU release
# and re-hydrates stage-2 with async H2D when stage-1 finishes.
self._capture_module_cpu_snapshot("transformer")
self._capture_module_cpu_snapshot("transformer_2")
self._pin_stage1_transformer_if_beneficial()
elif self.mode == "resident":
self._ensure_on_gpu("transformer")
self._ensure_on_gpu("transformer_2")
logger.info(
"Using resident LTX-2.3 two-stage transformers mode (both DiTs stay on GPU)"
)
self._active_phase = "stage1"
self._sync_refinement_stage_transformer("stage1")
self._record_phase_ready_event("stage1")
def switch_phase(self, phase: str) -> bool:
"""Switch active two-stage DiT with minimal transfer/sync overhead."""
if not self.should_use_premerged:
return False
if phase == self._active_phase:
return True
if self.mode == "resident":
self._sync_refinement_stage_transformer(phase)
self._active_phase = phase
return True
if self.server_args.dit_cpu_offload:
target_name = "transformer_2" if phase == "stage2" else "transformer"
target_module = self.pipeline.get_module(target_name)
if self.mode == "snapshot" and self._snapshot_low_vram_mode:
# Trade a bit of phase-switch latency for lower peak VRAM:
# evict stage-1 before stage-2 H2D.
if phase == "stage2" and phase not in self._phase_ready_events:
stage1_module = self.pipeline.get_module("transformer")
stage1_param = (
next(stage1_module.parameters(), None)
if stage1_module is not None
else None
)
if stage1_param is not None and stage1_param.device.type == "cuda":
self._release_module_to_cpu_snapshot("transformer")
if phase not in self._phase_ready_events:
if self._module_is_on_gpu(target_module):
self._record_phase_ready_event(phase)
else:
self._schedule_phase_prefetch(phase, target_module)
# Stage-2 is only consumed after stage-1 denoising + upsample.
# Kick off the H2D early in stage-1 to overlap transfer with compute.
if (
phase == "stage1"
and "stage2" not in self._phase_ready_events
and not self._snapshot_low_vram_mode
):
self._schedule_phase_prefetch(
"stage2", self.pipeline.get_module("transformer_2")
)
else:
self._record_phase_ready_event(phase)
self._sync_refinement_stage_transformer(phase)
self._active_phase = phase
return True
def prefetch_stage2_after_stage1(self) -> None:
"""Kick off stage-2 H2D right after stage-1 denoising to hide switch latency."""
if (
not self.should_use_premerged
or self.mode != "snapshot"
or not self.server_args.dit_cpu_offload
):
return
if "stage2" in self._phase_ready_events:
return
if self._snapshot_low_vram_mode:
stage1_module = self.pipeline.get_module("transformer")
stage1_param = (
next(stage1_module.parameters(), None)
if stage1_module is not None
else None
)
if stage1_param is not None and stage1_param.device.type == "cuda":
self._release_module_to_cpu_snapshot("transformer")
self._schedule_phase_prefetch(
"stage2", self.pipeline.get_module("transformer_2")
)
def ensure_phase_ready(self, phase: str | None) -> None:
if not self.should_use_premerged or phase not in ("stage1", "stage2"):
return
if self.mode == "resident":
return
ready_event = self._phase_ready_events.get(phase)
if ready_event is None or not current_platform.is_cuda():
return
torch.get_device_module().current_stream().wait_event(ready_event)
def release_premerged_transformers(self) -> None:
if not self.should_use_premerged or self.mode != "snapshot":
return
# Keep stage-1 resident across requests so the next request can start
# denoising immediately while stage-2 is prefetched in the background.
for module_name in ("transformer_2",):
module = self.pipeline.get_module(module_name)
param = next(module.parameters(), None) if module is not None else None
if param is not None and param.device.type == "cuda":
self._release_module_to_cpu_snapshot(module_name)
if (
self._snapshot_release_empty_cache
and torch.get_device_module().is_available()
):
torch.get_device_module().empty_cache()
self._record_phase_ready_event("stage1")
@staticmethod
def _clone_cpu_tensor_snapshot(
tensor: torch.Tensor, *, pin_memory: bool
) -> torch.Tensor:
snapshot = tensor.detach()
if snapshot.device.type == "cpu":
if pin_memory and not snapshot.is_pinned():
return snapshot.pin_memory()
return snapshot
cpu_tensor = snapshot.to("cpu")
if pin_memory:
return cpu_tensor.pin_memory()
return cpu_tensor
def _capture_module_cpu_snapshot(self, module_name: str) -> None:
if module_name in self._cpu_param_snapshots:
return
module = self.pipeline.get_module(module_name)
if module is None:
raise ValueError(f"Module {module_name} is not available.")
pin_memory = bool(
self.server_args.pin_cpu_memory and torch.get_device_module().is_available()
)
self._cpu_param_snapshots[module_name] = {
name: self._clone_cpu_tensor_snapshot(param.data, pin_memory=pin_memory)
for name, param in module.named_parameters()
}
self._cpu_buffer_snapshots[module_name] = {
name: self._clone_cpu_tensor_snapshot(buffer.data, pin_memory=pin_memory)
for name, buffer in module.named_buffers()
}
def _release_module_to_cpu_snapshot(self, module_name: str) -> None:
"""Replace module tensors with cached CPU snapshots to avoid D2H copies.
This does not call `module.to("cpu")`. Instead, parameter and buffer storages
are rebound to pre-captured CPU tensors so CUDA storages can be released by
the allocator without an explicit D2H transfer.
"""
module = self.pipeline.get_module(module_name)
if module is None:
return
param_snapshots = self._cpu_param_snapshots.get(module_name)
buffer_snapshots = self._cpu_buffer_snapshots.get(module_name)
if param_snapshots is None or buffer_snapshots is None:
module.to("cpu")
return
for name, param in module.named_parameters():
snapshot = param_snapshots.get(name)
if snapshot is None:
raise KeyError(
f"Missing CPU parameter snapshot for {module_name}.{name}"
)
param.data = snapshot
for name, buffer in module.named_buffers():
snapshot = buffer_snapshots.get(name)
if snapshot is None:
raise KeyError(f"Missing CPU buffer snapshot for {module_name}.{name}")
# Preserve runtime-updated buffers (e.g., lazily built caches) when
# releasing back to CPU snapshots.
if buffer.device.type == "cuda":
snapshot.copy_(buffer.detach().to(device="cpu", dtype=snapshot.dtype))
elif buffer.device.type == "cpu":
snapshot.copy_(buffer.detach().to(dtype=snapshot.dtype))
buffer.data = snapshot
phase = "stage2" if module_name == "transformer_2" else "stage1"
self._phase_ready_events.pop(phase, None)
def _ensure_on_gpu(self, module_name: str) -> None:
module = self.pipeline.get_module(module_name)
if module is None:
return
param = next(module.parameters(), None)
if param is not None and param.device.type == "cpu":
module.to(get_local_torch_device(), non_blocking=True)
@staticmethod
def _module_is_on_gpu(module: torch.nn.Module | None) -> bool:
if module is None:
return False
param = next(module.parameters(), None)
return param is not None and param.device.type == "cuda"
def _supports_async_phase_prefetch(self) -> bool:
return (
self.mode == "snapshot"
and self.server_args.dit_cpu_offload
and current_platform.is_cuda()
)
def _get_prefetch_stream(self):
if not self._supports_async_phase_prefetch():
return None
if self._prefetch_stream is None:
self._prefetch_stream = torch.get_device_module().Stream(
device=get_local_torch_device()
)
return self._prefetch_stream
def _record_phase_ready_event(self, phase: str) -> None:
if not current_platform.is_cuda():
self._phase_ready_events.pop(phase, None)
return
module_name = "transformer_2" if phase == "stage2" else "transformer"
module = self.pipeline.get_module(module_name)
if not self._module_is_on_gpu(module):
self._phase_ready_events.pop(phase, None)
return
event = torch.get_device_module().Event()
event.record(torch.get_device_module().current_stream())
self._phase_ready_events[phase] = event
def _schedule_phase_prefetch(
self, phase: str, module: torch.nn.Module | None
) -> None:
if module is None:
self._phase_ready_events.pop(phase, None)
return
prefetch_stream = self._get_prefetch_stream()
if prefetch_stream is None:
module.to(get_local_torch_device(), non_blocking=True)
self._record_phase_ready_event(phase)
return
with torch.get_device_module().stream(prefetch_stream):
module.to(get_local_torch_device(), non_blocking=True)
event = torch.get_device_module().Event()
event.record(prefetch_stream)
self._phase_ready_events[phase] = event
def _pin_stage1_transformer_if_beneficial(self) -> None:
"""Optionally pin stage-1 DiT on GPU to remove first-stage cold H2D stall.
We only do this on high-VRAM CUDA machines with CPU offload enabled and
without FSDP inference. It trades extra steady-state VRAM for lower
request latency before the first denoise step.
"""
if (
not self.server_args.dit_cpu_offload
or self.server_args.use_fsdp_inference
or not current_platform.is_cuda()
or current_platform.get_device_total_memory() / BYTES_PER_GB < 70
):
return
transformer = self.pipeline.get_module("transformer")
param = (
next(transformer.parameters(), None) if transformer is not None else None
)
if transformer is not None and param is not None and param.device.type == "cpu":
transformer.to(get_local_torch_device(), non_blocking=True)
logger.info(
"Pinned stage1 transformer on GPU for LTX-2.3 two-stage startup"
)
self._active_phase = "stage1"
def _sync_refinement_stage_transformer(self, phase: str) -> None:
"""Keep stage-2 refinement bound to the expected DiT for current phase."""
refinement_stage = self.pipeline.get_stage("LTX2RefinementStage")
if refinement_stage is None:
return
target_name = "transformer_2" if phase == "stage2" else "transformer"
target_transformer = self.pipeline.get_module(target_name)
if target_transformer is not None:
refinement_stage.transformer = target_transformer
class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
pipeline_name = "LTX2TwoStagePipeline"
STAGE_2_DISTILLED_SIGMA_VALUES = [0.909375, 0.725, 0.421875, 0.0]
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self._device_manager = LTX2TwoStageDeviceManager(self, self.server_args)
self._use_premerged_stage2_transformer = (
self._device_manager.should_use_premerged
)
self._device_manager.initialize()
@staticmethod
def _should_merge_stage2_distilled_lora(server_args: ServerArgs) -> bool:
return is_ltx23_native_variant(
@@ -311,11 +712,55 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
self._stage1_lora_path = server_args.lora_path
self._stage1_lora_scale = float(server_args.lora_scale)
self._active_lora_phase = None
self._use_premerged_stage2_transformer = False
def _initialize_premerged_stage2_transformer(self, server_args: ServerArgs) -> None:
transformer_path = self._resolve_component_path(
server_args, "transformer", "transformer"
)
module, memory_usage = PipelineComponentLoader.load_component(
component_name="transformer_2",
component_model_path=transformer_path,
transformers_or_diffusers="diffusers",
server_args=server_args,
)
self.modules["transformer_2"] = module
self.memory_usages["transformer_2"] = memory_usage
# Reuse the canonical LoRA path used by legacy switching to reduce
# precision drift between snapshot mode and origin/main behavior.
self.set_lora(
lora_nickname="ltx2_stage2_distilled",
lora_path=self._distilled_lora_path,
target="transformer_2",
strength=1.0,
merge_weights=True,
)
def release_premerged_transformers_to_cpu_snapshots(self) -> None:
"""Release inactive premerged DiTs according to the selected device mode."""
self._device_manager.release_premerged_transformers()
def ensure_ltx2_phase_ready(self, phase: str | None) -> None:
self._device_manager.ensure_phase_ready(phase)
def prefetch_ltx2_stage2_after_stage1(self) -> None:
self._device_manager.prefetch_stage2_after_stage1()
def should_skip_ltx2_lora_switch_stage(self) -> bool:
return self._use_premerged_stage2_transformer and self._device_manager.mode in (
"snapshot",
"resident",
)
def switch_lora_phase(self, phase: str) -> None:
if phase == self._active_lora_phase:
return
if self._device_manager.switch_phase(phase):
self._active_lora_phase = phase
return
if phase == "stage1":
if self._stage1_lora_path:
self.set_lora(
@@ -373,6 +818,7 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
spatial_upsampler=self.get_module("spatial_upsampler"),
vae=self.get_module("vae"),
audio_vae=self.get_module("audio_vae"),
pipeline=self,
),
(
LTX2LoRASwitchStage(pipeline=self, phase="stage2"),
@@ -390,6 +836,7 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
distilled_sigmas=self.STAGE_2_DISTILLED_SIGMA_VALUES,
vae=self.get_module("vae"),
audio_vae=self.get_module("audio_vae"),
pipeline=self,
),
]
)
@@ -69,6 +69,7 @@ class LoRAPipeline(ComposedPipelineBase):
def __init__(self, *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
# Initialize all mutable instance attributes to avoid sharing across instances
self.lora_adapters = defaultdict(dict)
self.loaded_adapter_paths = {}
@@ -82,6 +82,32 @@ class LTX2AVDenoisingStage(LTX2DenoisingStage):
batch.latents = latents
batch.audio_latents = audio_latents
pipeline = self.pipeline() if self.pipeline else None
current_phase = (
str(getattr(batch, "extra", {}).get("ltx2_phase", ""))
if hasattr(batch, "extra")
else ""
)
if (
pipeline is not None
and getattr(pipeline, "_use_premerged_stage2_transformer", False)
and server_args.dit_cpu_offload
and not server_args.use_fsdp_inference
and current_phase == "stage2"
):
release_to_snapshots = getattr(
pipeline, "release_premerged_transformers_to_cpu_snapshots", None
)
if callable(release_to_snapshots):
release_to_snapshots()
else:
for dit in filter(None, [self.transformer]):
param = next(dit.parameters(), None)
if param is not None and param.device.type == "cuda":
dit.to("cpu")
if torch.get_device_module().is_available():
torch.get_device_module().empty_cache()
if isinstance(self.transformer, OffloadableDiTMixin):
for manager in self.transformer.layerwise_offload_managers:
manager.release_all()
@@ -91,9 +117,15 @@ class LTX2RefinementStage(LTX2AVDenoisingStage):
"""Stage-2 refinement wrapper that re-noises distilled LTX latents once."""
def __init__(
self, transformer, scheduler, distilled_sigmas, vae=None, audio_vae=None
self,
transformer,
scheduler,
distilled_sigmas,
vae=None,
audio_vae=None,
pipeline=None,
):
super().__init__(transformer, scheduler, vae, audio_vae)
super().__init__(transformer, scheduler, vae, audio_vae, pipeline=pipeline)
self.distilled_sigmas = torch.tensor(distilled_sigmas)
@staticmethod
@@ -415,6 +415,22 @@ class LTX2DenoisingStage(DenoisingStage):
self, ctx: LTX2DenoisingContext, batch: Req, server_args: ServerArgs
) -> None:
"""Reset the mirrored audio scheduler before the shared loop begins."""
if ctx.stage in ("stage1", "stage2"):
pipeline = self.pipeline() if self.pipeline else None
switch_lora_phase = (
getattr(pipeline, "switch_lora_phase", None)
if pipeline is not None
else None
)
if callable(switch_lora_phase):
switch_lora_phase(ctx.stage)
ensure_phase_ready = (
getattr(pipeline, "ensure_ltx2_phase_ready", None)
if pipeline is not None
else None
)
if callable(ensure_phase_ready):
ensure_phase_ready(ctx.stage)
super()._before_denoising_loop(ctx, batch, server_args)
if ctx.audio_scheduler is None:
raise ValueError("LTX-2 audio scheduler was not prepared.")
@@ -45,6 +45,12 @@ class LTX2LoRASwitchStage(PipelineStage):
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
switch_fn = getattr(self.pipeline, "switch_lora_phase", None)
should_skip_switch_stage = getattr(
self.pipeline, "should_skip_ltx2_lora_switch_stage", None
)
if callable(should_skip_switch_stage) and should_skip_switch_stage():
batch.extra["ltx2_phase"] = self.phase
return batch
if not callable(switch_fn):
raise ValueError(
"LTX2LoRASwitchStage requires pipeline.switch_lora_phase()"
@@ -57,11 +63,12 @@ class LTX2LoRASwitchStage(PipelineStage):
class LTX2UpsampleStage(PipelineStage):
"""Upsample Stage-1 video latents and prepare Stage-2 inputs."""
def __init__(self, spatial_upsampler, vae, audio_vae=None):
def __init__(self, spatial_upsampler, vae, audio_vae=None, pipeline=None):
super().__init__()
self.spatial_upsampler = spatial_upsampler
self.vae = vae
self.audio_vae = audio_vae
self.pipeline = pipeline
def _upsample_video_latents(
self, latents: torch.Tensor, server_args: ServerArgs, device: torch.device
@@ -111,6 +118,14 @@ class LTX2UpsampleStage(PipelineStage):
)
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
prefetch_stage2 = (
getattr(self.pipeline, "prefetch_ltx2_stage2_after_stage1", None)
if self.pipeline is not None
else None
)
if callable(prefetch_stage2):
prefetch_stage2()
device = get_local_torch_device()
latents = self._upsample_video_latents(batch.latents, server_args, device)
logger.info("Upsampled video latents: %s", list(latents.shape))
@@ -22,6 +22,9 @@ import yaml
from sglang.multimodal_gen import envs
from sglang.multimodal_gen.configs.models.encoders import T5Config
from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
is_ltx23_native_variant,
)
from sglang.multimodal_gen.configs.quantization.nunchaku import NunchakuSVDQuantArgs
from sglang.multimodal_gen.runtime.disaggregation.disagg_args import (
DisaggArgsMixin,
@@ -69,6 +72,16 @@ logger = init_logger(__name__)
# our validated Wan/MOVA workloads, so use a 130 GiB cutoff to keep H200-class
# GPUs on the faster no-offload default while preserving some headroom.
WAN_LAYERWISE_OFFLOAD_AUTO_DISABLE_MEM_GB = 130
LTX2_TWO_STAGE_DEVICE_MODES = ("original", "snapshot", "resident")
# H200-class GPUs (>=130 GiB total) can usually keep both LTX2 DiTs resident.
LTX2_RESIDENT_AUTO_ENABLE_MEM_GB = 130
def _normalize_ltx2_two_stage_device_mode(mode: str | None) -> str | None:
if mode is None:
return None
mode = mode.lower()
return mode
class Backend(str, Enum):
@@ -175,6 +188,7 @@ class ServerArgs(DisaggArgsMixin):
vae_cpu_offload: bool | None = None
use_fsdp_inference: bool = False
pin_cpu_memory: bool = True
ltx2_two_stage_device_mode: str | None = None
# ComfyUI integration
comfyui_mode: bool = False
@@ -283,6 +297,7 @@ class ServerArgs(DisaggArgsMixin):
def _adjust_parameters(self):
"""set defaults and normalize values."""
self._adjust_offload()
self._adjust_ltx2_two_stage_device_mode()
self._adjust_path()
self._adjust_quant_config()
self._adjust_warmup()
@@ -378,6 +393,62 @@ class ServerArgs(DisaggArgsMixin):
if self.vae_cpu_offload is None:
self.vae_cpu_offload = True
def _adjust_ltx2_two_stage_device_mode(self):
is_ltx23_two_stage = self.pipeline_class_name == "LTX2TwoStagePipeline" and (
self._is_ltx23_model_path(self.model_path)
or is_ltx23_native_variant(self.pipeline_config.vae_config.arch_config)
)
if not is_ltx23_two_stage:
return
mode = self.ltx2_two_stage_device_mode
if mode is None:
env_mode = os.getenv("SGLANG_LTX2_TWO_STAGE_DEVICE_MODE")
mode = (
_normalize_ltx2_two_stage_device_mode(env_mode)
if env_mode
else self._resolve_default_ltx2_two_stage_device_mode()
)
else:
mode = _normalize_ltx2_two_stage_device_mode(mode)
if mode not in LTX2_TWO_STAGE_DEVICE_MODES:
raise ValueError(
f"Invalid ltx2_two_stage_device_mode={mode!r}. "
f"Expected one of {LTX2_TWO_STAGE_DEVICE_MODES}."
)
self.ltx2_two_stage_device_mode = mode
def _resolve_default_ltx2_two_stage_device_mode(self) -> str:
if not current_platform.is_cuda():
logger.info(
"Automatically set ltx2_two_stage_device_mode=snapshot on non-CUDA platform"
)
return "snapshot"
device_name = str(current_platform.get_device_name(0)).upper()
device_total_memory_gb = (
current_platform.get_device_total_memory() / BYTES_PER_GB
)
if (
"H200" in device_name
or device_total_memory_gb >= LTX2_RESIDENT_AUTO_ENABLE_MEM_GB
):
logger.info(
"Automatically set ltx2_two_stage_device_mode=resident for high-memory CUDA GPU (%s, %.2f GiB total)",
device_name,
device_total_memory_gb,
)
return "resident"
logger.info(
"Automatically set ltx2_two_stage_device_mode=snapshot for CUDA GPU (%s, %.2f GiB total)",
device_name,
device_total_memory_gb,
)
return "snapshot"
def _adjust_attention_backend(self):
if self.attention_backend in ["fa3", "fa4"]:
self.attention_backend = "fa"
@@ -894,6 +965,19 @@ class ServerArgs(DisaggArgsMixin):
help='Pin memory for CPU offload. Only added as a temp workaround if it throws "CUDA error: invalid argument". '
"Should be enabled in almost all cases",
)
parser.add_argument(
"--ltx2-two-stage-device-mode",
type=str,
choices=LTX2_TWO_STAGE_DEVICE_MODES,
default=ServerArgs.ltx2_two_stage_device_mode,
help=(
"LTX-2.3 two-stage device residency mode: "
"'original' keeps official two-stage semantics without premerged stage2, "
"'snapshot' keeps premerged stage2 with snapshot-based release, "
"'resident' keeps both transformers resident on GPU. "
"Default is auto: resident on H200/high-memory CUDA GPUs, otherwise snapshot."
),
)
parser.add_argument(
"--disable-autocast",
action=StoreBoolean,
@@ -522,7 +522,9 @@ TWO_GPU_CASES_A = [
"ltx_2_3_two_stage_ti2v_2gpus",
DiffusionServerArgs(
model_path="Lightricks/LTX-2.3",
extras=["--pipeline-class-name LTX2TwoStagePipeline"],
extras=[
"--pipeline-class-name LTX2TwoStagePipeline --ltx2-two-stage-device-mode original"
],
),
TI2V_sampling_params,
),
@@ -541,7 +543,10 @@ TWO_GPU_CASES_B = [
"ltx_2.3_two_stage_t2v_2gpus",
DiffusionServerArgs(
model_path="Lightricks/LTX-2.3",
extras=["--pipeline-class-name LTX2TwoStagePipeline"],
extras=[
"--pipeline-class-name LTX2TwoStagePipeline",
"--ltx2-two-stage-device-mode original",
],
),
T2V_sampling_params,
),