From 05b3fd0f44cddc0376fea116fa809b2d6b92b734 Mon Sep 17 00:00:00 2001 From: Mick Date: Thu, 18 Jun 2026 10:20:21 +0800 Subject: [PATCH] [diffusion] chore: remove ltx2 snapshot mode (#28533) --- .../cookbook/diffusion/LTX/LTX2 & LTX2.3.mdx | 7 +- .../sglang-diffusion/compatibility_matrix.mdx | 8 +- .../src/snippets/diffusion/ltx-deployment.jsx | 4 +- .../memory_managers/component_manager.py | 173 +--------- .../runtime/pipelines/ltx_2_pipeline.py | 302 ++---------------- .../executors/pipeline_executor.py | 4 - .../multimodal_gen/runtime/server_args.py | 33 +- .../runtime/server_args_auto_tune.py | 5 - .../multimodal_gen/test/server/gpu_cases.py | 3 +- .../test/server/perf_baselines.json | 70 ++-- .../sglang/multimodal_gen/test/test_utils.py | 2 +- .../test/unit/test_server_args.py | 69 +++- 12 files changed, 162 insertions(+), 518 deletions(-) diff --git a/docs_new/cookbook/diffusion/LTX/LTX2 & LTX2.3.mdx b/docs_new/cookbook/diffusion/LTX/LTX2 & LTX2.3.mdx index c73169bc1..8e82deaa5 100644 --- a/docs_new/cookbook/diffusion/LTX/LTX2 & LTX2.3.mdx +++ b/docs_new/cookbook/diffusion/LTX/LTX2 & LTX2.3.mdx @@ -61,10 +61,11 @@ For two-stage pipelines, `--ltx2-two-stage-device-mode` controls transformer res | 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. | | `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: - `--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 | | --- | --- | --- | | 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.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. | | 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. diff --git a/docs_new/docs/sglang-diffusion/compatibility_matrix.mdx b/docs_new/docs/sglang-diffusion/compatibility_matrix.mdx index 1eeb840d1..6dea94417 100644 --- a/docs_new/docs/sglang-diffusion/compatibility_matrix.mdx +++ b/docs_new/docs/sglang-diffusion/compatibility_matrix.mdx @@ -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). - 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 ...`. -4. LTX-2 / LTX-2.3 two-stage also supports `--ltx2-two-stage-device-mode {original,snapshot,resident}`: - - `snapshot` is the default and recommended mode. - - `resident` usually provides the best latency/throughput but uses much more VRAM. +4. LTX-2 / LTX-2.3 two-stage also supports `--ltx2-two-stage-device-mode {original,resident}`: - `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 `nvidia/Cosmos3-Super` (32B). Both share the same pipeline; the only difference is transformer depth and width, picked up from diff --git a/docs_new/src/snippets/diffusion/ltx-deployment.jsx b/docs_new/src/snippets/diffusion/ltx-deployment.jsx index 2f15a7e3c..1b6d6ab01 100644 --- a/docs_new/src/snippets/diffusion/ltx-deployment.jsx +++ b/docs_new/src/snippets/diffusion/ltx-deployment.jsx @@ -7,7 +7,7 @@ export const LTXDeployment = () => { { id: 'h200', label: '1x H200', subtitle: 'resident', default: true }, { id: 'h200-2gpu', label: '2 GPUs', subtitle: 'CFG parallel', 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 }, ], }, @@ -121,7 +121,7 @@ export const LTXDeployment = () => { if (values.hardware === 'official') { return 'original'; } - return 'snapshot'; + return 'original'; }; const getParallelFlags = () => { diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py index c5e25d073..57d89315c 100644 --- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py +++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/component_manager.py @@ -1,8 +1,8 @@ -from collections.abc import Callable, Iterator +from collections.abc import Iterator from contextlib import contextmanager from dataclasses import dataclass from functools import lru_cache -from typing import Mapping, MutableMapping, Protocol, Sequence, TypeVar +from typing import Mapping, MutableMapping, Protocol, Sequence import torch import torch.nn as nn @@ -29,8 +29,6 @@ from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import DiffusionNvtx logger = init_logger(__name__) -_T = TypeVar("_T") - @dataclass(slots=True) class ComponentUse: @@ -75,8 +73,6 @@ class ResidencyState: # the ComponentUses of the preceding stages future_uses: tuple[ComponentUse, ...] = () batch_is_warmup: bool = False - manager_mode: str = "static" - trace_enabled: bool = False class ResidencyBatch(Protocol): @@ -142,7 +138,7 @@ class ComponentResidencyManager: ) -> None: self.pipeline = pipeline self.server_args = server_args - self.state = ResidencyState(trace_enabled=False) + self.state = ResidencyState() self._stage_names_by_id: dict[int, str] = {} self._stage_uses_by_index: list[tuple[ComponentUse, ...]] = [] self._ordered_uses: tuple[ComponentUse, ...] = () @@ -159,10 +155,6 @@ class ComponentResidencyManager: ) self._uses_seen: dict[str, ComponentUse] = {} - @property - def enabled(self) -> bool: - return True - def refresh_pipeline(self, pipeline: ComponentResidencyPipeline) -> None: custom_strategies = dict(pipeline.component_residency_strategies) if pipeline is not self.pipeline: @@ -186,13 +178,6 @@ class ComponentResidencyManager: self.strategy_for.cache_clear() 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( self, stages: Sequence[ComponentResidencyStage], @@ -201,29 +186,19 @@ class ComponentResidencyManager: ) -> None: """A hook called before processing an actual request""" self.refresh_server_args(server_args) - self.state = ResidencyState( - stages=stages, batch_is_warmup=batch.is_warmup, trace_enabled=False - ) + self.state = ResidencyState(stages=stages, batch_is_warmup=batch.is_warmup) self._active_use = None self._active_use_module = None self._disable_active_nvtx() self._current_use_index = -1 self._prefetched_use_keys.clear() self._uses_seen.clear() - if self.enabled: - self._stage_uses_by_index = [ - tuple(stage.component_uses(server_args, self.stage_name(stage))) - for stage in stages - ] - self._ordered_uses = tuple( - 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)}", + self._stage_uses_by_index = [ + tuple(stage.component_uses(server_args, self.stage_name(stage))) + for stage in stages + ] + self._ordered_uses = tuple( + use for uses in self._stage_uses_by_index for use in uses ) def before_stage( @@ -234,25 +209,10 @@ class ComponentResidencyManager: server_args: ServerArgs, ) -> None: """called after stage starts""" - if not self.enabled: - return # update state before entering the stage self.state.stage_index = stage_index self.state.stage_name = self.stage_name(stage) 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: """Begin one sequential component use interval. this is idempotent @@ -321,26 +281,8 @@ class ComponentResidencyManager: finally: 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: """Prepare a shared component and wait without making it the active use.""" - if not self.enabled: - return self._prepare_forward_use(use, module=module) def remove_nvtx_hooks_for_module(self, module: nn.Module | None) -> None: @@ -357,19 +299,6 @@ class ComponentResidencyManager: hooks.remove_hooks() 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: """Finish the currently active sequential use, if any.""" 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.""" module = module or self.get_module(use.component_name) if module is None: - self._trace("skip_missing", use) return None strategy = self.strategy_for(use.component_name, module) self._uses_seen[use.component_name] = use self.state.current_use = use - self._trace("prepare", use, strategy, module) strategy.prepare_for_use(module, use, self.state) - self._trace("wait", use, strategy, module) strategy.wait_for_use(module, use, self.state) return module @@ -466,33 +392,24 @@ class ComponentResidencyManager: def _prefetch_use(self, use: ComponentUse) -> None: """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. - - manager will perform the prefetch at some checkpoints, if necessary + This is called for memory-intensive future uses where H2D placement can + overlap with the current stage. """ if not use.allow_prefetch: return module = self.get_module(use.component_name) if module is None: - self._trace("skip_missing", use) return strategy = self.strategy_for(use.component_name, module) if isinstance(strategy, VanillaD2HStrategy) and self._active_use is not None: # Avoid making two vanilla-offloaded heavy components resident before # a budget-aware planner can prove the overlap is safe. - self._trace("prefetch_skip_active_vanilla", use, strategy, module) return self._uses_seen[use.component_name] = use - self._trace("prefetch", use, strategy, module) if strategy.prefetch_for_use(module, use, self.state): 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( self, use: ComponentUse, @@ -503,28 +420,18 @@ class ComponentResidencyManager: """finish a specific use by keeping them resident or call finish_use hook""" module = module or self.get_module(use.component_name) if module is None: - self._trace("skip_missing", use) return should_keep = ( keep_on_warmup and self.state.batch_is_warmup ) or self._should_keep_after_use(use) if should_keep: - self._trace( - "keep", - use, - self.strategy_for(use.component_name, module), - module, - ) return strategy = self.strategy_for(use.component_name, module) - self._trace("finish", use, strategy, module) was_on_cuda = self._module_on_cuda(module) strategy.finish_use(module, use, self.state) self._empty_cache_after_large_release(use, strategy, module, was_on_cuda) 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. self.finish_active_use(prefetch_next=False) # 2. Pick components that should be ready for the next request. @@ -535,36 +442,19 @@ class ComponentResidencyManager: if module is None: continue 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 preferred = component_name in preferred_uses 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 strategy = self.strategy_for(component_name, module) if preferred and not self.state.batch_is_warmup: - self._trace("request_prefetch", use, strategy, module) strategy.prepare_after_request(module, use, self.state) else: - action = "request_resident" if preferred else "request_finish" - self._trace(action, use, strategy, module) was_on_cuda = self._module_on_cuda(module) strategy.finish_request(module, use, self.state, preferred=preferred) self._empty_cache_after_large_release( use, strategy, module, was_on_cuda ) - self._trace("request_end") def stage_name(self, stage: ComponentResidencyStage) -> str: return self._stage_names_by_id.get(id(stage), stage.__class__.__name__) @@ -593,12 +483,6 @@ class ComponentResidencyManager: 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: next_index = stage_index + 1 if next_index < 0 or next_index >= len(self.state.stages): @@ -629,7 +513,7 @@ class ComponentResidencyManager: continue if self._use_key(use) in self._prefetched_use_keys: return - self.prefetch_use(use) + self._prefetch_use(use) return 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]: 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: if module is None: return None @@ -743,7 +597,6 @@ class ComponentResidencyManager: if not torch.get_device_module().is_available(): return torch.get_device_module().empty_cache() - self._trace("empty_cache", use, strategy, module, detail="after_release") _GLOBAL_COMPONENT_RESIDENCY_MANAGER: ComponentResidencyManager | None = None diff --git a/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py index 463acad85..5acd524bb 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py @@ -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 ( PipelineComponentLoader, ) -from sglang.multimodal_gen.runtime.loader.utils import BYTES_PER_GB from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( ComponentResidencyStrategy, ComponentUse, 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 ( ComposedPipelineBase, ) @@ -51,12 +46,11 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.l LTX2TextConnectorStage, LTX2UpsampleStage, ) -from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.server_args import ( - LTX2_RESIDENT_AUTO_ENABLE_MEM_GB, + LTX2_TWO_STAGE_DEVICE_MODE_CHOICES, 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 logger = init_logger(__name__) @@ -413,10 +407,6 @@ class LTX2TwoStageResidencyStrategy(ComponentResidencyStrategy): 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: - return SnapshotModuleResidency.is_on_gpu(module) - class LTX2OriginalResidencyStrategy(LTX2TwoStageResidencyStrategy): pass @@ -442,265 +432,6 @@ class LTX2ResidentResidencyStrategy(LTX2TwoStageResidencyStrategy): 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: """ LTX-2.3 two-stage residency controller. @@ -709,11 +440,10 @@ class LTX2TwoStageResidencyController: Modes: - 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. """ - VALID_MODES = ("original", "snapshot", "resident") + VALID_MODES = ("original", "resident") def __init__(self, pipeline: "LTX2TwoStagePipeline", server_args: ServerArgs): self.pipeline = pipeline @@ -727,17 +457,21 @@ class LTX2TwoStageResidencyController: mode = server_args.ltx2_two_stage_device_mode if mode is None: 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: raise ValueError( 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 def _build_strategy(self) -> LTX2TwoStageResidencyStrategy: - if self.mode == "snapshot": - return LTX2SnapshotResidencyStrategy(self) if self.mode == "resident": return LTX2ResidentResidencyStrategy(self) return LTX2OriginalResidencyStrategy(self) @@ -750,11 +484,11 @@ class LTX2TwoStageResidencyController: 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 + We only enable this optimization for resident native LTX-2.3 two-stage + and when users did not explicitly provide a stage-1 LoRA path """ return ( - self.mode != "original" + self.mode == "resident" and self.pipeline._should_merge_stage2_distilled_lora(self.server_args) and self.pipeline._stage1_lora_path is None ) @@ -868,7 +602,7 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline): 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. + # precision drift against original two-stage behavior. self.set_lora( lora_nickname="ltx2_stage2_distilled", lora_path=self._distilled_lora_path, @@ -878,9 +612,9 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline): ) def should_skip_ltx2_lora_switch_stage(self) -> bool: - return self._use_premerged_stage2_transformer and self._ltx2_residency.mode in ( - "snapshot", - "resident", + return ( + self._use_premerged_stage2_transformer + and self._ltx2_residency.mode == "resident" ) def _get_stage_distilled_lora_strength( diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py index be94ccde5..f8ecad2b5 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py @@ -71,9 +71,6 @@ class PipelineExecutor(ABC): 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: self.component_residency_manager.finish_request() @@ -118,7 +115,6 @@ class PipelineExecutor(ABC): payload = self.run_stage_with_context( stage, payload, server_args, run_stage ) - self.after_stage(stage_index) return payload @staticmethod diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index b5bfde175..56b8e8afe 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -66,7 +66,8 @@ from sglang.multimodal_gen.utils import ( 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") # H200-class GPUs (>=130 GiB total) can usually keep both LTX2 DiTs resident. 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: return None 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 @@ -497,7 +505,7 @@ class ServerArgs(DisaggServerArgsMixin): 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}." + f"Expected one of {LTX2_TWO_STAGE_DEVICE_MODE_CHOICES}." ) self.ltx2_two_stage_device_mode = mode @@ -505,9 +513,9 @@ class ServerArgs(DisaggServerArgsMixin): 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" + "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_total_memory_gb = ( @@ -525,11 +533,11 @@ class ServerArgs(DisaggServerArgsMixin): return "resident" 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_total_memory_gb, ) - return "snapshot" + return "original" def _is_ltx23_two_stage_pipeline(self) -> bool: 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) ) - 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: if ( self.ltx2_two_stage_device_mode != "resident" @@ -1410,14 +1412,15 @@ class ServerArgs(DisaggServerArgsMixin): parser.add_argument( "--ltx2-two-stage-device-mode", type=str, - choices=LTX2_TWO_STAGE_DEVICE_MODES, + choices=LTX2_TWO_STAGE_DEVICE_MODE_CHOICES, 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." + "'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( diff --git a/python/sglang/multimodal_gen/runtime/server_args_auto_tune.py b/python/sglang/multimodal_gen/runtime/server_args_auto_tune.py index dfa4fec50..bceb24fe4 100644 --- a/python/sglang/multimodal_gen/runtime/server_args_auto_tune.py +++ b/python/sglang/multimodal_gen/runtime/server_args_auto_tune.py @@ -107,11 +107,6 @@ class ServerArgsAutoTuner: 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 ( args.dit_cpu_offload and "dit" in components diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py index bd6b2c4b4..828b8a647 100644 --- a/python/sglang/multimodal_gen/test/server/gpu_cases.py +++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py @@ -429,11 +429,10 @@ ONE_GPU_CASES: list[DiffusionTestCase] = [ DiffusionServerArgs( model_path="Lightricks/LTX-2.3", extras=[ - "--pipeline-class-name LTX2TwoStageHQPipeline --ltx2-two-stage-device-mode snapshot" + "--pipeline-class-name LTX2TwoStageHQPipeline --ltx2-two-stage-device-mode original" ], env_vars={ "PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True", - "SGLANG_LTX2_SNAPSHOT_RELEASE_EMPTY_CACHE": "true", }, ), run_component_accuracy_check=False, diff --git a/python/sglang/multimodal_gen/test/server/perf_baselines.json b/python/sglang/multimodal_gen/test/server/perf_baselines.json index 677c90dda..7d49fc138 100644 --- a/python/sglang/multimodal_gen/test/server/perf_baselines.json +++ b/python/sglang/multimodal_gen/test/server/perf_baselines.json @@ -2681,44 +2681,46 @@ }, "ltx_2_3_hq_pipeline": { "stages_ms": { - "InputValidationStage": 11.84, - "TextEncodingStage": 446.77, - "LTX2TextConnectorStage": 32.96, - "LTX2HalveResolutionStage": 0.07, - "LTX2LoRASwitchStage": 0.03, - "LTX2SigmaPreparationStage": 0.38, - "TimestepPreparationStage": 24.67, - "LTX2AVLatentPreparationStage": 0.16, - "LTX2ImageEncodingStage": 1133.15, - "LTX2AVDenoisingStage": 19539.95, - "LTX2UpsampleStage": 22.73, - "LTX2RefinementStage": 3059.95, - "LTX2AVDecodingStage": 1744.05, + "InputValidationStage": 4.63, + "TextEncodingStage": 401.75, + "LTX2TextConnectorStage": 27.48, + "LTX2HalveResolutionStage": 0.04, + "LTX2LoRASwitchStage": 180.0, + "LTX2SigmaPreparationStage": 0.26, + "TimestepPreparationStage": 14.45, + "LTX2AVLatentPreparationStage": 0.13, + "LTX2ImageEncodingStage": 57.62, + "LTX2AVDenoisingStage": 12162.51, + "LTX2UpsampleStage": 11.04, + "ltx2_lora_switch_stage2": 9155.58, + "ltx2_image_encoding_stage2": 64.54, + "LTX2RefinementStage": 3484.86, + "LTX2AVDecodingStage": 1054.91, "per_frame_generation": null }, "denoise_step_ms": { - "0": 1677.39, - "1": 778.3, - "2": 740.93, - "3": 1177.58, - "4": 751.77, - "5": 744.01, - "6": 742.02, - "7": 1264.88, - "8": 744.11, - "9": 739.84, - "10": 1209.91, - "11": 739.62, - "12": 736.7, - "13": 733.55, - "14": 590.44, - "15": 903.04, - "16": 876.98, - "17": 879.42 + "0": 785.16, + "1": 740.86, + "2": 740.7, + "3": 742.77, + "4": 772.73, + "5": 744.43, + "6": 743.14, + "7": 743.2, + "8": 743.76, + "9": 742.57, + "10": 741.31, + "11": 744.06, + "12": 760.17, + "13": 740.69, + "14": 369.53, + "15": 1193.58, + "16": 1144.87, + "17": 1143.19 }, - "expected_e2e_ms": 27309.99, - "expected_avg_denoise_ms": 1159.69, - "expected_median_denoise_ms": 1199, + "expected_e2e_ms": 26673.22, + "expected_avg_denoise_ms": 868.06, + "expected_median_denoise_ms": 747.62, "estimated_full_test_time_s": 363.2 }, "qwen_image_t2i_cache_dit_scm_config_diffusers_1gpu": { diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index 5db33bc58..ee8a79f96 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -34,7 +34,7 @@ if TYPE_CHECKING: logger = init_logger(__name__) -SGL_TEST_FILES_CI_DATA_REVISION = "f0fd96eab85baed5256d142c7659e0634e7e4410" +SGL_TEST_FILES_CI_DATA_REVISION = "51a6a6cd592983e1b8dadc9d7981fac63cd02800" if current_platform.is_npu(): SGL_TEST_FILES_CI_DATA_REVISION = "670d66a8a290b62c0c3c077b3e9b0f4a4d9a44e7" diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py index 670b01e3c..e0a6de9d4 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -882,7 +882,7 @@ class TestOffloadDefaults(unittest.TestCase): ["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, ): args = self._from_dict_with_pipeline_config( @@ -891,13 +891,12 @@ class TestOffloadDefaults(unittest.TestCase): kwargs={ "model_path": "Lightricks/LTX-2.3", "pipeline_class_name": "LTX2TwoStageHQPipeline", - "ltx2_two_stage_device_mode": "snapshot", "performance_mode": "auto", }, ) - self.assertEqual(args.ltx2_two_stage_device_mode, "snapshot") - self.assertTrue(args.dit_cpu_offload) + self.assertEqual(args.ltx2_two_stage_device_mode, "original") + self.assertFalse(args.dit_cpu_offload) self.assertTrue(args.layerwise_offload_components) self.assertFalse(args.text_encoder_cpu_offload) self.assertFalse(args.image_encoder_cpu_offload) @@ -1179,6 +1178,25 @@ class TestOffloadDefaults(unittest.TestCase): ["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): args = self._from_dict_with_pipeline_config( LTX2PipelineConfig(), @@ -1442,6 +1460,49 @@ class TestOffloadDefaults(unittest.TestCase): self.assertFalse(server_args.use_fsdp_inference) 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): def test_helpers_match_only_direct_block_entries(self):