Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -274,6 +274,16 @@ def transformer_class_name_matches(current_model: Any, needle: str) -> bool:
_PROMPT_PADDERS: list[tuple[Callable[[Any, dict], bool], PromptPadder]] = []


def _cosmos3_padder(
call_kwargs: dict, _current_model: Any, buckets: tuple[int, ...]
) -> dict:
out = pad_masked_prompt_kwargs(call_kwargs, buckets)
text_mask = first_tensor(out.get("text_mask"))
if text_mask is not None and out.get("max_text_seq_len") is not None:
out["max_text_seq_len"] = int(text_mask.shape[1])
return out


def register_prompt_padder(
predicate: Callable[[Any, dict], bool], padder: PromptPadder
) -> None:
Expand Down Expand Up @@ -305,3 +315,8 @@ def _ensure_model_padders_registered() -> None:
qwen_image,
zimage,
)

register_prompt_padder(
lambda model, _kwargs: transformer_class_name_matches(model, "cosmos3"),
_cosmos3_padder,
)
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,12 @@
import torch.nn as nn

from sglang.multimodal_gen.configs.sample.sampling_params import DataType
from sglang.multimodal_gen.runtime.breakable_cuda_graph import (
prompt_padding as bcg_utils,
)
from sglang.multimodal_gen.runtime.breakable_cuda_graph.runner import (
DiffusionBreakableCudaGraphRunner,
)
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.distributed.communication_op import (
cfg_model_parallel_all_reduce,
Expand Down Expand Up @@ -799,6 +805,12 @@ def __init__(self, transformer, scheduler, server_args: ServerArgs | None = None
self.scheduler = scheduler
self.server_args = server_args
self._logged_parallel_config = False
self._bcg_runner = None

if server_args is not None and server_args.enable_breakable_cuda_graph:
self._bcg_runner = DiffusionBreakableCudaGraphRunner(
transformer, get_local_torch_device()
)
self._logged_cfg_split = False

# Apply torch.compile if enabled
Expand Down Expand Up @@ -895,6 +907,7 @@ def _run_transformer(
action_noisy_mask: torch.Tensor | None = None,
action_fps: float | None = None,
action_start_frame_offset: int = 1,
forward_batch: Req | None = None,
) -> torch.Tensor | tuple[torch.Tensor, ...]:
"""Run transformer forward pass.

Expand All @@ -911,24 +924,56 @@ def _run_transformer(
"""
if current_timestep is None:
current_timestep = int(timestep.flatten()[0].item())
with set_forward_context(current_timestep=current_timestep, attn_metadata=None):
return self.transformer(
hidden_states=latents,
encoder_hidden_states=None, # Not used by Cosmos3
timestep=timestep,
text_ids=text_ids,
text_mask=text_mask,
fps=fps,
cache_key=cache_key,
noisy_frame_mask=noisy_frame_mask,
max_text_seq_len=max_text_seq_len,
sound_latents=sound_latents,
action_latents=action_latents,
action_domain_ids=action_domain_ids,
action_noisy_mask=action_noisy_mask,
action_fps=action_fps,
action_start_frame_offset=action_start_frame_offset,
)
call_kwargs = dict(
hidden_states=latents,
encoder_hidden_states=None,
timestep=timestep,
text_ids=text_ids,
text_mask=text_mask,
fps=fps,
cache_key=cache_key,
noisy_frame_mask=noisy_frame_mask,
max_text_seq_len=max_text_seq_len,
sound_latents=sound_latents,
action_latents=action_latents,
action_domain_ids=action_domain_ids,
action_noisy_mask=action_noisy_mask,
action_fps=action_fps,
action_start_frame_offset=action_start_frame_offset,
)
with set_forward_context(
current_timestep=current_timestep,
attn_metadata=None,
forward_batch=(
forward_batch
if forward_batch is not None
else getattr(self, "_current_forward_batch", None)
),
):
if self._bcg_runner is not None:
buckets = self.server_args.resolved_bcg_text_buckets()
if bool(
getattr(
(
forward_batch
if forward_batch is not None
else getattr(self, "_current_forward_batch", None)
),
"is_warmup",
False,
)
):
for bucket in buckets:
self._bcg_runner.capture(
**bcg_utils.select_prompt_padder(
self.transformer, call_kwargs
)(call_kwargs, self.transformer, (bucket,))
)
call_kwargs = bcg_utils.select_prompt_padder(
self.transformer, call_kwargs
)(call_kwargs, self.transformer, buckets)
transformer = self._bcg_runner or self.transformer
return transformer(**call_kwargs)

def _manage_device_placement(self, server_args: ServerArgs):
"""Move transformer to GPU if CPU offload is enabled."""
Expand Down Expand Up @@ -961,6 +1006,7 @@ def _cfg_active_at(t: torch.Tensor, interval: tuple[float, float] | None) -> boo

def forward(self, batch: Req, server_args: ServerArgs) -> Req:
"""Run the denoising loop with CFG and optional I2V conditioning."""
self._current_forward_batch = batch
self._manage_device_placement(server_args)

latents = batch.latents
Expand Down
16 changes: 12 additions & 4 deletions python/sglang/multimodal_gen/runtime/server_args/server_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -156,6 +156,7 @@ def choices(cls) -> list[str]:
"ltx-2",
"minimax-h3",
"minimaxai/minimax-h3",
"nvidia/cosmos3-nano",
"qwen/qwen-image",
"qwen/qwen-image-2512",
"qwen-image",
Expand All @@ -175,6 +176,7 @@ def choices(cls) -> list[str]:
"LTX2PipelineConfig",
"MiniMaxH3PipelineConfig",
"QwenImagePipelineConfig",
"Cosmos3Config",
"SanaPipelineConfig",
"ZImagePipelineConfig",
}
Expand Down Expand Up @@ -597,7 +599,7 @@ def _adjust_breakable_cuda_graph_support(self):
pipeline_config_name in BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS
and self._is_breakable_cuda_graph_supported_model()
):
if not self.warmup_resolutions:
if not self.warmup_resolutions and self.warmup_mode != "request":
self._default_bcg_warmup_resolution()
return

Expand Down Expand Up @@ -1032,9 +1034,15 @@ def _adjust_warmup(self):
if self.warmup_resolutions is not None and self.warmup_mode in (None, "off"):
self.warmup_mode = "request"

# BCG captures every graph during a synthetic warmup forward at startup
# so serving never records a fresh graph.
if self.enable_breakable_cuda_graph and self.disagg_role == RoleType.MONOLITHIC:
# BCG must capture before serving. Preserve an explicit request warmup:
# it is the only way to capture request-only dimensions such as the
# Cosmos3 image/video frame count. Otherwise use the synthetic startup
# warmup as the safe default.
if (
self.enable_breakable_cuda_graph
and self.disagg_role == RoleType.MONOLITHIC
and self.warmup_mode != "request"
):
self.warmup_mode = "server"

# Disaggregated roles do not host the HTTP startup request. Preserve
Expand Down
Loading