[diffusion] feat: introduce ltx-2-two-stage device manager (#22869)
This commit is contained in:
@@ -46,6 +46,10 @@ default parameters when initializing and generating videos.
|
|||||||
2. SageSLA is based on SpargeAttn. Install it first with `pip install git+https://github.com/thu-ml/SpargeAttn.git --no-build-isolation`
|
2. SageSLA is based on SpargeAttn. Install it first with `pip install git+https://github.com/thu-ml/SpargeAttn.git --no-build-isolation`
|
||||||
3. LTX-2 and LTX-2.3 two-stage generation uses `--pipeline-class-name LTX2TwoStagePipeline`. The spatial upsampler and distilled LoRA are auto-resolved from the model snapshot by default, and can still be overridden with `--spatial-upsampler-path` and `--distilled-lora-path`.
|
3. LTX-2 and LTX-2.3 two-stage generation uses `--pipeline-class-name LTX2TwoStagePipeline`. The spatial upsampler and distilled LoRA are auto-resolved from the model snapshot by default, and can still be overridden with `--spatial-upsampler-path` and `--distilled-lora-path`.
|
||||||
- For LTX models, the `Resolutions` column uses output video `width×height` semantics, matching `sglang generate --width ... --height ...`.
|
- For LTX models, the `Resolutions` column uses output video `width×height` semantics, matching `sglang generate --width ... --height ...`.
|
||||||
|
4. LTX-2.3 two-stage also supports `--ltx2-two-stage-device-mode {legacy,snapshot,resident}`:
|
||||||
|
- `snapshot` is the default and recommended mode.
|
||||||
|
- `resident` usually provides the best latency/throughput but uses much more VRAM.
|
||||||
|
- `legacy` preserves the historical switching path for fallback/debug.
|
||||||
|
|
||||||
### Image Generation Models
|
### Image Generation Models
|
||||||
|
|
||||||
|
|||||||
@@ -9,9 +9,11 @@ from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
|
|||||||
is_ltx23_native_variant,
|
is_ltx23_native_variant,
|
||||||
sync_ltx23_runtime_vae_markers,
|
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 (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
|
||||||
PipelineComponentLoader,
|
PipelineComponentLoader,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.utils import BYTES_PER_GB
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||||
ComposedPipelineBase,
|
ComposedPipelineBase,
|
||||||
)
|
)
|
||||||
@@ -31,7 +33,12 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages import (
|
|||||||
TextEncodingStage,
|
TextEncodingStage,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
|
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
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
@@ -270,10 +277,404 @@ class LTX2Pipeline(_BaseLTX2Pipeline):
|
|||||||
_add_ltx2_decoding_stage(self)
|
_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):
|
class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
|
||||||
pipeline_name = "LTX2TwoStagePipeline"
|
pipeline_name = "LTX2TwoStagePipeline"
|
||||||
STAGE_2_DISTILLED_SIGMA_VALUES = [0.909375, 0.725, 0.421875, 0.0]
|
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
|
@staticmethod
|
||||||
def _should_merge_stage2_distilled_lora(server_args: ServerArgs) -> bool:
|
def _should_merge_stage2_distilled_lora(server_args: ServerArgs) -> bool:
|
||||||
return is_ltx23_native_variant(
|
return is_ltx23_native_variant(
|
||||||
@@ -311,11 +712,55 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
|
|||||||
self._stage1_lora_path = server_args.lora_path
|
self._stage1_lora_path = server_args.lora_path
|
||||||
self._stage1_lora_scale = float(server_args.lora_scale)
|
self._stage1_lora_scale = float(server_args.lora_scale)
|
||||||
self._active_lora_phase = None
|
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:
|
def switch_lora_phase(self, phase: str) -> None:
|
||||||
if phase == self._active_lora_phase:
|
if phase == self._active_lora_phase:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
if self._device_manager.switch_phase(phase):
|
||||||
|
self._active_lora_phase = phase
|
||||||
|
return
|
||||||
|
|
||||||
if phase == "stage1":
|
if phase == "stage1":
|
||||||
if self._stage1_lora_path:
|
if self._stage1_lora_path:
|
||||||
self.set_lora(
|
self.set_lora(
|
||||||
@@ -373,6 +818,7 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
|
|||||||
spatial_upsampler=self.get_module("spatial_upsampler"),
|
spatial_upsampler=self.get_module("spatial_upsampler"),
|
||||||
vae=self.get_module("vae"),
|
vae=self.get_module("vae"),
|
||||||
audio_vae=self.get_module("audio_vae"),
|
audio_vae=self.get_module("audio_vae"),
|
||||||
|
pipeline=self,
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
LTX2LoRASwitchStage(pipeline=self, phase="stage2"),
|
LTX2LoRASwitchStage(pipeline=self, phase="stage2"),
|
||||||
@@ -390,6 +836,7 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
|
|||||||
distilled_sigmas=self.STAGE_2_DISTILLED_SIGMA_VALUES,
|
distilled_sigmas=self.STAGE_2_DISTILLED_SIGMA_VALUES,
|
||||||
vae=self.get_module("vae"),
|
vae=self.get_module("vae"),
|
||||||
audio_vae=self.get_module("audio_vae"),
|
audio_vae=self.get_module("audio_vae"),
|
||||||
|
pipeline=self,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -69,6 +69,7 @@ class LoRAPipeline(ComposedPipelineBase):
|
|||||||
|
|
||||||
def __init__(self, *args, **kwargs) -> None:
|
def __init__(self, *args, **kwargs) -> None:
|
||||||
super().__init__(*args, **kwargs)
|
super().__init__(*args, **kwargs)
|
||||||
|
|
||||||
# Initialize all mutable instance attributes to avoid sharing across instances
|
# Initialize all mutable instance attributes to avoid sharing across instances
|
||||||
self.lora_adapters = defaultdict(dict)
|
self.lora_adapters = defaultdict(dict)
|
||||||
self.loaded_adapter_paths = {}
|
self.loaded_adapter_paths = {}
|
||||||
|
|||||||
@@ -82,6 +82,32 @@ class LTX2AVDenoisingStage(LTX2DenoisingStage):
|
|||||||
batch.latents = latents
|
batch.latents = latents
|
||||||
batch.audio_latents = audio_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):
|
if isinstance(self.transformer, OffloadableDiTMixin):
|
||||||
for manager in self.transformer.layerwise_offload_managers:
|
for manager in self.transformer.layerwise_offload_managers:
|
||||||
manager.release_all()
|
manager.release_all()
|
||||||
@@ -91,9 +117,15 @@ class LTX2RefinementStage(LTX2AVDenoisingStage):
|
|||||||
"""Stage-2 refinement wrapper that re-noises distilled LTX latents once."""
|
"""Stage-2 refinement wrapper that re-noises distilled LTX latents once."""
|
||||||
|
|
||||||
def __init__(
|
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)
|
self.distilled_sigmas = torch.tensor(distilled_sigmas)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -415,6 +415,22 @@ class LTX2DenoisingStage(DenoisingStage):
|
|||||||
self, ctx: LTX2DenoisingContext, batch: Req, server_args: ServerArgs
|
self, ctx: LTX2DenoisingContext, batch: Req, server_args: ServerArgs
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Reset the mirrored audio scheduler before the shared loop begins."""
|
"""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)
|
super()._before_denoising_loop(ctx, batch, server_args)
|
||||||
if ctx.audio_scheduler is None:
|
if ctx.audio_scheduler is None:
|
||||||
raise ValueError("LTX-2 audio scheduler was not prepared.")
|
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:
|
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
|
||||||
switch_fn = getattr(self.pipeline, "switch_lora_phase", None)
|
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):
|
if not callable(switch_fn):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"LTX2LoRASwitchStage requires pipeline.switch_lora_phase()"
|
"LTX2LoRASwitchStage requires pipeline.switch_lora_phase()"
|
||||||
@@ -57,11 +63,12 @@ class LTX2LoRASwitchStage(PipelineStage):
|
|||||||
class LTX2UpsampleStage(PipelineStage):
|
class LTX2UpsampleStage(PipelineStage):
|
||||||
"""Upsample Stage-1 video latents and prepare Stage-2 inputs."""
|
"""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__()
|
super().__init__()
|
||||||
self.spatial_upsampler = spatial_upsampler
|
self.spatial_upsampler = spatial_upsampler
|
||||||
self.vae = vae
|
self.vae = vae
|
||||||
self.audio_vae = audio_vae
|
self.audio_vae = audio_vae
|
||||||
|
self.pipeline = pipeline
|
||||||
|
|
||||||
def _upsample_video_latents(
|
def _upsample_video_latents(
|
||||||
self, latents: torch.Tensor, server_args: ServerArgs, device: torch.device
|
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:
|
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()
|
device = get_local_torch_device()
|
||||||
latents = self._upsample_video_latents(batch.latents, server_args, device)
|
latents = self._upsample_video_latents(batch.latents, server_args, device)
|
||||||
logger.info("Upsampled video latents: %s", list(latents.shape))
|
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 import envs
|
||||||
from sglang.multimodal_gen.configs.models.encoders import T5Config
|
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.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.configs.quantization.nunchaku import NunchakuSVDQuantArgs
|
||||||
from sglang.multimodal_gen.runtime.disaggregation.disagg_args import (
|
from sglang.multimodal_gen.runtime.disaggregation.disagg_args import (
|
||||||
DisaggArgsMixin,
|
DisaggArgsMixin,
|
||||||
@@ -69,6 +72,16 @@ logger = init_logger(__name__)
|
|||||||
# our validated Wan/MOVA workloads, so use a 130 GiB cutoff to keep H200-class
|
# 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.
|
# GPUs on the faster no-offload default while preserving some headroom.
|
||||||
WAN_LAYERWISE_OFFLOAD_AUTO_DISABLE_MEM_GB = 130
|
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):
|
class Backend(str, Enum):
|
||||||
@@ -175,6 +188,7 @@ class ServerArgs(DisaggArgsMixin):
|
|||||||
vae_cpu_offload: bool | None = None
|
vae_cpu_offload: bool | None = None
|
||||||
use_fsdp_inference: bool = False
|
use_fsdp_inference: bool = False
|
||||||
pin_cpu_memory: bool = True
|
pin_cpu_memory: bool = True
|
||||||
|
ltx2_two_stage_device_mode: str | None = None
|
||||||
|
|
||||||
# ComfyUI integration
|
# ComfyUI integration
|
||||||
comfyui_mode: bool = False
|
comfyui_mode: bool = False
|
||||||
@@ -283,6 +297,7 @@ class ServerArgs(DisaggArgsMixin):
|
|||||||
def _adjust_parameters(self):
|
def _adjust_parameters(self):
|
||||||
"""set defaults and normalize values."""
|
"""set defaults and normalize values."""
|
||||||
self._adjust_offload()
|
self._adjust_offload()
|
||||||
|
self._adjust_ltx2_two_stage_device_mode()
|
||||||
self._adjust_path()
|
self._adjust_path()
|
||||||
self._adjust_quant_config()
|
self._adjust_quant_config()
|
||||||
self._adjust_warmup()
|
self._adjust_warmup()
|
||||||
@@ -378,6 +393,62 @@ class ServerArgs(DisaggArgsMixin):
|
|||||||
if self.vae_cpu_offload is None:
|
if self.vae_cpu_offload is None:
|
||||||
self.vae_cpu_offload = True
|
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):
|
def _adjust_attention_backend(self):
|
||||||
if self.attention_backend in ["fa3", "fa4"]:
|
if self.attention_backend in ["fa3", "fa4"]:
|
||||||
self.attention_backend = "fa"
|
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". '
|
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",
|
"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(
|
parser.add_argument(
|
||||||
"--disable-autocast",
|
"--disable-autocast",
|
||||||
action=StoreBoolean,
|
action=StoreBoolean,
|
||||||
|
|||||||
@@ -522,7 +522,9 @@ TWO_GPU_CASES_A = [
|
|||||||
"ltx_2_3_two_stage_ti2v_2gpus",
|
"ltx_2_3_two_stage_ti2v_2gpus",
|
||||||
DiffusionServerArgs(
|
DiffusionServerArgs(
|
||||||
model_path="Lightricks/LTX-2.3",
|
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,
|
TI2V_sampling_params,
|
||||||
),
|
),
|
||||||
@@ -541,7 +543,10 @@ TWO_GPU_CASES_B = [
|
|||||||
"ltx_2.3_two_stage_t2v_2gpus",
|
"ltx_2.3_two_stage_t2v_2gpus",
|
||||||
DiffusionServerArgs(
|
DiffusionServerArgs(
|
||||||
model_path="Lightricks/LTX-2.3",
|
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,
|
T2V_sampling_params,
|
||||||
),
|
),
|
||||||
|
|||||||
Reference in New Issue
Block a user