diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising.py index c8baaa338..8b3ebf601 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/denoising.py @@ -101,6 +101,48 @@ class LTX2GuidancePassSpec: disable_v2a_cross_attn: bool = False +def _prepare_ltx2_rope_coords_for_bcg( + *, + enabled: bool, + current_model, + latent_model_input: torch.Tensor, + audio_latent_model_input: torch.Tensor, + video_coords: torch.Tensor | None, + audio_coords: torch.Tensor | None, + num_frames: int, + height: int, + width: int, + audio_num_frames: int, + fps: int, +) -> tuple[torch.Tensor | None, torch.Tensor | None]: + """Build missing LTX-2 RoPE coordinates before CUDA graph capture. + + The legacy LTX-2.3 one-stage path normally builds these tensors inside the + model forward. That performs an unpinned host-to-device copy, which CUDA + graph capture rejects. Preparing the same coordinates before entering the + graph preserves the eager values and makes both legacy and two-stage LTX + forwards capturable. + """ + if not enabled: + return video_coords, audio_coords + if video_coords is None: + video_coords = current_model.rope.prepare_video_coords( + batch_size=int(latent_model_input.shape[0]), + num_frames=num_frames, + height=height, + width=width, + device=latent_model_input.device, + fps=fps, + ) + if audio_coords is None: + audio_coords = current_model.audio_rope.prepare_audio_coords( + batch_size=int(audio_latent_model_input.shape[0]), + num_frames=audio_num_frames, + device=audio_latent_model_input.device, + ) + return video_coords, audio_coords + + class LTX2DenoisingStage(DenoisingStage): """ LTX-2 specific denoising stage that handles joint video and audio generation. @@ -1169,28 +1211,19 @@ class LTX2DenoisingStage(DenoisingStage): audio_latent_model_input, num_frames=audio_num_frames_latent, ) - if server_args.enable_breakable_cuda_graph: - # The in-model RoPE coordinate construction builds host - # tensors (torch.tensor(list, device=cuda)), which is an - # unpinned H2D copy and therefore illegal inside CUDA graph - # capture. Build the coords outside the captured region with - # the exact same rope helpers (start_frame=0 == the sp<=1 - # in-model path), so values are bit-identical. - if video_coords is None: - video_coords = step.current_model.rope.prepare_video_coords( - batch_size=int(latent_model_input.shape[0]), - num_frames=ctx.latent_num_frames_for_model, - height=ctx.latent_height, - width=ctx.latent_width, - device=latent_model_input.device, - fps=batch.fps, - ) - if audio_coords is None: - audio_coords = step.current_model.audio_rope.prepare_audio_coords( - batch_size=int(audio_latent_model_input.shape[0]), - num_frames=audio_num_frames_latent, - device=audio_latent_model_input.device, - ) + video_coords, audio_coords = _prepare_ltx2_rope_coords_for_bcg( + enabled=server_args.enable_breakable_cuda_graph, + current_model=step.current_model, + latent_model_input=latent_model_input, + audio_latent_model_input=audio_latent_model_input, + video_coords=video_coords, + audio_coords=audio_coords, + num_frames=ctx.latent_num_frames_for_model, + height=ctx.latent_height, + width=ctx.latent_width, + audio_num_frames=audio_num_frames_latent, + fps=batch.fps, + ) batch_size = int(latent_model_input.shape[0]) use_raw_sigma_timestep = ctx.use_ltx23_hq_timestep_semantics diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index 043572260..f983f7455 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -158,7 +158,9 @@ BREAKABLE_CUDA_GRAPH_SUPPORTED_MODEL_IDS = frozenset( "ideogram-ai/ideogram-4-fp8", "ideogram-ai/ideogram-4-nf4", "lightricks/ltx-2", + "lightricks/ltx-2.3", "ltx-2", + "ltx-2.3", "minimax-h3", "minimaxai/minimax-h3", "qwen/qwen-image", @@ -178,6 +180,7 @@ BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS = frozenset( "GlmImagePipelineConfig", "Ideogram4PipelineConfig", "LTX2PipelineConfig", + "LTX23PipelineConfig", "MiniMaxH3PipelineConfig", "QwenImagePipelineConfig", "SanaPipelineConfig", diff --git a/python/sglang/multimodal_gen/test/unit/test_diffusion_bcg_padding.py b/python/sglang/multimodal_gen/test/unit/test_diffusion_bcg_padding.py index 07ab70cb0..8ce7abb43 100644 --- a/python/sglang/multimodal_gen/test/unit/test_diffusion_bcg_padding.py +++ b/python/sglang/multimodal_gen/test/unit/test_diffusion_bcg_padding.py @@ -388,6 +388,20 @@ class TestDiffusionBCGPadding(unittest.TestCase): ): self.assertIn(config_name, BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS) + def test_ltx_models_are_registered_as_bcg_supported(self): + for model_id in ( + "lightricks/ltx-2", + "lightricks/ltx-2.3", + ): + self.assertIn(model_id, BREAKABLE_CUDA_GRAPH_SUPPORTED_MODEL_IDS) + + self.assertIn( + "LTX2PipelineConfig", BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS + ) + self.assertIn( + "LTX23PipelineConfig", BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS + ) + def test_dynamic_varlen_mask_meta_rebuilds_once_per_replay_token(self): builder = DynamicVarlenMaskMeta() mask = torch.tensor([[True, True, False, False]]) diff --git a/python/sglang/multimodal_gen/test/unit/test_ltx2_bcg_coords.py b/python/sglang/multimodal_gen/test/unit/test_ltx2_bcg_coords.py new file mode 100644 index 000000000..b433b8347 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_ltx2_bcg_coords.py @@ -0,0 +1,87 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import Mock + +import torch + +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ltx_2.denoising import ( + _prepare_ltx2_rope_coords_for_bcg, +) +from sglang.test.test_utils import CustomTestCase + + +class TestLTX2BCGCoords(CustomTestCase): + def setUp(self): + self.video_result = torch.ones(1, 4) + self.audio_result = torch.ones(1, 2) + self.model = SimpleNamespace( + rope=SimpleNamespace( + prepare_video_coords=Mock(return_value=self.video_result) + ), + audio_rope=SimpleNamespace( + prepare_audio_coords=Mock(return_value=self.audio_result) + ), + ) + self.video_latents = torch.zeros(2, 8, 16) + self.audio_latents = torch.zeros(2, 4, 16) + + def _call(self, *, enabled, video_coords=None, audio_coords=None): + return _prepare_ltx2_rope_coords_for_bcg( + enabled=enabled, + current_model=self.model, + latent_model_input=self.video_latents, + audio_latent_model_input=self.audio_latents, + video_coords=video_coords, + audio_coords=audio_coords, + num_frames=16, + height=32, + width=48, + audio_num_frames=126, + fps=24, + ) + + def test_prepares_missing_coords_before_bcg_capture(self): + video, audio = self._call(enabled=True) + + self.assertIs(video, self.video_result) + self.assertIs(audio, self.audio_result) + self.model.rope.prepare_video_coords.assert_called_once_with( + batch_size=2, + num_frames=16, + height=32, + width=48, + device=self.video_latents.device, + fps=24, + ) + self.model.audio_rope.prepare_audio_coords.assert_called_once_with( + batch_size=2, + num_frames=126, + device=self.audio_latents.device, + ) + + def test_disabled_bcg_keeps_legacy_none_coords(self): + video, audio = self._call(enabled=False) + + self.assertIsNone(video) + self.assertIsNone(audio) + self.model.rope.prepare_video_coords.assert_not_called() + self.model.audio_rope.prepare_audio_coords.assert_not_called() + + def test_existing_parallel_coords_are_preserved(self): + existing_video = torch.zeros(3) + existing_audio = torch.zeros(5) + + video, audio = self._call( + enabled=True, + video_coords=existing_video, + audio_coords=existing_audio, + ) + + self.assertIs(video, existing_video) + self.assertIs(audio, existing_audio) + self.model.rope.prepare_video_coords.assert_not_called() + self.model.audio_rope.prepare_audio_coords.assert_not_called() + + +if __name__ == "__main__": + unittest.main()