[diffusion] fix: fix LTX2 resident defaults and stage profiling (#25596)

This commit is contained in:
Mick
2026-05-19 10:41:28 +08:00
committed by GitHub
parent 87c3c96bc8
commit a7b3ced334
7 changed files with 170 additions and 6 deletions
@@ -718,6 +718,8 @@ class DiffusersPipeline(ComposedPipelineBase):
if stage_name in self._stage_name_mapping:
raise ValueError(f"Duplicate stage name detected: {stage_name}")
stage.set_registered_stage_name(stage_name)
stage.set_profile_stage_name(self._profile_stage_name(stage, stage_name))
self._stages.append(stage)
self._stage_name_mapping[stage_name] = stage
return self
@@ -502,6 +502,12 @@ class ComposedPipelineBase(ABC):
def _infer_stage_name(stage: PipelineStage) -> str:
return stage.__class__.__name__
def _profile_stage_name(self, stage: PipelineStage, stage_name: str) -> str:
class_name = stage.__class__.__name__
if any(existing.__class__.__name__ == class_name for existing in self._stages):
return stage_name
return class_name
def add_stage(
self, stage: PipelineStage, stage_name: str | None = None
) -> "ComposedPipelineBase":
@@ -526,6 +532,8 @@ class ComposedPipelineBase(ABC):
if stage_name in self._stage_name_mapping:
raise ValueError(f"Duplicate stage name detected: {stage_name}")
stage.set_registered_stage_name(stage_name)
stage.set_profile_stage_name(self._profile_stage_name(stage, stage_name))
self._stages.append(stage)
self._stage_name_mapping[stage_name] = stage
return self
@@ -60,6 +60,8 @@ class PipelineStage(StageDedupMixin, ABC):
def __init__(self):
self.server_args = get_global_server_args()
self._component_residency_manager = None
self._registered_stage_name: str | None = None
self._profile_stage_name: str | None = None
def log_info(self, msg, *args):
"""Logs an informational message with the stage name as a prefix."""
@@ -115,14 +117,29 @@ class PipelineStage(StageDedupMixin, ABC):
def set_component_residency_manager(self, manager) -> None:
self._component_residency_manager = manager
def set_registered_stage_name(self, stage_name: str) -> None:
self._registered_stage_name = stage_name
def set_profile_stage_name(self, stage_name: str) -> None:
self._profile_stage_name = stage_name
def _component_stage_name(self, stage_name: str | None = None) -> str:
return stage_name or self.__class__.__name__
return (
stage_name
or getattr(self, "_registered_stage_name", None)
or self.__class__.__name__
)
def _active_component_stage_name(self) -> str:
manager = self._component_residency_manager
if manager is not None and manager.state.stage_name is not None:
return manager.state.stage_name
return self.__class__.__name__
manager = getattr(self, "_component_residency_manager", None)
manager_state = getattr(manager, "state", None)
manager_stage_name = getattr(manager_state, "stage_name", None)
if manager_stage_name is not None:
return manager_stage_name
return self._component_stage_name()
def _active_profile_stage_name(self) -> str:
return getattr(self, "_profile_stage_name", None) or self.__class__.__name__
def _finish_active_component_use(self) -> None:
if self._component_residency_manager is not None:
@@ -267,7 +284,7 @@ class PipelineStage(StageDedupMixin, ABC):
Returns:
The updated batch information after this stage's processing.
"""
stage_name = self.__class__.__name__
stage_name = self._active_profile_stage_name()
# Check if verification is enabled (simple approach for prototype)
# Pre-execution input verification
@@ -505,6 +505,18 @@ class ServerArgs(DisaggArgsMixin):
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"
or not self._is_ltx23_two_stage_pipeline()
or not current_platform.is_cuda()
):
return False
return (
current_platform.get_device_total_memory() / BYTES_PER_GB
>= LTX2_RESIDENT_AUTO_ENABLE_MEM_GB
)
def _adjust_attention_backend(self):
if self.attention_backend in ["fa3", "fa4"]:
self.attention_backend = "fa"
@@ -148,6 +148,38 @@ class ServerArgsAutoTuner:
", ".join(changed),
)
self._maybe_keep_ltx23_resident_aux_components_resident()
def _maybe_keep_ltx23_resident_aux_components_resident(self) -> None:
args = self.server_args
if not args._uses_ltx23_high_memory_resident_two_stage_mode():
return
changed: list[str] = []
if (
args.layerwise_offload_components is not None
and not args.is_arg_explicitly_set("layerwise_offload_components")
):
args.layerwise_offload_components = None
changed.append("layerwise_offload_components=None")
# high-memory resident mode keeps both DiTs on GPU; unset auxiliary
# placement should stay resident instead of using default layerwise
for arg_name in (
"text_encoder_cpu_offload",
"image_encoder_cpu_offload",
"vae_cpu_offload",
):
if getattr(args, arg_name) and not args.is_arg_explicitly_set(arg_name):
setattr(args, arg_name, False)
changed.append(f"{arg_name}=False")
if changed:
logger.info(
"Keeping LTX2 high-memory two-stage auxiliary components resident: %s",
", ".join(changed),
)
def maybe_adjust_auto_fsdp_with_offload_enabled(self) -> None:
args = self.server_args
if (
@@ -0,0 +1,39 @@
import unittest
from types import SimpleNamespace
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
class NamedNoOpStage(PipelineStage):
def __init__(self):
self.server_args = SimpleNamespace(comfyui_mode=True)
def forward(self, batch: Req, server_args) -> Req:
return batch
class TestPipelineStageProfiling(unittest.TestCase):
def test_profiler_uses_profile_stage_name(self):
stage = NamedNoOpStage()
stage.set_profile_stage_name("profile_stage")
batch = Req(perf_dump_path="/tmp/unused_perf.json")
stage(batch, SimpleNamespace())
self.assertIn("profile_stage", batch.metrics.stages)
self.assertNotIn("NamedNoOpStage", batch.metrics.stages)
def test_registered_stage_name_does_not_change_profile_name(self):
stage = NamedNoOpStage()
stage.set_registered_stage_name("prompt_encoding_stage_primary")
batch = Req(perf_dump_path="/tmp/unused_perf.json")
stage(batch, SimpleNamespace())
self.assertIn("NamedNoOpStage", batch.metrics.stages)
self.assertNotIn("prompt_encoding_stage_primary", batch.metrics.stages)
if __name__ == "__main__":
unittest.main()
@@ -603,6 +603,60 @@ class TestOffloadDefaults(unittest.TestCase):
["text_encoder", "image_encoder", "vae"],
)
def test_auto_high_memory_ltx23_resident_keeps_aux_components_resident(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",
},
)
self.assertEqual(args.ltx2_two_stage_device_mode, "resident")
self.assertFalse(args.use_fsdp_inference)
self.assertFalse(args.dit_cpu_offload)
self.assertFalse(args.text_encoder_cpu_offload)
self.assertFalse(args.image_encoder_cpu_offload)
self.assertFalse(args.vae_cpu_offload)
self.assertIsNone(args.layerwise_offload_components)
def test_auto_high_memory_ltx23_original_keeps_default_layerwise_components(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": "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(),
memory_gb=140,
available_memory_gb=134,
kwargs={
"model_path": "Lightricks/LTX-2.3",
"num_gpus": 2,
"pipeline_class_name": "LTX2TwoStagePipeline",
"layerwise_offload_components": ["text_encoder"],
},
)
self.assertEqual(args.ltx2_two_stage_device_mode, "resident")
self.assertEqual(args.layerwise_offload_components, ["text_encoder"])
def test_auto_multi_gpu_qwen_replaces_text_encoder_offload_with_cfg(self):
args = self._from_dict_with_pipeline_config(
QwenImagePipelineConfig(),