[diffusion] fix: fix LTX2 resident defaults and stage profiling (#25596)
This commit is contained in:
@@ -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(),
|
||||
|
||||
Reference in New Issue
Block a user