[diffusion] chore: remove ltx2 snapshot mode (#28533)
This commit is contained in:
@@ -61,10 +61,11 @@ For two-stage pipelines, `--ltx2-two-stage-device-mode` controls transformer res
|
|||||||
|
|
||||||
| Mode | When to use it |
|
| Mode | When to use it |
|
||||||
| --- | --- |
|
| --- | --- |
|
||||||
| `snapshot` | Recommended default. Balances latency and VRAM. |
|
|
||||||
| `resident` | Best latency on high-VRAM GPUs because both DiTs can stay resident. |
|
| `resident` | Best latency on high-VRAM GPUs because both DiTs can stay resident. |
|
||||||
| `original` | Closest to the original two-stage switching semantics. |
|
| `original` | Closest to the original two-stage switching semantics. |
|
||||||
|
|
||||||
|
`snapshot` is kept only as a deprecated compatibility alias for `original` and may be removed after two release cycles; use `original` or `resident` in new configs.
|
||||||
|
|
||||||
Other deployment flags:
|
Other deployment flags:
|
||||||
|
|
||||||
- `--lora-path`: Preload a community LoRA adapter.
|
- `--lora-path`: Preload a community LoRA adapter.
|
||||||
@@ -81,13 +82,13 @@ For latency-oriented LTX serving, prefer CFG parallel over sequence parallelism.
|
|||||||
| Target | Recommended server flags | Notes |
|
| Target | Recommended server flags | Notes |
|
||||||
| --- | --- | --- |
|
| --- | --- | --- |
|
||||||
| LTX-2.3, 1 high-VRAM GPU | `--ltx2-two-stage-device-mode resident` | Fastest two-stage setup when both DiTs fit. |
|
| LTX-2.3, 1 high-VRAM GPU | `--ltx2-two-stage-device-mode resident` | Fastest two-stage setup when both DiTs fit. |
|
||||||
| LTX-2.3, 1 standard GPU | `--ltx2-two-stage-device-mode snapshot` | Lower VRAM than `resident`; use this when H100-class memory is tight. |
|
| LTX-2.3, 1 standard GPU | `--ltx2-two-stage-device-mode original` | Lower VRAM than `resident`; use this when H100-class memory is tight. |
|
||||||
| LTX-2, 2 GPUs | `--num-gpus 2 --enable-cfg-parallel` | Fastest verified 2-GPU setup; keep `--dit-layerwise-offload` disabled unless memory is tight. |
|
| LTX-2, 2 GPUs | `--num-gpus 2 --enable-cfg-parallel` | Fastest verified 2-GPU setup; keep `--dit-layerwise-offload` disabled unless memory is tight. |
|
||||||
| LTX-2.3, 2 GPUs | `--num-gpus 2 --enable-cfg-parallel --ltx2-two-stage-device-mode resident` | Fastest common 2-GPU setup. |
|
| LTX-2.3, 2 GPUs | `--num-gpus 2 --enable-cfg-parallel --ltx2-two-stage-device-mode resident` | Fastest common 2-GPU setup. |
|
||||||
| LTX-2.3, 4 GPUs | `--num-gpus 4 --tp-size 2 --enable-cfg-parallel --ltx2-two-stage-device-mode resident` | Fastest common 4-GPU layout: TP2 inside each CFG branch. |
|
| LTX-2.3, 4 GPUs | `--num-gpus 4 --tp-size 2 --enable-cfg-parallel --ltx2-two-stage-device-mode resident` | Fastest common 4-GPU layout: TP2 inside each CFG branch. |
|
||||||
| Official comparison | `--ltx2-two-stage-device-mode original` | Use this only when matching the original LTX-2.3 stage-switch semantics matters. |
|
| Official comparison | `--ltx2-two-stage-device-mode original` | Use this only when matching the original LTX-2.3 stage-switch semantics matters. |
|
||||||
|
|
||||||
Use `--enable-cfg-parallel` for degree-2 CFG parallel. Use `--cfg-parallel-size` only when you explicitly need a different CFG branch count. If `resident` exceeds available VRAM, keep the same parallelism preset and switch only the device mode to `snapshot`.
|
Use `--enable-cfg-parallel` for degree-2 CFG parallel. Use `--cfg-parallel-size` only when you explicitly need a different CFG branch count. If `resident` exceeds available VRAM, keep the same parallelism preset and switch only the device mode to `original`.
|
||||||
|
|
||||||
On high-VRAM GPUs, add `--text-encoder-cpu-offload false` if text encoding latency matters and you have enough memory.
|
On high-VRAM GPUs, add `--text-encoder-cpu-offload false` if text encoding latency matters and you have enough memory.
|
||||||
|
|
||||||
|
|||||||
@@ -586,11 +586,11 @@ Optimization columns are abbreviated to keep the matrix readable:
|
|||||||
- LTX-2 and LTX-2.3 support both T2V and TI2V (`--image-path`) on one-stage and two-stage pipelines (including HQ).
|
- LTX-2 and LTX-2.3 support both T2V and TI2V (`--image-path`) on one-stage and two-stage pipelines (including HQ).
|
||||||
- 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`.
|
- 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 / LTX-2.3 two-stage also supports `--ltx2-two-stage-device-mode {original,snapshot,resident}`:
|
4. LTX-2 / LTX-2.3 two-stage also supports `--ltx2-two-stage-device-mode {original,resident}`:
|
||||||
- `snapshot` is the default and recommended mode.
|
|
||||||
- `resident` usually provides the best latency/throughput but uses much more VRAM.
|
|
||||||
- `original` keeps official two-stage semantics without the premerged stage-2 transformer path.
|
- `original` keeps official two-stage semantics without the premerged stage-2 transformer path.
|
||||||
- Example (one prior run): `original` `154.67s`, `snapshot` `114.05s`, `resident` `75.71s`; peak VRAM trend is `original < snapshot < resident`.
|
- `resident` usually provides the best latency/throughput but uses much more VRAM.
|
||||||
|
- Default is auto: `resident` on H200/high-memory CUDA GPUs, otherwise `original`.
|
||||||
|
- Deprecated compatibility: `snapshot` is accepted as an alias for `original` and may be removed after two release cycles.
|
||||||
5. Cosmos3 ships in two sizes — `nvidia/Cosmos3-Nano` (8B) and
|
5. Cosmos3 ships in two sizes — `nvidia/Cosmos3-Nano` (8B) and
|
||||||
`nvidia/Cosmos3-Super` (32B). Both share the same pipeline; the only
|
`nvidia/Cosmos3-Super` (32B). Both share the same pipeline; the only
|
||||||
difference is transformer depth and width, picked up from
|
difference is transformer depth and width, picked up from
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ export const LTXDeployment = () => {
|
|||||||
{ id: 'h200', label: '1x H200', subtitle: 'resident', default: true },
|
{ id: 'h200', label: '1x H200', subtitle: 'resident', default: true },
|
||||||
{ id: 'h200-2gpu', label: '2 GPUs', subtitle: 'CFG parallel', default: false },
|
{ id: 'h200-2gpu', label: '2 GPUs', subtitle: 'CFG parallel', default: false },
|
||||||
{ id: 'h200-4gpu', label: '4 GPUs', subtitle: 'TP2 + CFG', default: false },
|
{ id: 'h200-4gpu', label: '4 GPUs', subtitle: 'TP2 + CFG', default: false },
|
||||||
{ id: 'standard', label: 'Standard CUDA', subtitle: 'Snapshot mode', default: false },
|
{ id: 'standard', label: 'Standard CUDA', subtitle: 'Original mode', default: false },
|
||||||
{ id: 'official', label: 'Official Match', subtitle: 'Original switching', default: false },
|
{ id: 'official', label: 'Official Match', subtitle: 'Original switching', default: false },
|
||||||
],
|
],
|
||||||
},
|
},
|
||||||
@@ -121,7 +121,7 @@ export const LTXDeployment = () => {
|
|||||||
if (values.hardware === 'official') {
|
if (values.hardware === 'official') {
|
||||||
return 'original';
|
return 'original';
|
||||||
}
|
}
|
||||||
return 'snapshot';
|
return 'original';
|
||||||
};
|
};
|
||||||
|
|
||||||
const getParallelFlags = () => {
|
const getParallelFlags = () => {
|
||||||
|
|||||||
+13
-160
@@ -1,8 +1,8 @@
|
|||||||
from collections.abc import Callable, Iterator
|
from collections.abc import Iterator
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from functools import lru_cache
|
from functools import lru_cache
|
||||||
from typing import Mapping, MutableMapping, Protocol, Sequence, TypeVar
|
from typing import Mapping, MutableMapping, Protocol, Sequence
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
@@ -29,8 +29,6 @@ from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import DiffusionNvtx
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
_T = TypeVar("_T")
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
@dataclass(slots=True)
|
||||||
class ComponentUse:
|
class ComponentUse:
|
||||||
@@ -75,8 +73,6 @@ class ResidencyState:
|
|||||||
# the ComponentUses of the preceding stages
|
# the ComponentUses of the preceding stages
|
||||||
future_uses: tuple[ComponentUse, ...] = ()
|
future_uses: tuple[ComponentUse, ...] = ()
|
||||||
batch_is_warmup: bool = False
|
batch_is_warmup: bool = False
|
||||||
manager_mode: str = "static"
|
|
||||||
trace_enabled: bool = False
|
|
||||||
|
|
||||||
|
|
||||||
class ResidencyBatch(Protocol):
|
class ResidencyBatch(Protocol):
|
||||||
@@ -142,7 +138,7 @@ class ComponentResidencyManager:
|
|||||||
) -> None:
|
) -> None:
|
||||||
self.pipeline = pipeline
|
self.pipeline = pipeline
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
self.state = ResidencyState(trace_enabled=False)
|
self.state = ResidencyState()
|
||||||
self._stage_names_by_id: dict[int, str] = {}
|
self._stage_names_by_id: dict[int, str] = {}
|
||||||
self._stage_uses_by_index: list[tuple[ComponentUse, ...]] = []
|
self._stage_uses_by_index: list[tuple[ComponentUse, ...]] = []
|
||||||
self._ordered_uses: tuple[ComponentUse, ...] = ()
|
self._ordered_uses: tuple[ComponentUse, ...] = ()
|
||||||
@@ -159,10 +155,6 @@ class ComponentResidencyManager:
|
|||||||
)
|
)
|
||||||
self._uses_seen: dict[str, ComponentUse] = {}
|
self._uses_seen: dict[str, ComponentUse] = {}
|
||||||
|
|
||||||
@property
|
|
||||||
def enabled(self) -> bool:
|
|
||||||
return True
|
|
||||||
|
|
||||||
def refresh_pipeline(self, pipeline: ComponentResidencyPipeline) -> None:
|
def refresh_pipeline(self, pipeline: ComponentResidencyPipeline) -> None:
|
||||||
custom_strategies = dict(pipeline.component_residency_strategies)
|
custom_strategies = dict(pipeline.component_residency_strategies)
|
||||||
if pipeline is not self.pipeline:
|
if pipeline is not self.pipeline:
|
||||||
@@ -186,13 +178,6 @@ class ComponentResidencyManager:
|
|||||||
self.strategy_for.cache_clear()
|
self.strategy_for.cache_clear()
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
|
|
||||||
def register_strategy(
|
|
||||||
self, component_name: str, strategy: ComponentResidencyStrategy
|
|
||||||
) -> None:
|
|
||||||
self.pipeline.component_residency_strategies[component_name] = strategy
|
|
||||||
self._custom_strategies[component_name] = strategy
|
|
||||||
self.strategy_for.cache_clear()
|
|
||||||
|
|
||||||
def begin_request(
|
def begin_request(
|
||||||
self,
|
self,
|
||||||
stages: Sequence[ComponentResidencyStage],
|
stages: Sequence[ComponentResidencyStage],
|
||||||
@@ -201,29 +186,19 @@ class ComponentResidencyManager:
|
|||||||
) -> None:
|
) -> None:
|
||||||
"""A hook called before processing an actual request"""
|
"""A hook called before processing an actual request"""
|
||||||
self.refresh_server_args(server_args)
|
self.refresh_server_args(server_args)
|
||||||
self.state = ResidencyState(
|
self.state = ResidencyState(stages=stages, batch_is_warmup=batch.is_warmup)
|
||||||
stages=stages, batch_is_warmup=batch.is_warmup, trace_enabled=False
|
|
||||||
)
|
|
||||||
self._active_use = None
|
self._active_use = None
|
||||||
self._active_use_module = None
|
self._active_use_module = None
|
||||||
self._disable_active_nvtx()
|
self._disable_active_nvtx()
|
||||||
self._current_use_index = -1
|
self._current_use_index = -1
|
||||||
self._prefetched_use_keys.clear()
|
self._prefetched_use_keys.clear()
|
||||||
self._uses_seen.clear()
|
self._uses_seen.clear()
|
||||||
if self.enabled:
|
self._stage_uses_by_index = [
|
||||||
self._stage_uses_by_index = [
|
tuple(stage.component_uses(server_args, self.stage_name(stage)))
|
||||||
tuple(stage.component_uses(server_args, self.stage_name(stage)))
|
for stage in stages
|
||||||
for stage in stages
|
]
|
||||||
]
|
self._ordered_uses = tuple(
|
||||||
self._ordered_uses = tuple(
|
use for uses in self._stage_uses_by_index for use in uses
|
||||||
use for uses in self._stage_uses_by_index for use in uses
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
self._stage_uses_by_index = []
|
|
||||||
self._ordered_uses = ()
|
|
||||||
self._trace(
|
|
||||||
"request_start",
|
|
||||||
detail=f"stages={len(stages)} uses={len(self._ordered_uses)}",
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def before_stage(
|
def before_stage(
|
||||||
@@ -234,25 +209,10 @@ class ComponentResidencyManager:
|
|||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""called after stage starts"""
|
"""called after stage starts"""
|
||||||
if not self.enabled:
|
|
||||||
return
|
|
||||||
# update state before entering the stage
|
# update state before entering the stage
|
||||||
self.state.stage_index = stage_index
|
self.state.stage_index = stage_index
|
||||||
self.state.stage_name = self.stage_name(stage)
|
self.state.stage_name = self.stage_name(stage)
|
||||||
self.state.next_stage_name = self._next_stage_name(stage_index)
|
self.state.next_stage_name = self._next_stage_name(stage_index)
|
||||||
self._trace("stage_enter", detail=f"index={stage_index}")
|
|
||||||
|
|
||||||
def after_stage(self, stage_index: int) -> None:
|
|
||||||
"""called after stage exits"""
|
|
||||||
if not self.enabled:
|
|
||||||
return
|
|
||||||
self._trace("stage_exit", detail=f"index={stage_index}")
|
|
||||||
|
|
||||||
def before_use(self, use: ComponentUse, module: nn.Module | None = None) -> None:
|
|
||||||
"""component use-site starts"""
|
|
||||||
if not self.enabled:
|
|
||||||
return
|
|
||||||
self.begin_use(use, module=module)
|
|
||||||
|
|
||||||
def begin_use(self, use: ComponentUse, module: nn.Module | None = None) -> None:
|
def begin_use(self, use: ComponentUse, module: nn.Module | None = None) -> None:
|
||||||
"""Begin one sequential component use interval. this is idempotent
|
"""Begin one sequential component use interval. this is idempotent
|
||||||
@@ -321,26 +281,8 @@ class ComponentResidencyManager:
|
|||||||
finally:
|
finally:
|
||||||
self.end_use(use, module=module)
|
self.end_use(use, module=module)
|
||||||
|
|
||||||
def call_component(
|
|
||||||
self,
|
|
||||||
use: ComponentUse,
|
|
||||||
module: Callable[..., _T],
|
|
||||||
*args,
|
|
||||||
**kwargs,
|
|
||||||
) -> _T:
|
|
||||||
with self.use_component(use):
|
|
||||||
return module(*args, **kwargs)
|
|
||||||
|
|
||||||
def prefetch_use(self, use: ComponentUse) -> None:
|
|
||||||
"""Prepare a future use without blocking the current use."""
|
|
||||||
if not self.enabled:
|
|
||||||
return
|
|
||||||
self._prefetch_use(use)
|
|
||||||
|
|
||||||
def ensure_ready(self, use: ComponentUse, module: nn.Module | None = None) -> None:
|
def ensure_ready(self, use: ComponentUse, module: nn.Module | None = None) -> None:
|
||||||
"""Prepare a shared component and wait without making it the active use."""
|
"""Prepare a shared component and wait without making it the active use."""
|
||||||
if not self.enabled:
|
|
||||||
return
|
|
||||||
self._prepare_forward_use(use, module=module)
|
self._prepare_forward_use(use, module=module)
|
||||||
|
|
||||||
def remove_nvtx_hooks_for_module(self, module: nn.Module | None) -> None:
|
def remove_nvtx_hooks_for_module(self, module: nn.Module | None) -> None:
|
||||||
@@ -357,19 +299,6 @@ class ComponentResidencyManager:
|
|||||||
hooks.remove_hooks()
|
hooks.remove_hooks()
|
||||||
del self._nvtx_hooks_by_use_key[key]
|
del self._nvtx_hooks_by_use_key[key]
|
||||||
|
|
||||||
def prefetch_checkpoint(self, anchor: ComponentUse | None = None) -> None:
|
|
||||||
"""Give the manager a timeline overlap point.
|
|
||||||
|
|
||||||
1. Locate the anchor or current use in the ordered timeline.
|
|
||||||
2. Find the next prefetchable memory-intensive use.
|
|
||||||
3. Prepare it opportunistically without waiting.
|
|
||||||
"""
|
|
||||||
if not self.enabled:
|
|
||||||
return
|
|
||||||
if anchor is not None:
|
|
||||||
self._mark_current_use(anchor)
|
|
||||||
self._prefetch_next_memory_intensive_use()
|
|
||||||
|
|
||||||
def finish_active_use(self, *, prefetch_next: bool = True) -> None:
|
def finish_active_use(self, *, prefetch_next: bool = True) -> None:
|
||||||
"""Finish the currently active sequential use, if any."""
|
"""Finish the currently active sequential use, if any."""
|
||||||
if self._active_use is None:
|
if self._active_use is None:
|
||||||
@@ -393,14 +322,11 @@ class ComponentResidencyManager:
|
|||||||
"""Prepare a component that is about to run and wait until it is ready."""
|
"""Prepare a component that is about to run and wait until it is ready."""
|
||||||
module = module or self.get_module(use.component_name)
|
module = module or self.get_module(use.component_name)
|
||||||
if module is None:
|
if module is None:
|
||||||
self._trace("skip_missing", use)
|
|
||||||
return None
|
return None
|
||||||
strategy = self.strategy_for(use.component_name, module)
|
strategy = self.strategy_for(use.component_name, module)
|
||||||
self._uses_seen[use.component_name] = use
|
self._uses_seen[use.component_name] = use
|
||||||
self.state.current_use = use
|
self.state.current_use = use
|
||||||
self._trace("prepare", use, strategy, module)
|
|
||||||
strategy.prepare_for_use(module, use, self.state)
|
strategy.prepare_for_use(module, use, self.state)
|
||||||
self._trace("wait", use, strategy, module)
|
|
||||||
strategy.wait_for_use(module, use, self.state)
|
strategy.wait_for_use(module, use, self.state)
|
||||||
return module
|
return module
|
||||||
|
|
||||||
@@ -466,33 +392,24 @@ class ComponentResidencyManager:
|
|||||||
def _prefetch_use(self, use: ComponentUse) -> None:
|
def _prefetch_use(self, use: ComponentUse) -> None:
|
||||||
"""Prepare a future component opportunistically without waiting.
|
"""Prepare a future component opportunistically without waiting.
|
||||||
|
|
||||||
This is called when the component is memory-intensive so it may takes a long time to prefetch.
|
This is called for memory-intensive future uses where H2D placement can
|
||||||
|
overlap with the current stage.
|
||||||
manager will perform the prefetch at some checkpoints, if necessary
|
|
||||||
"""
|
"""
|
||||||
if not use.allow_prefetch:
|
if not use.allow_prefetch:
|
||||||
return
|
return
|
||||||
module = self.get_module(use.component_name)
|
module = self.get_module(use.component_name)
|
||||||
if module is None:
|
if module is None:
|
||||||
self._trace("skip_missing", use)
|
|
||||||
return
|
return
|
||||||
strategy = self.strategy_for(use.component_name, module)
|
strategy = self.strategy_for(use.component_name, module)
|
||||||
if isinstance(strategy, VanillaD2HStrategy) and self._active_use is not None:
|
if isinstance(strategy, VanillaD2HStrategy) and self._active_use is not None:
|
||||||
# Avoid making two vanilla-offloaded heavy components resident before
|
# Avoid making two vanilla-offloaded heavy components resident before
|
||||||
# a budget-aware planner can prove the overlap is safe.
|
# a budget-aware planner can prove the overlap is safe.
|
||||||
self._trace("prefetch_skip_active_vanilla", use, strategy, module)
|
|
||||||
return
|
return
|
||||||
|
|
||||||
self._uses_seen[use.component_name] = use
|
self._uses_seen[use.component_name] = use
|
||||||
self._trace("prefetch", use, strategy, module)
|
|
||||||
if strategy.prefetch_for_use(module, use, self.state):
|
if strategy.prefetch_for_use(module, use, self.state):
|
||||||
self._prefetched_use_keys.add(self._use_key(use))
|
self._prefetched_use_keys.add(self._use_key(use))
|
||||||
|
|
||||||
def after_use(self, use: ComponentUse) -> None:
|
|
||||||
if not self.enabled:
|
|
||||||
return
|
|
||||||
self.end_use(use)
|
|
||||||
|
|
||||||
def _finish_use(
|
def _finish_use(
|
||||||
self,
|
self,
|
||||||
use: ComponentUse,
|
use: ComponentUse,
|
||||||
@@ -503,28 +420,18 @@ class ComponentResidencyManager:
|
|||||||
"""finish a specific use by keeping them resident or call finish_use hook"""
|
"""finish a specific use by keeping them resident or call finish_use hook"""
|
||||||
module = module or self.get_module(use.component_name)
|
module = module or self.get_module(use.component_name)
|
||||||
if module is None:
|
if module is None:
|
||||||
self._trace("skip_missing", use)
|
|
||||||
return
|
return
|
||||||
should_keep = (
|
should_keep = (
|
||||||
keep_on_warmup and self.state.batch_is_warmup
|
keep_on_warmup and self.state.batch_is_warmup
|
||||||
) or self._should_keep_after_use(use)
|
) or self._should_keep_after_use(use)
|
||||||
if should_keep:
|
if should_keep:
|
||||||
self._trace(
|
|
||||||
"keep",
|
|
||||||
use,
|
|
||||||
self.strategy_for(use.component_name, module),
|
|
||||||
module,
|
|
||||||
)
|
|
||||||
return
|
return
|
||||||
strategy = self.strategy_for(use.component_name, module)
|
strategy = self.strategy_for(use.component_name, module)
|
||||||
self._trace("finish", use, strategy, module)
|
|
||||||
was_on_cuda = self._module_on_cuda(module)
|
was_on_cuda = self._module_on_cuda(module)
|
||||||
strategy.finish_use(module, use, self.state)
|
strategy.finish_use(module, use, self.state)
|
||||||
self._empty_cache_after_large_release(use, strategy, module, was_on_cuda)
|
self._empty_cache_after_large_release(use, strategy, module, was_on_cuda)
|
||||||
|
|
||||||
def finish_request(self) -> None:
|
def finish_request(self) -> None:
|
||||||
if not self.enabled and not self._uses_seen and self._active_use is None:
|
|
||||||
return
|
|
||||||
# 1. Close the currently active sequential use.
|
# 1. Close the currently active sequential use.
|
||||||
self.finish_active_use(prefetch_next=False)
|
self.finish_active_use(prefetch_next=False)
|
||||||
# 2. Pick components that should be ready for the next request.
|
# 2. Pick components that should be ready for the next request.
|
||||||
@@ -535,36 +442,19 @@ class ComponentResidencyManager:
|
|||||||
if module is None:
|
if module is None:
|
||||||
continue
|
continue
|
||||||
if self.state.batch_is_warmup and use.keep_ready_after_warmup:
|
if self.state.batch_is_warmup and use.keep_ready_after_warmup:
|
||||||
self._trace(
|
|
||||||
"request_keep_warmup",
|
|
||||||
use,
|
|
||||||
self.strategy_for(component_name, module),
|
|
||||||
module,
|
|
||||||
)
|
|
||||||
continue
|
continue
|
||||||
preferred = component_name in preferred_uses
|
preferred = component_name in preferred_uses
|
||||||
if not preferred and self._should_keep_single_dit(component_name):
|
if not preferred and self._should_keep_single_dit(component_name):
|
||||||
self._trace(
|
|
||||||
"keep",
|
|
||||||
use,
|
|
||||||
self.strategy_for(component_name, module),
|
|
||||||
module,
|
|
||||||
detail="single_dit",
|
|
||||||
)
|
|
||||||
continue
|
continue
|
||||||
strategy = self.strategy_for(component_name, module)
|
strategy = self.strategy_for(component_name, module)
|
||||||
if preferred and not self.state.batch_is_warmup:
|
if preferred and not self.state.batch_is_warmup:
|
||||||
self._trace("request_prefetch", use, strategy, module)
|
|
||||||
strategy.prepare_after_request(module, use, self.state)
|
strategy.prepare_after_request(module, use, self.state)
|
||||||
else:
|
else:
|
||||||
action = "request_resident" if preferred else "request_finish"
|
|
||||||
self._trace(action, use, strategy, module)
|
|
||||||
was_on_cuda = self._module_on_cuda(module)
|
was_on_cuda = self._module_on_cuda(module)
|
||||||
strategy.finish_request(module, use, self.state, preferred=preferred)
|
strategy.finish_request(module, use, self.state, preferred=preferred)
|
||||||
self._empty_cache_after_large_release(
|
self._empty_cache_after_large_release(
|
||||||
use, strategy, module, was_on_cuda
|
use, strategy, module, was_on_cuda
|
||||||
)
|
)
|
||||||
self._trace("request_end")
|
|
||||||
|
|
||||||
def stage_name(self, stage: ComponentResidencyStage) -> str:
|
def stage_name(self, stage: ComponentResidencyStage) -> str:
|
||||||
return self._stage_names_by_id.get(id(stage), stage.__class__.__name__)
|
return self._stage_names_by_id.get(id(stage), stage.__class__.__name__)
|
||||||
@@ -593,12 +483,6 @@ class ComponentResidencyManager:
|
|||||||
component_name, module, self.server_args
|
component_name, module, self.server_args
|
||||||
)
|
)
|
||||||
|
|
||||||
def _stage_uses(self, stage_index: int) -> tuple[ComponentUse, ...]:
|
|
||||||
"""Returns the ComponentUse(s) of a specific stage"""
|
|
||||||
if stage_index < 0 or stage_index >= len(self._stage_uses_by_index):
|
|
||||||
return ()
|
|
||||||
return self._stage_uses_by_index[stage_index]
|
|
||||||
|
|
||||||
def _next_stage_name(self, stage_index: int) -> str | None:
|
def _next_stage_name(self, stage_index: int) -> str | None:
|
||||||
next_index = stage_index + 1
|
next_index = stage_index + 1
|
||||||
if next_index < 0 or next_index >= len(self.state.stages):
|
if next_index < 0 or next_index >= len(self.state.stages):
|
||||||
@@ -629,7 +513,7 @@ class ComponentResidencyManager:
|
|||||||
continue
|
continue
|
||||||
if self._use_key(use) in self._prefetched_use_keys:
|
if self._use_key(use) in self._prefetched_use_keys:
|
||||||
return
|
return
|
||||||
self.prefetch_use(use)
|
self._prefetch_use(use)
|
||||||
return
|
return
|
||||||
|
|
||||||
def _should_keep_after_use(self, use: ComponentUse) -> bool:
|
def _should_keep_after_use(self, use: ComponentUse) -> bool:
|
||||||
@@ -684,36 +568,6 @@ class ComponentResidencyManager:
|
|||||||
def _use_key(use: ComponentUse) -> tuple[str, str, str | None]:
|
def _use_key(use: ComponentUse) -> tuple[str, str, str | None]:
|
||||||
return (use.stage_name, use.component_name, use.phase)
|
return (use.stage_name, use.component_name, use.phase)
|
||||||
|
|
||||||
def _trace(
|
|
||||||
self,
|
|
||||||
action: str,
|
|
||||||
use: ComponentUse | None = None,
|
|
||||||
strategy: ComponentResidencyStrategy | None = None,
|
|
||||||
module: nn.Module | None = None,
|
|
||||||
*,
|
|
||||||
component_name: str | None = None,
|
|
||||||
detail: str = "",
|
|
||||||
) -> None:
|
|
||||||
if not self.state.trace_enabled:
|
|
||||||
return
|
|
||||||
if use is not None:
|
|
||||||
component_name = use.component_name
|
|
||||||
device = self._module_device(module)
|
|
||||||
logger.info(
|
|
||||||
"[component_residency] action=%s stage=%s next_stage=%s component=%s "
|
|
||||||
"strategy=%s phase=%s device=%s warmup=%s mode=%s %s",
|
|
||||||
action,
|
|
||||||
self.state.stage_name,
|
|
||||||
self.state.next_stage_name,
|
|
||||||
component_name,
|
|
||||||
strategy.name if strategy is not None else None,
|
|
||||||
use.phase if use is not None else None,
|
|
||||||
device,
|
|
||||||
self.state.batch_is_warmup,
|
|
||||||
self.state.manager_mode,
|
|
||||||
detail,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _module_device(self, module: nn.Module | None) -> str | None:
|
def _module_device(self, module: nn.Module | None) -> str | None:
|
||||||
if module is None:
|
if module is None:
|
||||||
return None
|
return None
|
||||||
@@ -743,7 +597,6 @@ class ComponentResidencyManager:
|
|||||||
if not torch.get_device_module().is_available():
|
if not torch.get_device_module().is_available():
|
||||||
return
|
return
|
||||||
torch.get_device_module().empty_cache()
|
torch.get_device_module().empty_cache()
|
||||||
self._trace("empty_cache", use, strategy, module, detail="after_release")
|
|
||||||
|
|
||||||
|
|
||||||
_GLOBAL_COMPONENT_RESIDENCY_MANAGER: ComponentResidencyManager | None = None
|
_GLOBAL_COMPONENT_RESIDENCY_MANAGER: ComponentResidencyManager | None = None
|
||||||
|
|||||||
@@ -18,16 +18,11 @@ 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.managers.memory_managers.component_manager import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
||||||
ComponentResidencyStrategy,
|
ComponentResidencyStrategy,
|
||||||
ComponentUse,
|
ComponentUse,
|
||||||
ResidencyState,
|
ResidencyState,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_resident_strategies import (
|
|
||||||
SnapshotModuleResidency,
|
|
||||||
SnapshotStrategy,
|
|
||||||
)
|
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||||
ComposedPipelineBase,
|
ComposedPipelineBase,
|
||||||
)
|
)
|
||||||
@@ -51,12 +46,11 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.l
|
|||||||
LTX2TextConnectorStage,
|
LTX2TextConnectorStage,
|
||||||
LTX2UpsampleStage,
|
LTX2UpsampleStage,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
|
||||||
from sglang.multimodal_gen.runtime.server_args import (
|
from sglang.multimodal_gen.runtime.server_args import (
|
||||||
LTX2_RESIDENT_AUTO_ENABLE_MEM_GB,
|
LTX2_TWO_STAGE_DEVICE_MODE_CHOICES,
|
||||||
ServerArgs,
|
ServerArgs,
|
||||||
|
_normalize_ltx2_two_stage_device_mode,
|
||||||
)
|
)
|
||||||
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__)
|
||||||
@@ -413,10 +407,6 @@ class LTX2TwoStageResidencyStrategy(ComponentResidencyStrategy):
|
|||||||
if param is not None and param.device.type == "cpu":
|
if param is not None and param.device.type == "cpu":
|
||||||
module.to(get_local_torch_device(), non_blocking=True)
|
module.to(get_local_torch_device(), non_blocking=True)
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _module_is_on_gpu(module: torch.nn.Module | None) -> bool:
|
|
||||||
return SnapshotModuleResidency.is_on_gpu(module)
|
|
||||||
|
|
||||||
|
|
||||||
class LTX2OriginalResidencyStrategy(LTX2TwoStageResidencyStrategy):
|
class LTX2OriginalResidencyStrategy(LTX2TwoStageResidencyStrategy):
|
||||||
pass
|
pass
|
||||||
@@ -442,265 +432,6 @@ class LTX2ResidentResidencyStrategy(LTX2TwoStageResidencyStrategy):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
class LTX2SnapshotResidencyStrategy(LTX2TwoStageResidencyStrategy):
|
|
||||||
"""
|
|
||||||
Snapshot mode keeps CPU snapshots and prefetches the target DiT with async H2D. (only with pre-merged lora enabled)
|
|
||||||
|
|
||||||
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`.
|
|
||||||
"""
|
|
||||||
|
|
||||||
name = "ltx2_snapshot"
|
|
||||||
|
|
||||||
def __init__(self, manager: "LTX2TwoStageResidencyController") -> None:
|
|
||||||
super().__init__(manager)
|
|
||||||
self._snapshot_strategy = SnapshotStrategy(
|
|
||||||
pin_cpu_memory=manager.server_args.pin_cpu_memory,
|
|
||||||
enable_async_prefetch=manager.server_args.dit_cpu_offload,
|
|
||||||
)
|
|
||||||
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",
|
|
||||||
)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _module_name_for_phase(phase: str | None) -> str | None:
|
|
||||||
if phase == "stage1":
|
|
||||||
return "transformer"
|
|
||||||
if phase == "stage2":
|
|
||||||
return "transformer_2"
|
|
||||||
return None
|
|
||||||
|
|
||||||
def _resolve_snapshot_low_vram_mode(self) -> bool:
|
|
||||||
if 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
|
|
||||||
|
|
||||||
def initialize(self) -> None:
|
|
||||||
# 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()
|
|
||||||
self.manager._sync_refinement_stage_transformer("stage1")
|
|
||||||
self._record_component_ready("transformer")
|
|
||||||
|
|
||||||
def enter_phase(self, phase: str) -> bool:
|
|
||||||
if self.server_args.dit_cpu_offload:
|
|
||||||
target_module_name = self._module_name_for_phase(phase)
|
|
||||||
if target_module_name is None:
|
|
||||||
return False
|
|
||||||
target_module = self.pipeline.get_module(target_module_name)
|
|
||||||
if 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 not self._snapshot_strategy.is_ready(
|
|
||||||
target_module_name
|
|
||||||
):
|
|
||||||
self._release_stage1_for_low_vram()
|
|
||||||
|
|
||||||
# make sure the component is pre-fetched
|
|
||||||
if not self._snapshot_strategy.is_ready(target_module_name):
|
|
||||||
if self._module_is_on_gpu(target_module):
|
|
||||||
self._record_component_ready(target_module_name)
|
|
||||||
else:
|
|
||||||
self._snapshot_strategy.prefetch_component(
|
|
||||||
target_module_name, target_module
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
component_name = self._module_name_for_phase(phase)
|
|
||||||
if component_name is not None:
|
|
||||||
self._record_component_ready(component_name)
|
|
||||||
|
|
||||||
self.manager._sync_refinement_stage_transformer(phase)
|
|
||||||
self.manager._active_phase = phase
|
|
||||||
return True
|
|
||||||
|
|
||||||
def prepare_after_request(
|
|
||||||
self,
|
|
||||||
module: torch.nn.Module,
|
|
||||||
use: ComponentUse,
|
|
||||||
state: ResidencyState,
|
|
||||||
) -> None:
|
|
||||||
phase = self._phase(use)
|
|
||||||
if phase != "stage1":
|
|
||||||
return
|
|
||||||
if self.server_args.dit_cpu_offload:
|
|
||||||
target_module = self.pipeline.get_module("transformer")
|
|
||||||
if self._module_is_on_gpu(target_module):
|
|
||||||
self._record_component_ready("transformer")
|
|
||||||
elif not self._snapshot_strategy.is_ready("transformer"):
|
|
||||||
if self._snapshot_low_vram_mode:
|
|
||||||
self._release_stage2_for_low_vram()
|
|
||||||
self._snapshot_strategy.prefetch_component("transformer", target_module)
|
|
||||||
else:
|
|
||||||
self._record_component_ready("transformer")
|
|
||||||
self.manager._sync_refinement_stage_transformer("stage1")
|
|
||||||
self.manager._active_phase = "stage1"
|
|
||||||
|
|
||||||
def finish_request(
|
|
||||||
self,
|
|
||||||
module: torch.nn.Module,
|
|
||||||
use: ComponentUse,
|
|
||||||
state: ResidencyState,
|
|
||||||
*,
|
|
||||||
preferred: bool,
|
|
||||||
) -> None:
|
|
||||||
if (
|
|
||||||
preferred
|
|
||||||
and state.batch_is_warmup
|
|
||||||
and self._snapshot_low_vram_mode
|
|
||||||
and self._phase(use) == "stage1"
|
|
||||||
):
|
|
||||||
# keep the text encoder warm, but avoid stage1 DiT overlap before the first real request
|
|
||||||
self.manager._active_phase = None
|
|
||||||
return
|
|
||||||
super().finish_request(module, use, state, preferred=preferred)
|
|
||||||
|
|
||||||
def finish_use(
|
|
||||||
self,
|
|
||||||
module: torch.nn.Module,
|
|
||||||
use: ComponentUse,
|
|
||||||
state: ResidencyState,
|
|
||||||
) -> None:
|
|
||||||
phase = self._phase(use)
|
|
||||||
if self.server_args.dit_cpu_offload:
|
|
||||||
# release cuda storage
|
|
||||||
self._snapshot_strategy.release_component(use.component_name, module)
|
|
||||||
if (
|
|
||||||
phase == "stage2"
|
|
||||||
and self._snapshot_release_empty_cache
|
|
||||||
and torch.get_device_module().is_available()
|
|
||||||
):
|
|
||||||
torch.get_device_module().empty_cache()
|
|
||||||
|
|
||||||
def ensure_phase_ready(self, phase: str | None) -> None:
|
|
||||||
component_name = self._module_name_for_phase(phase)
|
|
||||||
if component_name is None:
|
|
||||||
return
|
|
||||||
self._snapshot_strategy.wait_component_ready(component_name)
|
|
||||||
|
|
||||||
def _capture_module_cpu_snapshot(self, module_name: str) -> None:
|
|
||||||
module = self.pipeline.get_module(module_name)
|
|
||||||
if module is None:
|
|
||||||
raise ValueError(f"Module {module_name} is not available.")
|
|
||||||
self._snapshot_strategy.capture(module_name, module)
|
|
||||||
|
|
||||||
def _release_module_to_cpu_snapshot(self, module_name: str) -> None:
|
|
||||||
module = self.pipeline.get_module(module_name)
|
|
||||||
if module is None:
|
|
||||||
return
|
|
||||||
self._snapshot_strategy.release_component(module_name, module)
|
|
||||||
|
|
||||||
def _release_stage1_for_low_vram(self) -> None:
|
|
||||||
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")
|
|
||||||
|
|
||||||
def _release_stage2_for_low_vram(self) -> None:
|
|
||||||
stage2_module = self.pipeline.get_module("transformer_2")
|
|
||||||
stage2_param = (
|
|
||||||
next(stage2_module.parameters(), None)
|
|
||||||
if stage2_module is not None
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
if stage2_param is not None and stage2_param.device.type == "cuda":
|
|
||||||
self._release_module_to_cpu_snapshot("transformer_2")
|
|
||||||
|
|
||||||
def _record_component_ready(self, module_name: str) -> None:
|
|
||||||
self._snapshot_strategy.record_ready(
|
|
||||||
module_name, self.pipeline.get_module(module_name)
|
|
||||||
)
|
|
||||||
|
|
||||||
def prefetch_for_use(
|
|
||||||
self,
|
|
||||||
module: torch.nn.Module,
|
|
||||||
use: ComponentUse,
|
|
||||||
state: ResidencyState,
|
|
||||||
) -> bool:
|
|
||||||
if not self.server_args.dit_cpu_offload:
|
|
||||||
return True
|
|
||||||
phase = self._phase(use)
|
|
||||||
if (
|
|
||||||
self._snapshot_low_vram_mode
|
|
||||||
and phase == "stage1"
|
|
||||||
and state.current_use is not None
|
|
||||||
and state.current_use.component_name.startswith("text_encoder")
|
|
||||||
):
|
|
||||||
return False
|
|
||||||
if phase == "stage2":
|
|
||||||
if self._snapshot_strategy.is_ready("transformer_2"):
|
|
||||||
return True
|
|
||||||
if self._snapshot_low_vram_mode and state.current_use is not None:
|
|
||||||
return False
|
|
||||||
if self._snapshot_low_vram_mode:
|
|
||||||
self._release_stage1_for_low_vram()
|
|
||||||
self._snapshot_strategy.prefetch_component(use.component_name, module)
|
|
||||||
return True
|
|
||||||
|
|
||||||
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 outside low-VRAM mode 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 self._snapshot_low_vram_mode
|
|
||||||
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.manager._active_phase = "stage1"
|
|
||||||
|
|
||||||
|
|
||||||
class LTX2TwoStageResidencyController:
|
class LTX2TwoStageResidencyController:
|
||||||
"""
|
"""
|
||||||
LTX-2.3 two-stage residency controller.
|
LTX-2.3 two-stage residency controller.
|
||||||
@@ -709,11 +440,10 @@ class LTX2TwoStageResidencyController:
|
|||||||
|
|
||||||
Modes:
|
Modes:
|
||||||
- resident: keep both DiTs on GPU; phase switch is pointer rebinding only.
|
- resident: keep both DiTs on GPU; phase switch is pointer rebinding only.
|
||||||
- snapshot: keep CPU snapshots and prefetch the target DiT.
|
|
||||||
- original: official two-stage semantics without premerged stage-2.
|
- original: official two-stage semantics without premerged stage-2.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
VALID_MODES = ("original", "snapshot", "resident")
|
VALID_MODES = ("original", "resident")
|
||||||
|
|
||||||
def __init__(self, pipeline: "LTX2TwoStagePipeline", server_args: ServerArgs):
|
def __init__(self, pipeline: "LTX2TwoStagePipeline", server_args: ServerArgs):
|
||||||
self.pipeline = pipeline
|
self.pipeline = pipeline
|
||||||
@@ -727,17 +457,21 @@ class LTX2TwoStageResidencyController:
|
|||||||
mode = server_args.ltx2_two_stage_device_mode
|
mode = server_args.ltx2_two_stage_device_mode
|
||||||
if mode is None:
|
if mode is None:
|
||||||
env_mode = os.getenv("SGLANG_LTX2_TWO_STAGE_DEVICE_MODE")
|
env_mode = os.getenv("SGLANG_LTX2_TWO_STAGE_DEVICE_MODE")
|
||||||
mode = env_mode.lower() if env_mode else "snapshot"
|
mode = (
|
||||||
|
_normalize_ltx2_two_stage_device_mode(env_mode)
|
||||||
|
if env_mode
|
||||||
|
else "original"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
mode = _normalize_ltx2_two_stage_device_mode(mode)
|
||||||
if mode not in cls.VALID_MODES:
|
if mode not in cls.VALID_MODES:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Invalid ltx2_two_stage_device_mode={mode!r}. "
|
f"Invalid ltx2_two_stage_device_mode={mode!r}. "
|
||||||
f"Expected one of {cls.VALID_MODES}."
|
f"Expected one of {LTX2_TWO_STAGE_DEVICE_MODE_CHOICES}."
|
||||||
)
|
)
|
||||||
return mode
|
return mode
|
||||||
|
|
||||||
def _build_strategy(self) -> LTX2TwoStageResidencyStrategy:
|
def _build_strategy(self) -> LTX2TwoStageResidencyStrategy:
|
||||||
if self.mode == "snapshot":
|
|
||||||
return LTX2SnapshotResidencyStrategy(self)
|
|
||||||
if self.mode == "resident":
|
if self.mode == "resident":
|
||||||
return LTX2ResidentResidencyStrategy(self)
|
return LTX2ResidentResidencyStrategy(self)
|
||||||
return LTX2OriginalResidencyStrategy(self)
|
return LTX2OriginalResidencyStrategy(self)
|
||||||
@@ -750,11 +484,11 @@ class LTX2TwoStageResidencyController:
|
|||||||
def should_use_premerged(self) -> bool:
|
def should_use_premerged(self) -> bool:
|
||||||
"""Whether to keep a pre-merged stage-2 DiT for LTX-2.3 two-stage.
|
"""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
|
We only enable this optimization for resident native LTX-2.3 two-stage
|
||||||
users did not explicitly provide a stage-1 LoRA path
|
and when users did not explicitly provide a stage-1 LoRA path
|
||||||
"""
|
"""
|
||||||
return (
|
return (
|
||||||
self.mode != "original"
|
self.mode == "resident"
|
||||||
and self.pipeline._should_merge_stage2_distilled_lora(self.server_args)
|
and self.pipeline._should_merge_stage2_distilled_lora(self.server_args)
|
||||||
and self.pipeline._stage1_lora_path is None
|
and self.pipeline._stage1_lora_path is None
|
||||||
)
|
)
|
||||||
@@ -868,7 +602,7 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
|
|||||||
self.memory_usages["transformer_2"] = memory_usage
|
self.memory_usages["transformer_2"] = memory_usage
|
||||||
|
|
||||||
# Reuse the canonical LoRA path used by legacy switching to reduce
|
# Reuse the canonical LoRA path used by legacy switching to reduce
|
||||||
# precision drift between snapshot mode and origin/main behavior.
|
# precision drift against original two-stage behavior.
|
||||||
self.set_lora(
|
self.set_lora(
|
||||||
lora_nickname="ltx2_stage2_distilled",
|
lora_nickname="ltx2_stage2_distilled",
|
||||||
lora_path=self._distilled_lora_path,
|
lora_path=self._distilled_lora_path,
|
||||||
@@ -878,9 +612,9 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def should_skip_ltx2_lora_switch_stage(self) -> bool:
|
def should_skip_ltx2_lora_switch_stage(self) -> bool:
|
||||||
return self._use_premerged_stage2_transformer and self._ltx2_residency.mode in (
|
return (
|
||||||
"snapshot",
|
self._use_premerged_stage2_transformer
|
||||||
"resident",
|
and self._ltx2_residency.mode == "resident"
|
||||||
)
|
)
|
||||||
|
|
||||||
def _get_stage_distilled_lora_strength(
|
def _get_stage_distilled_lora_strength(
|
||||||
|
|||||||
@@ -71,9 +71,6 @@ class PipelineExecutor(ABC):
|
|||||||
stage, stage_index, batch, server_args
|
stage, stage_index, batch, server_args
|
||||||
)
|
)
|
||||||
|
|
||||||
def after_stage(self, stage_index: int) -> None:
|
|
||||||
self.component_residency_manager.after_stage(stage_index)
|
|
||||||
|
|
||||||
def finish_component_residency_request(self) -> None:
|
def finish_component_residency_request(self) -> None:
|
||||||
self.component_residency_manager.finish_request()
|
self.component_residency_manager.finish_request()
|
||||||
|
|
||||||
@@ -118,7 +115,6 @@ class PipelineExecutor(ABC):
|
|||||||
payload = self.run_stage_with_context(
|
payload = self.run_stage_with_context(
|
||||||
stage, payload, server_args, run_stage
|
stage, payload, server_args, run_stage
|
||||||
)
|
)
|
||||||
self.after_stage(stage_index)
|
|
||||||
return payload
|
return payload
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -66,7 +66,8 @@ from sglang.multimodal_gen.utils import (
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
LTX2_TWO_STAGE_DEVICE_MODES = ("original", "snapshot", "resident")
|
LTX2_TWO_STAGE_DEVICE_MODES = ("original", "resident")
|
||||||
|
LTX2_TWO_STAGE_DEVICE_MODE_CHOICES = (*LTX2_TWO_STAGE_DEVICE_MODES, "snapshot")
|
||||||
LTX2_TWO_STAGE_PIPELINE_NAMES = ("LTX2TwoStagePipeline", "LTX2TwoStageHQPipeline")
|
LTX2_TWO_STAGE_PIPELINE_NAMES = ("LTX2TwoStagePipeline", "LTX2TwoStageHQPipeline")
|
||||||
# H200-class GPUs (>=130 GiB total) can usually keep both LTX2 DiTs resident.
|
# H200-class GPUs (>=130 GiB total) can usually keep both LTX2 DiTs resident.
|
||||||
LTX2_RESIDENT_AUTO_ENABLE_MEM_GB = 130
|
LTX2_RESIDENT_AUTO_ENABLE_MEM_GB = 130
|
||||||
@@ -77,6 +78,13 @@ def _normalize_ltx2_two_stage_device_mode(mode: str | None) -> str | None:
|
|||||||
if mode is None:
|
if mode is None:
|
||||||
return None
|
return None
|
||||||
mode = mode.lower()
|
mode = mode.lower()
|
||||||
|
if mode == "snapshot":
|
||||||
|
logger.warning(
|
||||||
|
"ltx2_two_stage_device_mode=snapshot is deprecated and is treated "
|
||||||
|
"as original. Please use ltx2_two_stage_device_mode=original or "
|
||||||
|
"resident instead. This alias may be removed after two release cycles."
|
||||||
|
)
|
||||||
|
return "original"
|
||||||
return mode
|
return mode
|
||||||
|
|
||||||
|
|
||||||
@@ -497,7 +505,7 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
if mode not in LTX2_TWO_STAGE_DEVICE_MODES:
|
if mode not in LTX2_TWO_STAGE_DEVICE_MODES:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Invalid ltx2_two_stage_device_mode={mode!r}. "
|
f"Invalid ltx2_two_stage_device_mode={mode!r}. "
|
||||||
f"Expected one of {LTX2_TWO_STAGE_DEVICE_MODES}."
|
f"Expected one of {LTX2_TWO_STAGE_DEVICE_MODE_CHOICES}."
|
||||||
)
|
)
|
||||||
|
|
||||||
self.ltx2_two_stage_device_mode = mode
|
self.ltx2_two_stage_device_mode = mode
|
||||||
@@ -505,9 +513,9 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
def _resolve_default_ltx2_two_stage_device_mode(self) -> str:
|
def _resolve_default_ltx2_two_stage_device_mode(self) -> str:
|
||||||
if not current_platform.is_cuda():
|
if not current_platform.is_cuda():
|
||||||
logger.info(
|
logger.info(
|
||||||
"Automatically set ltx2_two_stage_device_mode=snapshot on non-CUDA platform"
|
"Automatically set ltx2_two_stage_device_mode=original on non-CUDA platform"
|
||||||
)
|
)
|
||||||
return "snapshot"
|
return "original"
|
||||||
|
|
||||||
device_name = str(current_platform.get_device_name(0)).upper()
|
device_name = str(current_platform.get_device_name(0)).upper()
|
||||||
device_total_memory_gb = (
|
device_total_memory_gb = (
|
||||||
@@ -525,11 +533,11 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
return "resident"
|
return "resident"
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"Automatically set ltx2_two_stage_device_mode=snapshot for CUDA GPU (%s, %.2f GiB total)",
|
"Automatically set ltx2_two_stage_device_mode=original for CUDA GPU (%s, %.2f GiB total)",
|
||||||
device_name,
|
device_name,
|
||||||
device_total_memory_gb,
|
device_total_memory_gb,
|
||||||
)
|
)
|
||||||
return "snapshot"
|
return "original"
|
||||||
|
|
||||||
def _is_ltx23_two_stage_pipeline(self) -> bool:
|
def _is_ltx23_two_stage_pipeline(self) -> bool:
|
||||||
return is_ltx2_two_stage_pipeline_name(self.pipeline_class_name) and (
|
return is_ltx2_two_stage_pipeline_name(self.pipeline_class_name) and (
|
||||||
@@ -537,12 +545,6 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
or is_ltx23_native_variant(self.pipeline_config.vae_config.arch_config)
|
or is_ltx23_native_variant(self.pipeline_config.vae_config.arch_config)
|
||||||
)
|
)
|
||||||
|
|
||||||
def _uses_ltx23_snapshot_two_stage_residency(self) -> bool:
|
|
||||||
return (
|
|
||||||
self.ltx2_two_stage_device_mode == "snapshot"
|
|
||||||
and self._is_ltx23_two_stage_pipeline()
|
|
||||||
)
|
|
||||||
|
|
||||||
def _uses_ltx23_high_memory_resident_two_stage_mode(self) -> bool:
|
def _uses_ltx23_high_memory_resident_two_stage_mode(self) -> bool:
|
||||||
if (
|
if (
|
||||||
self.ltx2_two_stage_device_mode != "resident"
|
self.ltx2_two_stage_device_mode != "resident"
|
||||||
@@ -1410,14 +1412,15 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--ltx2-two-stage-device-mode",
|
"--ltx2-two-stage-device-mode",
|
||||||
type=str,
|
type=str,
|
||||||
choices=LTX2_TWO_STAGE_DEVICE_MODES,
|
choices=LTX2_TWO_STAGE_DEVICE_MODE_CHOICES,
|
||||||
default=ServerArgs.ltx2_two_stage_device_mode,
|
default=ServerArgs.ltx2_two_stage_device_mode,
|
||||||
help=(
|
help=(
|
||||||
"LTX-2.3 two-stage device residency mode: "
|
"LTX-2.3 two-stage device residency mode: "
|
||||||
"'original' keeps official two-stage semantics without premerged stage2, "
|
"'original' keeps official two-stage semantics without premerged stage2, "
|
||||||
"'snapshot' keeps premerged stage2 with snapshot-based release, "
|
|
||||||
"'resident' keeps both transformers resident on GPU. "
|
"'resident' keeps both transformers resident on GPU. "
|
||||||
"Default is auto: resident on H200/high-memory CUDA GPUs, otherwise snapshot."
|
"'snapshot' is deprecated, treated as 'original', and may be "
|
||||||
|
"removed after two release cycles. "
|
||||||
|
"Default is auto: resident on H200/high-memory CUDA GPUs, otherwise original."
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
|
|||||||
@@ -107,11 +107,6 @@ class ServerArgsAutoTuner:
|
|||||||
components = (
|
components = (
|
||||||
self._deployment_config().auto_disable_component_offload_components
|
self._deployment_config().auto_disable_component_offload_components
|
||||||
)
|
)
|
||||||
if args._uses_ltx23_snapshot_two_stage_residency():
|
|
||||||
# ltx2 snapshot mode uses DiT offload to release/prefetch stage DiTs between phases
|
|
||||||
components = tuple(
|
|
||||||
component for component in components if component != "dit"
|
|
||||||
)
|
|
||||||
if (
|
if (
|
||||||
args.dit_cpu_offload
|
args.dit_cpu_offload
|
||||||
and "dit" in components
|
and "dit" in components
|
||||||
|
|||||||
@@ -429,11 +429,10 @@ ONE_GPU_CASES: list[DiffusionTestCase] = [
|
|||||||
DiffusionServerArgs(
|
DiffusionServerArgs(
|
||||||
model_path="Lightricks/LTX-2.3",
|
model_path="Lightricks/LTX-2.3",
|
||||||
extras=[
|
extras=[
|
||||||
"--pipeline-class-name LTX2TwoStageHQPipeline --ltx2-two-stage-device-mode snapshot"
|
"--pipeline-class-name LTX2TwoStageHQPipeline --ltx2-two-stage-device-mode original"
|
||||||
],
|
],
|
||||||
env_vars={
|
env_vars={
|
||||||
"PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True",
|
"PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True",
|
||||||
"SGLANG_LTX2_SNAPSHOT_RELEASE_EMPTY_CACHE": "true",
|
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
run_component_accuracy_check=False,
|
run_component_accuracy_check=False,
|
||||||
|
|||||||
@@ -2681,44 +2681,46 @@
|
|||||||
},
|
},
|
||||||
"ltx_2_3_hq_pipeline": {
|
"ltx_2_3_hq_pipeline": {
|
||||||
"stages_ms": {
|
"stages_ms": {
|
||||||
"InputValidationStage": 11.84,
|
"InputValidationStage": 4.63,
|
||||||
"TextEncodingStage": 446.77,
|
"TextEncodingStage": 401.75,
|
||||||
"LTX2TextConnectorStage": 32.96,
|
"LTX2TextConnectorStage": 27.48,
|
||||||
"LTX2HalveResolutionStage": 0.07,
|
"LTX2HalveResolutionStage": 0.04,
|
||||||
"LTX2LoRASwitchStage": 0.03,
|
"LTX2LoRASwitchStage": 180.0,
|
||||||
"LTX2SigmaPreparationStage": 0.38,
|
"LTX2SigmaPreparationStage": 0.26,
|
||||||
"TimestepPreparationStage": 24.67,
|
"TimestepPreparationStage": 14.45,
|
||||||
"LTX2AVLatentPreparationStage": 0.16,
|
"LTX2AVLatentPreparationStage": 0.13,
|
||||||
"LTX2ImageEncodingStage": 1133.15,
|
"LTX2ImageEncodingStage": 57.62,
|
||||||
"LTX2AVDenoisingStage": 19539.95,
|
"LTX2AVDenoisingStage": 12162.51,
|
||||||
"LTX2UpsampleStage": 22.73,
|
"LTX2UpsampleStage": 11.04,
|
||||||
"LTX2RefinementStage": 3059.95,
|
"ltx2_lora_switch_stage2": 9155.58,
|
||||||
"LTX2AVDecodingStage": 1744.05,
|
"ltx2_image_encoding_stage2": 64.54,
|
||||||
|
"LTX2RefinementStage": 3484.86,
|
||||||
|
"LTX2AVDecodingStage": 1054.91,
|
||||||
"per_frame_generation": null
|
"per_frame_generation": null
|
||||||
},
|
},
|
||||||
"denoise_step_ms": {
|
"denoise_step_ms": {
|
||||||
"0": 1677.39,
|
"0": 785.16,
|
||||||
"1": 778.3,
|
"1": 740.86,
|
||||||
"2": 740.93,
|
"2": 740.7,
|
||||||
"3": 1177.58,
|
"3": 742.77,
|
||||||
"4": 751.77,
|
"4": 772.73,
|
||||||
"5": 744.01,
|
"5": 744.43,
|
||||||
"6": 742.02,
|
"6": 743.14,
|
||||||
"7": 1264.88,
|
"7": 743.2,
|
||||||
"8": 744.11,
|
"8": 743.76,
|
||||||
"9": 739.84,
|
"9": 742.57,
|
||||||
"10": 1209.91,
|
"10": 741.31,
|
||||||
"11": 739.62,
|
"11": 744.06,
|
||||||
"12": 736.7,
|
"12": 760.17,
|
||||||
"13": 733.55,
|
"13": 740.69,
|
||||||
"14": 590.44,
|
"14": 369.53,
|
||||||
"15": 903.04,
|
"15": 1193.58,
|
||||||
"16": 876.98,
|
"16": 1144.87,
|
||||||
"17": 879.42
|
"17": 1143.19
|
||||||
},
|
},
|
||||||
"expected_e2e_ms": 27309.99,
|
"expected_e2e_ms": 26673.22,
|
||||||
"expected_avg_denoise_ms": 1159.69,
|
"expected_avg_denoise_ms": 868.06,
|
||||||
"expected_median_denoise_ms": 1199,
|
"expected_median_denoise_ms": 747.62,
|
||||||
"estimated_full_test_time_s": 363.2
|
"estimated_full_test_time_s": 363.2
|
||||||
},
|
},
|
||||||
"qwen_image_t2i_cache_dit_scm_config_diffusers_1gpu": {
|
"qwen_image_t2i_cache_dit_scm_config_diffusers_1gpu": {
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
SGL_TEST_FILES_CI_DATA_REVISION = "f0fd96eab85baed5256d142c7659e0634e7e4410"
|
SGL_TEST_FILES_CI_DATA_REVISION = "51a6a6cd592983e1b8dadc9d7981fac63cd02800"
|
||||||
|
|
||||||
if current_platform.is_npu():
|
if current_platform.is_npu():
|
||||||
SGL_TEST_FILES_CI_DATA_REVISION = "670d66a8a290b62c0c3c077b3e9b0f4a4d9a44e7"
|
SGL_TEST_FILES_CI_DATA_REVISION = "670d66a8a290b62c0c3c077b3e9b0f4a4d9a44e7"
|
||||||
|
|||||||
@@ -882,7 +882,7 @@ class TestOffloadDefaults(unittest.TestCase):
|
|||||||
["text_encoder", "image_encoder", "vae"],
|
["text_encoder", "image_encoder", "vae"],
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_auto_ltx_snapshot_keeps_dit_offload_and_replaces_encoder_cpu_offload(
|
def test_auto_ltx_original_replaces_component_cpu_offload(
|
||||||
self,
|
self,
|
||||||
):
|
):
|
||||||
args = self._from_dict_with_pipeline_config(
|
args = self._from_dict_with_pipeline_config(
|
||||||
@@ -891,13 +891,12 @@ class TestOffloadDefaults(unittest.TestCase):
|
|||||||
kwargs={
|
kwargs={
|
||||||
"model_path": "Lightricks/LTX-2.3",
|
"model_path": "Lightricks/LTX-2.3",
|
||||||
"pipeline_class_name": "LTX2TwoStageHQPipeline",
|
"pipeline_class_name": "LTX2TwoStageHQPipeline",
|
||||||
"ltx2_two_stage_device_mode": "snapshot",
|
|
||||||
"performance_mode": "auto",
|
"performance_mode": "auto",
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(args.ltx2_two_stage_device_mode, "snapshot")
|
self.assertEqual(args.ltx2_two_stage_device_mode, "original")
|
||||||
self.assertTrue(args.dit_cpu_offload)
|
self.assertFalse(args.dit_cpu_offload)
|
||||||
self.assertTrue(args.layerwise_offload_components)
|
self.assertTrue(args.layerwise_offload_components)
|
||||||
self.assertFalse(args.text_encoder_cpu_offload)
|
self.assertFalse(args.text_encoder_cpu_offload)
|
||||||
self.assertFalse(args.image_encoder_cpu_offload)
|
self.assertFalse(args.image_encoder_cpu_offload)
|
||||||
@@ -1179,6 +1178,25 @@ class TestOffloadDefaults(unittest.TestCase):
|
|||||||
["text_encoder", "image_encoder", "vae"],
|
["text_encoder", "image_encoder", "vae"],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_ltx23_snapshot_device_mode_is_deprecated_alias_for_original(self):
|
||||||
|
args = self._from_dict_with_pipeline_config(
|
||||||
|
LTX2PipelineConfig(),
|
||||||
|
memory_gb=140,
|
||||||
|
available_memory_gb=134,
|
||||||
|
kwargs={
|
||||||
|
"model_path": "Lightricks/LTX-2.3",
|
||||||
|
"num_gpus": 2,
|
||||||
|
"pipeline_class_name": "LTX2TwoStagePipeline",
|
||||||
|
"ltx2_two_stage_device_mode": "snapshot",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(args.ltx2_two_stage_device_mode, "original")
|
||||||
|
self.assertEqual(
|
||||||
|
args.layerwise_offload_components,
|
||||||
|
["text_encoder", "image_encoder", "vae"],
|
||||||
|
)
|
||||||
|
|
||||||
def test_explicit_layerwise_components_preserved_in_ltx23_resident(self):
|
def test_explicit_layerwise_components_preserved_in_ltx23_resident(self):
|
||||||
args = self._from_dict_with_pipeline_config(
|
args = self._from_dict_with_pipeline_config(
|
||||||
LTX2PipelineConfig(),
|
LTX2PipelineConfig(),
|
||||||
@@ -1442,6 +1460,49 @@ class TestOffloadDefaults(unittest.TestCase):
|
|||||||
self.assertFalse(server_args.use_fsdp_inference)
|
self.assertFalse(server_args.use_fsdp_inference)
|
||||||
self.assertFalse(server_args.enable_cfg_parallel)
|
self.assertFalse(server_args.enable_cfg_parallel)
|
||||||
|
|
||||||
|
def test_ltx23_snapshot_device_mode_cli_alias_is_accepted(self):
|
||||||
|
parser = FlexibleArgumentParser()
|
||||||
|
ServerArgs.add_cli_args(parser)
|
||||||
|
argv = [
|
||||||
|
"--model-path",
|
||||||
|
"Lightricks/LTX-2.3",
|
||||||
|
"--pipeline-class-name",
|
||||||
|
"LTX2TwoStagePipeline",
|
||||||
|
"--ltx2-two-stage-device-mode",
|
||||||
|
"snapshot",
|
||||||
|
]
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(sys, "argv", ["sglang"] + argv),
|
||||||
|
patch.object(
|
||||||
|
PipelineConfig, "from_kwargs", return_value=LTX2PipelineConfig()
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.server_args.current_platform.is_cpu",
|
||||||
|
return_value=False,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.server_args.current_platform.is_mps",
|
||||||
|
return_value=False,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.server_args.current_platform.is_cuda",
|
||||||
|
return_value=True,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.server_args.current_platform.get_device_total_memory",
|
||||||
|
return_value=140 * 1024**3,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.multimodal_gen.runtime.server_args.current_platform.get_available_gpu_memory",
|
||||||
|
return_value=134,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
args, unknown_args = parser.parse_known_args(argv)
|
||||||
|
server_args = ServerArgs.from_cli_args(args, unknown_args)
|
||||||
|
|
||||||
|
self.assertEqual(server_args.ltx2_two_stage_device_mode, "original")
|
||||||
|
|
||||||
|
|
||||||
class TestFSDPShardConditions(unittest.TestCase):
|
class TestFSDPShardConditions(unittest.TestCase):
|
||||||
def test_helpers_match_only_direct_block_entries(self):
|
def test_helpers_match_only_direct_block_entries(self):
|
||||||
|
|||||||
Reference in New Issue
Block a user