From 803b4fb31c30229ebde1ea3b95aa087e10b0cfd0 Mon Sep 17 00:00:00 2001 From: Mick Date: Fri, 28 Aug 2026 19:08:30 +0800 Subject: [PATCH] [diffusion] refactor: scope model-specific API parameters (#35613) --- docs/cookbook/diffusion/Cosmos/Cosmos3.mdx | 9 +- docs/cookbook/diffusion/LTX/LTX2.5.mdx | 25 ++ .../diffusion/LongCat/LongCat-Image.mdx | 92 +++++++ docs/cookbook/diffusion/intro.mdx | 6 + docs/docs.json | 7 + docs/docs/sglang-diffusion/api/openai_api.mdx | 23 ++ python/sglang/multimodal_gen/README.md | 2 +- .../multimodal_gen/configs/sample/cosmos3.py | 186 ++++++++++++-- .../configs/sample/ernie_image.py | 4 + .../multimodal_gen/configs/sample/ideogram.py | 4 + .../configs/sample/longcat_image.py | 12 +- .../multimodal_gen/configs/sample/ltx_2.py | 1 + .../multimodal_gen/configs/sample/ltx_2_5.py | 20 +- .../configs/sample/sampling_params.py | 79 +++--- .../runtime/entrypoints/openai/image_api.py | 79 +++--- .../runtime/entrypoints/openai/protocol.py | 4 + .../runtime/entrypoints/openai/utils.py | 106 +++++++- .../runtime/entrypoints/openai/video_api.py | 237 +++--------------- .../multimodal_gen/test/unit/test_cosmos3.py | 29 ++- .../test/unit/test_openai_image_api.py | 89 +++++++ .../test/unit/test_video_api_profiling.py | 35 +++ 21 files changed, 732 insertions(+), 317 deletions(-) create mode 100644 docs/cookbook/diffusion/LongCat/LongCat-Image.mdx diff --git a/docs/cookbook/diffusion/Cosmos/Cosmos3.mdx b/docs/cookbook/diffusion/Cosmos/Cosmos3.mdx index 0dcb84c85..76ff862d9 100644 --- a/docs/cookbook/diffusion/Cosmos/Cosmos3.mdx +++ b/docs/cookbook/diffusion/Cosmos/Cosmos3.mdx @@ -123,7 +123,7 @@ curl -sS -X POST http://127.0.0.1:30010/v1/images/generations \ "guidance_scale": 6.0, "flow_shift": 3.0, "seed": 0, - "extra_args": { + "extra_body": { "use_resolution_template": false, "guardrails": true } @@ -141,7 +141,7 @@ curl -sS -X POST http://127.0.0.1:30010/v1/images/generations \ "n": 1, "guidance_scale": 1.0, "seed": 0, - "extra_args": { + "extra_body": { "use_resolution_template": false, "guardrails": true } @@ -390,7 +390,10 @@ Cosmos3 omnimodal fields are accepted as extra JSON fields or multipart form fie - `action_view_point`: viewpoint used in the structured action caption. - `action_normalization`: dataset normalization mode, such as `quantile`, `meanstd`, or `minmax`. -Put model-specific compatibility knobs in `extra_params` for video requests, or `extra_args` for image requests: +Pass model-specific controls through `extra_body` with the OpenAI Python SDK. +Raw JSON may keep them at the top level; multipart video requests should put +them in the `extra_params` JSON object. The legacy image `extra_args` container +remains accepted for compatibility, but new clients should use `extra_body`: - `use_duration_template`: whether to append SGLang's generated duration suffix to video prompts. - `use_resolution_template`: accepted for vLLM-Omni request compatibility. diff --git a/docs/cookbook/diffusion/LTX/LTX2.5.mdx b/docs/cookbook/diffusion/LTX/LTX2.5.mdx index 69ff7e031..94def6262 100644 --- a/docs/cookbook/diffusion/LTX/LTX2.5.mdx +++ b/docs/cookbook/diffusion/LTX/LTX2.5.mdx @@ -215,6 +215,23 @@ The prediction is clamped to `--auto-duration-min-seconds` / `--auto-duration-max-seconds` (default 1–20 s) and snapped to the VAE's temporal grid, so the result is always a valid frame count. It overrides `--num-frames`. +For an online server, pass the same LTX-2.5-only controls through `extra_body`: + +```python Python +from openai import OpenAI + +client = OpenAI(api_key="EMPTY", base_url="http://localhost:30010/v1") +video = client.videos.create( + model="Lightricks/LTX-2.5-Diffusers", + prompt="A red fox walking through a snowy forest at dawn.", + extra_body={ + "auto_duration": True, + "auto_duration_min_seconds": 2.0, + "auto_duration_max_seconds": 8.0, + }, +) +``` + ### 4.4 Two-stage (higher quality) Stage 1 runs at half the requested resolution, the latents are upsampled 2x, and @@ -290,6 +307,14 @@ sglang serve \ --load-diffusion-decoder ``` +```python Python +video = client.videos.create( + model="Lightricks/LTX-2.5-Diffusers", + prompt="A red fox walking through a snowy forest at dawn.", + extra_body={"use_diffusion_decoder": True}, +) +``` + This keeps the default server footprint unchanged while still allowing VAE and diffusion-decoder requests to share one server. When GPU memory is constrained, `--cpu-offload-components diffusion_decoder` keeps the optional decoder on CPU diff --git a/docs/cookbook/diffusion/LongCat/LongCat-Image.mdx b/docs/cookbook/diffusion/LongCat/LongCat-Image.mdx new file mode 100644 index 000000000..75192f720 --- /dev/null +++ b/docs/cookbook/diffusion/LongCat/LongCat-Image.mdx @@ -0,0 +1,92 @@ +--- +title: LongCat-Image +metatags: + description: "Deploy LongCat-Image with SGLang Diffusion and its native in-process Qwen2.5-VL prompt rewriter." +--- + +import { DiffusionModelTags } from '/src/snippets/diffusion/model-tags.jsx'; + + + +## 1. Model Introduction + +[LongCat-Image](https://huggingface.co/meituan-longcat/LongCat-Image) is a +text-to-image model from Meituan. SGLang runs its Qwen2.5-VL prompt rewriter +in process with the native SGLang runtime before text encoding and denoising. + +The native pipeline keeps prompt rewriting and diffusion behind one OpenAI-compatible +image endpoint. Rewriting is enabled by default for stronger prompt expansion, but +each request can disable it when lower latency matters more than the rewritten prompt. + +## 2. Installation + +Install SGLang with the diffusion dependencies: + +```bash Command +pip install -e "python[diffusion]" +``` + +For other installation options, see the +[SGLang Diffusion installation guide](/docs/sglang-diffusion/installation). + +## 3. Serve the model + +```bash Command +sglang serve \ + --model-path meituan-longcat/LongCat-Image \ + --performance-mode auto \ + --port 30010 +``` + +Prompt rewriting is enabled by default for LongCat-Image. It adds an +autoregressive Qwen2.5-VL pass before diffusion; set +`enable_prompt_rewrite=false` on a request when lower latency is more important +than rewritten prompt quality. + +## 4. Generate an image + +```python Python +import base64 +from openai import OpenAI + +client = OpenAI(api_key="EMPTY", base_url="http://127.0.0.1:30010/v1") + +response = client.images.generate( + model="meituan-longcat/LongCat-Image", + prompt="A quiet bookshop on a rainy evening, warm light in the windows", + n=1, + response_format="b64_json", +) + +image_bytes = base64.b64decode(response.data[0].b64_json) +with open("longcat_image.png", "wb") as f: + f.write(image_bytes) +``` + +To skip prompt rewriting with the OpenAI client, pass the model-specific request +field through `extra_body`: + +```python Python +response = client.images.generate( + model="meituan-longcat/LongCat-Image", + prompt="A quiet bookshop on a rainy evening", + extra_body={"enable_prompt_rewrite": False}, +) +``` + +## 5. Memory placement + +Use the unified component-residency selector when the complete pipeline does +not fit on the accelerator. For example, keep the repeatedly used DiT resident +while moving auxiliary components to CPU between stages: + +```bash Command +sglang serve \ + --model-path meituan-longcat/LongCat-Image \ + --component-residency dit=resident text_encoder=component-offload vae=component-offload \ + --pin-cpu-memory \ + --port 30010 +``` + +See [Component Residency](/docs/sglang-diffusion/api/cli#component-residency) +for mode semantics and compatibility with the existing CPU-offload flags. diff --git a/docs/cookbook/diffusion/intro.mdx b/docs/cookbook/diffusion/intro.mdx index 5ecc49011..37772941b 100644 --- a/docs/cookbook/diffusion/intro.mdx +++ b/docs/cookbook/diffusion/intro.mdx @@ -31,6 +31,12 @@ Image models generate one image request as a bounded denoising job, usually with href="/cookbook/diffusion/Qwen-Image/Qwen-Image" img="/cards/logos/qwen.png" /> + Any: + if not isinstance(value, str): + return value + try: + return json.loads(value) + except (json.JSONDecodeError, TypeError, ValueError): + return value + + +def _optional_int_list(value: Any) -> list[int] | None: + value = _parse_request_value(value) + if value is None or (isinstance(value, str) and not value.strip()): + return None + if isinstance(value, (list, tuple)): + return [int(item) for item in value] + return [int(value)] + + @dataclass class Cosmos3SamplingParams(SamplingParams): """Cosmos3 sampling parameters (T2V defaults; also used for I2V / V2V / T2I). @@ -63,6 +82,12 @@ class Cosmos3SamplingParams(SamplingParams): negative_prompt: str = "" + use_duration_template: bool | None = None + use_resolution_template: bool | None = None + use_system_prompt: bool | None = None + use_guardrails: bool | None = None + sound_duration: float = 0.0 + # Optional CFG window — T2I requests typically pass e.g. ``(400, 1000)`` to # skip guidance at low noise levels. T2V / I2V / V2V leave it unset. guidance_interval: tuple[float, float] | None = None @@ -159,33 +184,54 @@ class Cosmos3SamplingParams(SamplingParams): action_normalization: str = "quantile" @classmethod - def video_request_extra_fields(cls) -> frozenset[str]: + def image_request_extra_fields(cls) -> frozenset[str]: return frozenset( { - "generate_sound", - "sound_duration", + "guidance_interval", + "use_duration_template", + "use_guardrails", + "use_resolution_template", + "use_system_prompt", + } + ) + + @classmethod + def default_image_output_format(cls) -> str: + return "png" + + @classmethod + def default_image_response_format(cls) -> str: + return "b64_json" + + @classmethod + def video_request_extra_fields(cls) -> frozenset[str]: + return cls.image_request_extra_fields() | frozenset( + { + "action", + "action_fps", + "action_mode", + "action_normalization", + "action_view_point", "condition_frame_indexes", "condition_frame_indexes_vision", "condition_video_keep", - "control_path", - "control_hint", "control_guidance", "control_guidance_interval", - "num_video_frames_per_chunk", - "num_conditional_frames", - "num_first_chunk_conditional_frames", - "max_frames", - "show_control_condition", - "show_input", - "share_vision_temporal_positions", - "action_mode", + "control_hint", + "control_path", "domain_id", "domain_name", + "generate_sound", + "guardrails", + "max_frames", + "num_conditional_frames", + "num_first_chunk_conditional_frames", + "num_video_frames_per_chunk", "raw_action_dim", - "action_fps", - "action", - "action_view_point", - "action_normalization", + "share_vision_temporal_positions", + "show_control_condition", + "show_input", + "sound_duration", } ) @@ -247,16 +293,114 @@ class Cosmos3SamplingParams(SamplingParams): def lower_video_request_kwargs( cls, request: Any, kwargs: dict[str, Any] ) -> dict[str, Any]: - """Apply defaults that the generic video endpoint pre-resolves.""" + kwargs = super().lower_video_request_kwargs(request, dict(kwargs)) + extras = getattr(request, "model_extra", None) or {} + + if "use_guardrails" not in kwargs and extras.get("guardrails") is not None: + kwargs["use_guardrails"] = _parse_request_value(extras["guardrails"]) + + condition_indexes = kwargs.get("condition_frame_indexes") + if condition_indexes is None: + condition_indexes = extras.get("condition_frame_indexes_vision") + condition_indexes = _optional_int_list(condition_indexes) + if condition_indexes is not None: + kwargs["condition_frame_indexes"] = condition_indexes + + if "sound_duration" in kwargs: + kwargs["sound_duration"] = float( + _parse_request_value(kwargs["sound_duration"]) + ) + generate_sound = _parse_request_value(extras.get("generate_sound")) + if generate_sound is False: + kwargs["sound_duration"] = 0.0 + elif generate_sound is True and "sound_duration" not in kwargs: + kwargs["sound_duration"] = float(kwargs["num_frames"]) / float( + kwargs["fps"] + ) + + for name in ("control_path", "control_hint"): + value = _parse_request_value(kwargs.get(name)) + if isinstance(value, (list, tuple)): + value = [str(item) for item in value if str(item).strip()] + elif value is not None and not isinstance(value, str): + value = str(value) + if isinstance(value, str): + value = value if value.strip() else None + if value: + kwargs[name] = value + else: + kwargs.pop(name, None) + + if "control_guidance" in kwargs: + kwargs["control_guidance"] = float( + _parse_request_value(kwargs["control_guidance"]) + ) + if "control_guidance_interval" in kwargs: + interval = _parse_request_value(kwargs["control_guidance_interval"]) + if interval is None or (isinstance(interval, str) and not interval.strip()): + kwargs.pop("control_guidance_interval") + else: + if not isinstance(interval, (list, tuple)): + interval = [interval] + kwargs["control_guidance_interval"] = tuple( + float(item) for item in interval + ) + + for name in ( + "num_video_frames_per_chunk", + "num_conditional_frames", + "num_first_chunk_conditional_frames", + "max_frames", + ): + value = _parse_request_value(kwargs.get(name)) + if value is not None and value != "": + kwargs[name] = int(value) + + for name in ( + "show_control_condition", + "show_input", + "share_vision_temporal_positions", + ): + value = _parse_request_value(kwargs.get(name)) + if value is None or (isinstance(value, str) and not value.strip()): + kwargs.pop(name, None) + elif isinstance(value, str): + normalized = value.strip().lower() + if normalized in {"1", "true", "yes", "on"}: + kwargs[name] = True + elif normalized in {"0", "false", "no", "off"}: + kwargs[name] = False + else: + raise ValueError(f"Invalid boolean value: {value!r}") + else: + kwargs[name] = bool(value) + + for name in ( + "condition_video_keep", + "action_mode", + "domain_id", + "domain_name", + "raw_action_dim", + "action_fps", + "action", + "action_view_point", + "action_normalization", + ): + value = _parse_request_value(kwargs.get(name)) + if isinstance(value, str) and not value.strip(): + kwargs.pop(name, None) + elif value is not None: + kwargs[name] = value + hint = kwargs.get("control_hint") paths = kwargs.get("control_path") hints = [hint] if isinstance(hint, str) else list(hint or []) control_paths = [paths] if isinstance(paths, str) else list(paths or []) if len(control_paths) == 1 and hints == ["wsm"]: defaults = cls._TRANSFER_DEFAULTS["wsm"] - if getattr(request, "num_frames", None) is None: + if request.num_frames is None: kwargs["num_frames"] = defaults["num_frames"] - if getattr(request, "fps", None) is None: + if request.fps is None: kwargs["fps"] = defaults["fps"] return kwargs @@ -264,6 +408,8 @@ class Cosmos3SamplingParams(SamplingParams): # adjust distil and edge args — read from the pre-computed config fields # so no checkpoint download happens at request time. pipeline_config = server_args.pipeline_config + if self.action_stats_path is None: + self.action_stats_path = getattr(pipeline_config, "action_stats_path", None) distilled_sigmas = pipeline_config.distilled_sigmas if distilled_sigmas is not None: self.num_inference_steps = len(distilled_sigmas) diff --git a/python/sglang/multimodal_gen/configs/sample/ernie_image.py b/python/sglang/multimodal_gen/configs/sample/ernie_image.py index d985180d4..0ef35c09e 100644 --- a/python/sglang/multimodal_gen/configs/sample/ernie_image.py +++ b/python/sglang/multimodal_gen/configs/sample/ernie_image.py @@ -13,3 +13,7 @@ class ErnieImageSamplingParams(SamplingParams): guidance_scale: float = 5.0 num_inference_steps: int = 50 use_pe: bool = True + + @classmethod + def image_request_extra_fields(cls) -> frozenset[str]: + return frozenset({"use_pe"}) diff --git a/python/sglang/multimodal_gen/configs/sample/ideogram.py b/python/sglang/multimodal_gen/configs/sample/ideogram.py index 62afbc5f1..561a3b845 100644 --- a/python/sglang/multimodal_gen/configs/sample/ideogram.py +++ b/python/sglang/multimodal_gen/configs/sample/ideogram.py @@ -53,6 +53,10 @@ class Ideogram4SamplingParams(SamplingParams): guidance_scale: float | None = None preset: str = "V4_DEFAULT_20" + @classmethod + def image_request_extra_fields(cls) -> frozenset[str]: + return frozenset({"preset"}) + def __post_init__(self) -> None: if self.preset not in IDEOGRAM4_PRESETS: raise ValueError( diff --git a/python/sglang/multimodal_gen/configs/sample/longcat_image.py b/python/sglang/multimodal_gen/configs/sample/longcat_image.py index 91d35b5c2..49890eaf5 100644 --- a/python/sglang/multimodal_gen/configs/sample/longcat_image.py +++ b/python/sglang/multimodal_gen/configs/sample/longcat_image.py @@ -9,10 +9,20 @@ class LongCatImageSamplingParams(SamplingParams): guidance_scale: float = 4.5 height: int = 1024 width: int = 1024 - # Override base class defaults to enable LongCat-specific features by default enable_cfg_renorm: bool = True + cfg_renorm_min: float = 0.0 enable_prompt_rewrite: bool = True + @classmethod + def image_request_extra_fields(cls) -> frozenset[str]: + return frozenset( + { + "cfg_renorm_min", + "enable_cfg_renorm", + "enable_prompt_rewrite", + } + ) + @dataclass class LongCatImageEditSamplingParams(SamplingParams): diff --git a/python/sglang/multimodal_gen/configs/sample/ltx_2.py b/python/sglang/multimodal_gen/configs/sample/ltx_2.py index 23989b820..2962c7545 100644 --- a/python/sglang/multimodal_gen/configs/sample/ltx_2.py +++ b/python/sglang/multimodal_gen/configs/sample/ltx_2.py @@ -22,6 +22,7 @@ class LTX2SamplingParams(SamplingParams): # Audio specific generate_audio: bool = True + use_diffusion_decoder: bool = False # Denoising parameters guidance_scale: float = 4.0 diff --git a/python/sglang/multimodal_gen/configs/sample/ltx_2_5.py b/python/sglang/multimodal_gen/configs/sample/ltx_2_5.py index 8d7bad1b5..b525b02c6 100644 --- a/python/sglang/multimodal_gen/configs/sample/ltx_2_5.py +++ b/python/sglang/multimodal_gen/configs/sample/ltx_2_5.py @@ -26,8 +26,24 @@ class LTX25SamplingParams(LTX2SamplingParams): guidance_scale: float = 1.0 - # `auto_duration` on the base class has the duration head predict this - # instead, overriding `num_frames`. + use_diffusion_decoder: bool = False + auto_duration: bool = False + auto_duration_min_seconds: float = 1.0 + auto_duration_max_seconds: float = 20.0 + + # `auto_duration` has the duration head predict this instead, overriding + # `num_frames`. # The schedule is pinned by the pipeline config; this only keeps the # reported step count honest. num_inference_steps: int = 8 + + @classmethod + def video_request_extra_fields(cls) -> frozenset[str]: + return frozenset( + { + "auto_duration", + "auto_duration_max_seconds", + "auto_duration_min_seconds", + "use_diffusion_decoder", + } + ) diff --git a/python/sglang/multimodal_gen/configs/sample/sampling_params.py b/python/sglang/multimodal_gen/configs/sample/sampling_params.py index f28227fec..68bca19ca 100644 --- a/python/sglang/multimodal_gen/configs/sample/sampling_params.py +++ b/python/sglang/multimodal_gen/configs/sample/sampling_params.py @@ -97,10 +97,17 @@ class DataType(Enum): @dataclass class SamplingParams: """ - Sampling parameters for generation. + Model-agnostic sampling parameters for generation. Dynamic batching compares these fields for compatibility, except fields marked with `batch_sig_exclude`. + + New fields in this base class must be shared across model families; legacy + compatibility fields are not precedent. A model-specific field belongs on + that model's SamplingParams subclass and, when accepted by an online + endpoint, must also be declared via ``image_request_extra_fields`` or + ``video_request_extra_fields``. Do not add model fields here merely to make + the common API transport accept them. """ data_type: DataType = DataType.VIDEO @@ -180,24 +187,11 @@ class SamplingParams: width: int | None = None fps: int = 24 - # LTX-2.5 duration head. Ignored by other models, so the flags stay - # universally accepted. - # Decode with the diffusion decoder instead of the VAE one. Ignored by - # models that ship no such decoder. - use_diffusion_decoder: bool = False - - auto_duration: bool = False - auto_duration_min_seconds: float = 1.0 - auto_duration_max_seconds: float = 20.0 - # Resolution validation supported_resolutions: list[tuple[int, int]] | None = field( default=None, metadata={"batch_sig_exclude": True} ) # None means all resolutions allowed - # Output audio duration in seconds (models without an audio modality ignore this). - sound_duration: float = 0.0 - # Denoising parameters num_inference_steps: int = None guidance_scale: float = 1.0 @@ -215,11 +209,6 @@ class SamplingParams: progressive_levels: int = 1 progressive_delta: float = 0.01 - # LongCat-Image parameters - enable_cfg_renorm: bool = False - cfg_renorm_min: float = 0.0 - enable_prompt_rewrite: bool = False - # TeaCache parameters enable_teacache: bool = False teacache_params: Any = ( @@ -292,17 +281,9 @@ class SamplingParams: max_sequence_length: int | None = None flow_shift: float | None = None - # cosmos-related - use_duration_template: bool | None = None - use_resolution_template: bool | None = None - use_system_prompt: bool | None = None - use_guardrails: bool | None = None condition_inputs: dict[str, Any] = field(default_factory=dict) realtime_chunk_size: int | None = None - # Prompt enhancement (ErnieImage) - use_pe: bool | None = None - def _set_output_file_ext(self): # add extension if needed output_extensions = (".mp4", ".jpg", ".png", ".webp", ".obj", ".glb", ".json") @@ -395,11 +376,39 @@ class SamplingParams: req.realtime_chunk_size = self.realtime_chunk_size @classmethod - def video_request_extra_fields(cls) -> frozenset[str]: - """Declare model-specific multipart video fields accepted by this type.""" + def image_request_extra_fields(cls) -> frozenset[str]: + """Declare model-owned JSON fields accepted by the image API. + + Every returned name must be an init field on ``cls``. The common + endpoint resolves the active subclass before reading these fields, so + model-specific extraction and defaults stay out of the API layer. + """ return frozenset() + @classmethod + def video_request_extra_fields(cls) -> frozenset[str]: + """Declare model-owned JSON or multipart fields accepted by the video API. + + Dataclass-backed names are forwarded to ``cls``. Transport-only aliases + may also be declared so multipart parsing preserves them, but the + subclass must consume those aliases in ``lower_video_request_kwargs``. + """ + + return frozenset() + + @classmethod + def default_image_output_format(cls) -> str | None: + """Return a model-owned default format for the image API, if any.""" + + return None + + @classmethod + def default_image_response_format(cls) -> str | None: + """Return a model-owned default response format for the image API, if any.""" + + return None + @classmethod def lower_video_request_kwargs( cls, @@ -904,7 +913,13 @@ class SamplingParams: @staticmethod def add_cli_args(parser: Any) -> Any: - """Add CLI arguments for SamplingParam fields""" + """Add CLI arguments for SamplingParam fields. + + This shared parser still contains legacy model-specific flags because + argparse is constructed before the active model is resolved. Do not add + new model-specific dataclass fields to ``SamplingParams`` or new API + special cases here; model request ownership remains on subclasses. + """ def add_argument(*name_or_flags, **kwargs): kwargs.setdefault("default", argparse.SUPPRESS) @@ -1008,7 +1023,7 @@ class SamplingParams: add_argument( "--enable-cfg-renorm", action=StoreBoolean, - help="Enable CFG renormalization for LongCat-Image (default: false).", + help="Enable CFG renormalization for LongCat-Image (enabled by default).", ) add_argument( "--cfg-renorm-min", @@ -1018,7 +1033,7 @@ class SamplingParams: add_argument( "--enable-prompt-rewrite", action=StoreBoolean, - help="Enable prompt rewriting via Qwen2.5-VL before encoding for LongCat-Image (default: false).", + help="Enable prompt rewriting via Qwen2.5-VL before encoding for LongCat-Image (enabled by default).", ) # profiling diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py index f4db1d78c..88b40d0e7 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py @@ -3,7 +3,6 @@ import asyncio import base64 import contextlib -import json import os import time from typing import Any, List, Optional @@ -37,9 +36,11 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( add_common_data_to_response, build_sampling_params, choose_output_image_ext, - flatten_extra_params, + get_sampling_request_extra_fields, merge_image_input_list, process_generation_batch, + request_extra_value, + resolve_sampling_params_cls, save_image_to_path, temp_dir_if_disabled, ) @@ -53,20 +54,7 @@ router = APIRouter(prefix="/v1/images", tags=["images"]) def _get_extra_field(request, field_name): - """Get a field from model_extra, with fallback to nested extra_body dict.""" - extra = request.model_extra or {} - value = extra.get(field_name) - if value is not None: - return value - if field_name == "use_guardrails" and extra.get("guardrails") is not None: - return extra["guardrails"] - - for container_name in ("extra_body", "extra_json", "extra_args", "extra_params"): - value = _parse_extra_container(extra.get(container_name)).get(field_name) - if value is not None: - return value - - return value + return request_extra_value(request, field_name) def _get_request_field_or_extra(request, field_name): @@ -76,22 +64,25 @@ def _get_request_field_or_extra(request, field_name): return _get_extra_field(request, field_name) +def _image_request_model_kwargs( + request: ImageGenerationsRequest, + sampling_params_cls: type[SamplingParams], +) -> dict[str, Any]: + """Extract fields owned and declared by the active model contract.""" + + kwargs = {} + for field_name in get_sampling_request_extra_fields(sampling_params_cls, "image"): + value = _get_extra_field(request, field_name) + if value is not None: + kwargs[field_name] = value + return kwargs + + def _runtime_sampling_quality(quality: str | None) -> str | None: """Keep OpenAI's automatic default out of SGLang's sampling contract.""" return None if quality in (None, "auto") else quality -def _parse_extra_container(value: Any) -> dict[str, Any]: - if isinstance(value, str): - try: - value = json.loads(value) - except Exception: - return {} - if isinstance(value, dict): - return flatten_extra_params(dict(value)) - return {} - - def _read_b64_for_paths(paths: list[str]) -> list[str]: """Read and base64-encode each file. Must be called before cloud upload deletes them.""" result = [] @@ -275,12 +266,14 @@ async def generations( ): request_id = generate_request_id() server_args = get_global_server_args() - is_cosmos3 = "cosmos3" in (server_args.model_path or "").lower() - ext = ( - "png" - if is_cosmos3 and request.output_format is None - else choose_output_image_ext(request.output_format, request.background) + sampling_params_cls = resolve_sampling_params_cls(server_args) + model_kwargs = _image_request_model_kwargs(request, sampling_params_cls) + output_format = ( + request.output_format + if request.output_format is not None + else sampling_params_cls.default_image_output_format() ) + ext = choose_output_image_ext(output_format, request.background) with temp_dir_if_disabled(server_args.output_path) as output_dir: sampling = build_sampling_params( @@ -309,12 +302,6 @@ async def generations( if request.flow_shift is not None else _get_extra_field(request, "flow_shift") ), - use_duration_template=_get_extra_field(request, "use_duration_template"), - use_resolution_template=_get_extra_field( - request, "use_resolution_template" - ), - use_system_prompt=_get_extra_field(request, "use_system_prompt"), - use_guardrails=_get_extra_field(request, "use_guardrails"), enable_teacache=request.enable_teacache, enable_cache_dit=_get_extra_field(request, "enable_cache_dit"), cache_dit_params=_get_extra_field(request, "cache_dit_params"), @@ -330,13 +317,12 @@ async def generations( upscaling_model_path=request.upscaling_model_path, upscaling_scale=request.upscaling_scale, perf_dump_path=request.perf_dump_path, - use_pe=_get_extra_field(request, "use_pe"), - preset=_get_extra_field(request, "preset"), progressive_mode=_get_request_field_or_extra(request, "progressive_mode"), progressive_levels=_get_request_field_or_extra( request, "progressive_levels" ), progressive_delta=_get_request_field_or_extra(request, "progressive_delta"), + **model_kwargs, ) trace_headers = extract_trace_headers(raw_request.headers) batch = prepare_request( @@ -353,13 +339,12 @@ async def generations( ) save_file_path = save_file_path_list[0] response_resize = _get_response_resize(sampling, save_file_path) - resp_format = (request.response_format or "b64_json").lower() - if ( - is_cosmos3 - and "response_format" not in request.model_fields_set - and request.response_format == "url" - ): - resp_format = "b64_json" + response_format = request.response_format + if "response_format" not in request.model_fields_set: + response_format = ( + sampling_params_cls.default_image_response_format() or response_format + ) + resp_format = (response_format or "b64_json").lower() # read b64 before cloud upload may delete the local file b64_list = ( diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py index ac6895035..6db73b316 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py @@ -38,6 +38,10 @@ class ImageResponse(BaseModel): usage: Optional[ImageUsage] = None +# Keep request schemas limited to OpenAI fields and stable cross-model SGLang +# extensions. Model-owned controls travel as allowed extras and are interpreted +# only after the active SamplingParams subclass is resolved; do not add them to +# these shared protocol models. class ImageGenerationsRequest(BaseModel): model_config = ConfigDict(extra="allow") diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py index 6b5ef964d..bb2c3abe7 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py @@ -1,5 +1,6 @@ # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo import asyncio +import dataclasses import inspect import json import os @@ -7,7 +8,8 @@ import shutil import tempfile import time from contextlib import contextmanager -from typing import Any, Generator, List, Optional, Union +from functools import cache +from typing import Any, Generator, List, Literal, Optional, Union import httpx from fastapi import HTTPException, UploadFile @@ -95,6 +97,81 @@ def flatten_extra_params(payload: Any) -> dict[str, Any]: return payload +_REQUEST_EXTRA_CONTAINERS = ( + "extra_body", + "extra_json", + "extra_args", + "extra_params", +) + + +def _parse_request_extra_container(value: Any) -> dict[str, Any]: + if isinstance(value, str): + try: + value = json.loads(value) + except (json.JSONDecodeError, TypeError, ValueError): + return {} + if not isinstance(value, dict): + return {} + return flatten_extra_params(dict(value)) + + +def request_extra_value(request: Any, field_name: str) -> Any: + """Read an extension field while preserving top-level precedence. + + This function only handles transport compatibility. Callers must first use + the active SamplingParams subclass to decide which model-owned fields are + valid; transport helpers must not introduce per-model allowlists. + """ + + extra = dict(getattr(request, "model_extra", None) or {}) + direct = { + key: value + for key, value in extra.items() + if key not in _REQUEST_EXTRA_CONTAINERS + } + direct = flatten_extra_params(direct) + if field_name in direct and direct[field_name] is not None: + return direct[field_name] + + for container_name in _REQUEST_EXTRA_CONTAINERS: + nested = _parse_request_extra_container(extra.get(container_name)) + if field_name in nested and nested[field_name] is not None: + return nested[field_name] + return None + + +@cache +def get_declared_request_extra_fields( + sampling_params_cls: type[SamplingParams], + api: Literal["image", "video"], +) -> frozenset[str]: + """Return the active model's accepted fields, including transport aliases.""" + + if api == "image": + return sampling_params_cls.image_request_extra_fields() + return sampling_params_cls.video_request_extra_fields() + + +@cache +def get_sampling_request_extra_fields( + sampling_params_cls: type[SamplingParams], + api: Literal["image", "video"], +) -> frozenset[str]: + """Return declared extension fields that can initialize SamplingParams. + + A video declaration may also contain transport-only aliases. Those remain + on the request for the model's lowering hook instead of being passed to the + dataclass constructor. + """ + + declared = get_declared_request_extra_fields(sampling_params_cls, api) + init_fields = { + field.name for field in dataclasses.fields(sampling_params_cls) if field.init + } + return declared & init_fields + + @contextmanager def temp_dir_if_disabled( configured_path: str | None, @@ -179,6 +256,33 @@ def build_sampling_params(request_id: str, **kwargs) -> SamplingParams: return sampling_params +def resolve_sampling_params_cls(server_args: Any) -> type[SamplingParams]: + """Resolve the model-owned sampling contract selected for this server. + + Shared API code must dispatch through this type instead of branching on a + model ID or importing individual model configurations. + """ + + sampling_params_cls = SamplingParams + if server_args.pipeline_class_name: + from sglang.multimodal_gen.registry import get_pipeline_config_classes + + config_classes = get_pipeline_config_classes(server_args.pipeline_class_name) + if config_classes is not None: + _, sampling_params_cls = config_classes + if sampling_params_cls is SamplingParams: + from sglang.multimodal_gen.registry import get_model_info + + model_info = get_model_info( + server_args.model_path, + backend=server_args.backend, + model_id=server_args.model_id, + ) + if model_info is not None: + sampling_params_cls = model_info.sampling_param_cls + return sampling_params_cls + + async def save_image_to_path( image: Union[UploadFile, bytes, str], target_path: str, diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py index b0f2b6718..47d884546 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py @@ -40,8 +40,12 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( add_common_data_to_response, build_sampling_params, flatten_extra_params, + get_declared_request_extra_fields, + get_sampling_request_extra_fields, merge_image_input_list, process_generation_batch, + request_extra_value, + resolve_sampling_params_cls, save_image_to_path, ) from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request @@ -90,7 +94,7 @@ async def shutdown_video_jobs() -> None: def _extra_value(request: VideoGenerationsRequest, name: str) -> Any: - return (request.model_extra or {}).get(name) + return request_extra_value(request, name) def _request_value(request: VideoGenerationsRequest, name: str) -> Any: @@ -100,6 +104,20 @@ def _request_value(request: VideoGenerationsRequest, name: str) -> Any: return _extra_value(request, name) +def _video_request_model_kwargs( + request: VideoGenerationsRequest, + sampling_params_cls: type[SamplingParams], +) -> dict[str, Any]: + """Extract fields owned and declared by the active model contract.""" + + kwargs = {} + for field_name in get_sampling_request_extra_fields(sampling_params_cls, "video"): + value = _extra_value(request, field_name) + if value is not None: + kwargs[field_name] = value + return kwargs + + def _parse_form_extra_value(value: Any) -> Any: if not isinstance(value, str): return value @@ -110,40 +128,14 @@ def _parse_form_extra_value(value: Any) -> Any: _MULTIPART_EXTRA_FORM_FIELDS = ( - "use_duration_template", - "use_resolution_template", - "use_system_prompt", - "use_guardrails", - "guardrails", - "video_path", - "video_url", + "attention_backend_override", + "cache_dit_params", + "cfg_gate_step", + "enable_cache_dit", "quality", ) -def _video_sampling_params_cls(server_args) -> type[SamplingParams]: - """Resolve the params type selected for the current server.""" - - sampling_params_cls = SamplingParams - if server_args.pipeline_class_name: - from sglang.multimodal_gen.registry import get_pipeline_config_classes - - config_classes = get_pipeline_config_classes(server_args.pipeline_class_name) - if config_classes is not None: - _, sampling_params_cls = config_classes - if sampling_params_cls is SamplingParams: - from sglang.multimodal_gen.registry import get_model_info - - model_info = get_model_info( - server_args.model_path, - backend=server_args.backend, - model_id=server_args.model_id, - ) - if model_info is not None: - sampling_params_cls = model_info.sampling_param_cls - return sampling_params_cls - - def _multipart_extra_form_keys( sampling_params_cls: type[SamplingParams], ) -> tuple[str, ...]: @@ -152,7 +144,9 @@ def _multipart_extra_form_keys( ( *VideoGenerationsRequest.model_fields, *_MULTIPART_EXTRA_FORM_FIELDS, - *sorted(sampling_params_cls.video_request_extra_fields()), + *sorted( + get_declared_request_extra_fields(sampling_params_cls, "video") + ), ) ) ) @@ -172,7 +166,7 @@ def _merge_multipart_extra_form_fields( sampling_params_cls: type[SamplingParams], ) -> None: for key in _multipart_extra_form_keys(sampling_params_cls): - if key in raw_form and key not in extra_from_form: + if key in raw_form: extra_from_form[key] = _parse_form_extra_value(raw_form[key]) @@ -229,72 +223,6 @@ def _is_probably_video_source(source: Any) -> bool: return os.path.splitext(source_name)[1].lower() in _VIDEO_EXTENSIONS -def _is_cosmos3_server(server_args) -> bool: - from sglang.multimodal_gen.configs.pipeline_configs.cosmos3 import Cosmos3Config - - return isinstance(server_args.pipeline_config, Cosmos3Config) - - -def _normalize_optional_string(value: Any) -> Any: - if isinstance(value, str) and not value.strip(): - return None - return value - - -def _coerce_optional_int_list(value: Any) -> list[int] | None: - value = _parse_form_extra_value(value) - if value is None: - return None - if isinstance(value, str) and not value.strip(): - return None - if isinstance(value, (list, tuple)): - return [int(item) for item in value] - return [int(value)] - - -def _coerce_optional_float_list(value: Any) -> list[float] | None: - value = _parse_form_extra_value(value) - if value is None: - return None - if isinstance(value, str) and not value.strip(): - return None - if isinstance(value, (list, tuple)): - return [float(item) for item in value] - return [float(value)] - - -def _coerce_optional_bool(value: Any) -> bool | None: - value = _parse_form_extra_value(value) - if value is None or (isinstance(value, str) and not value.strip()): - return None - if isinstance(value, str): - normalized = value.strip().lower() - if normalized in {"1", "true", "yes", "on"}: - return True - if normalized in {"0", "false", "no", "off"}: - return False - raise ValueError(f"Invalid boolean value: {value!r}") - return bool(value) - - -def _coerce_optional_str_list(value: Any) -> str | list[str] | None: - """Coerce a control_path/control_hint value to str or list[str]. - - Accepts a JSON list (``["edge.mp4", "depth.mp4"]``), a plain string, or a - native list. Empty values resolve to ``None`` so unset fields don't override - sampling-param defaults. - """ - value = _parse_form_extra_value(value) - if value is None: - return None - if isinstance(value, (list, tuple)): - items = [str(item) for item in value if str(item).strip()] - return items or None - if isinstance(value, str): - return value if value.strip() else None - return str(value) - - def _resolve_video_path(req: VideoGenerationsRequest) -> str | None: video_path = _request_value(req, "video_path") or _request_value(req, "video_url") if video_path: @@ -322,96 +250,11 @@ def _resolve_image_path( return image_path -def _resolve_sound_duration( - req: VideoGenerationsRequest, *, num_frames: int, fps: int -) -> float | None: - generate_sound = _request_value(req, "generate_sound") - sound_duration = _request_value(req, "sound_duration") - - if generate_sound is False: - return 0.0 - if sound_duration is not None: - return float(sound_duration) - if generate_sound is True: - return float(num_frames) / float(fps) - return None - - -def _cosmos3_sampling_param_kwargs( - req: VideoGenerationsRequest, *, num_frames: int, fps: int -) -> Dict[str, Any]: - """Map HTTP/API aliases to Cosmos3SamplingParams field names.""" - kwargs: Dict[str, Any] = {} - - sound_duration = _resolve_sound_duration(req, num_frames=num_frames, fps=fps) - if sound_duration is not None: - kwargs["sound_duration"] = sound_duration - - condition_frame_indexes = _request_value(req, "condition_frame_indexes") - if condition_frame_indexes is None: - condition_frame_indexes = _request_value(req, "condition_frame_indexes_vision") - condition_frame_indexes = _coerce_optional_int_list(condition_frame_indexes) - if condition_frame_indexes is not None: - kwargs["condition_frame_indexes"] = condition_frame_indexes - - # Transfer (control-video) conditioning. - control_path = _coerce_optional_str_list(_request_value(req, "control_path")) - if control_path is not None: - kwargs["control_path"] = control_path - control_hint = _coerce_optional_str_list(_request_value(req, "control_hint")) - if control_hint is not None: - kwargs["control_hint"] = control_hint - control_guidance = _request_value(req, "control_guidance") - if control_guidance is not None: - kwargs["control_guidance"] = float(control_guidance) - control_guidance_interval = _coerce_optional_float_list( - _request_value(req, "control_guidance_interval") - ) - if control_guidance_interval is not None: - kwargs["control_guidance_interval"] = tuple(control_guidance_interval) - - for name in ( - "num_video_frames_per_chunk", - "num_conditional_frames", - "num_first_chunk_conditional_frames", - "max_frames", - ): - value = _parse_form_extra_value(_request_value(req, name)) - if value is not None and value != "": - kwargs[name] = int(value) - - for name in ( - "show_control_condition", - "show_input", - "share_vision_temporal_positions", - ): - value = _coerce_optional_bool(_request_value(req, name)) - if value is not None: - kwargs[name] = value - - for name in ( - "condition_video_keep", - "action_mode", - "domain_id", - "domain_name", - "raw_action_dim", - "action_fps", - "action", - "action_view_point", - "action_normalization", - ): - value = _parse_form_extra_value(_request_value(req, name)) - value = _normalize_optional_string(value) - if value is not None: - kwargs[name] = value - - return kwargs - - def _build_video_sampling_params(request_id: str, request: VideoGenerationsRequest): """Resolve video-specific defaults (fps, seconds → num_frames) then delegate to the shared build_sampling_params.""" server_args = get_global_server_args() + sampling_params_cls = resolve_sampling_params_cls(server_args) seconds = request.seconds if request.seconds is not None else DEFAULT_VIDEO_SECONDS fps = request.fps if request.fps is not None else DEFAULT_FPS num_frames = request.num_frames if request.num_frames is not None else fps * seconds @@ -420,15 +263,6 @@ def _build_video_sampling_params(request_id: str, request: VideoGenerationsReque num_outputs = request.n or 1 video_path = _resolve_video_path(request) image_path = _resolve_image_path(request, video_path) - cosmos3_kwargs = {} - if _is_cosmos3_server(server_args): - cosmos3_kwargs = _cosmos3_sampling_param_kwargs( - request, num_frames=num_frames, fps=fps - ) - if server_args.pipeline_config.action_stats_path is not None: - cosmos3_kwargs["action_stats_path"] = ( - server_args.pipeline_config.action_stats_path - ) kwargs = { "prompt": request.prompt, @@ -450,10 +284,6 @@ def _build_video_sampling_params(request_id: str, request: VideoGenerationsReque "negative_prompt": request.negative_prompt, "max_sequence_length": request.max_sequence_length, "flow_shift": request.flow_shift, - "use_duration_template": _extra_value(request, "use_duration_template"), - "use_resolution_template": _extra_value(request, "use_resolution_template"), - "use_system_prompt": _extra_value(request, "use_system_prompt"), - "use_guardrails": _extra_value(request, "use_guardrails"), "enable_teacache": request.enable_teacache, "enable_cache_dit": _extra_value(request, "enable_cache_dit"), "cache_dit_params": _extra_value(request, "cache_dit_params"), @@ -477,10 +307,9 @@ def _build_video_sampling_params(request_id: str, request: VideoGenerationsReque "num_profiled_timesteps": request.num_profiled_timesteps, "profile_all_stages": request.profile_all_stages, "diffusers_kwargs": request.diffusers_kwargs, - **cosmos3_kwargs, + **_video_request_model_kwargs(request, sampling_params_cls), } - sampling_params_cls = _video_sampling_params_cls(server_args) kwargs = sampling_params_cls.lower_video_request_kwargs(request, kwargs) sampling_params = build_sampling_params(request_id, **kwargs) if ( @@ -687,7 +516,7 @@ async def create_video( raw_form, extra_body=extra_body, extra_params=extra_params, - sampling_params_cls=_video_sampling_params_cls(server_args), + sampling_params_cls=resolve_sampling_params_cls(server_args), ) # Resolve input upload directory (may be a temp dir when saving is disabled) @@ -831,13 +660,15 @@ async def create_video( if isinstance(extra, str): extra = json.loads(extra) if isinstance(extra, dict): - payload.update(flatten_extra_params(extra)) + for key, value in flatten_extra_params(extra).items(): + payload.setdefault(key, value) # openai may turn extra_body to extra_json extra_json = payload.pop("extra_json", None) if isinstance(extra_json, str): extra_json = json.loads(extra_json) if isinstance(extra_json, dict): - payload.update(flatten_extra_params(extra_json)) + for key, value in flatten_extra_params(extra_json).items(): + payload.setdefault(key, value) flatten_extra_params(payload) # Validate image input based on model task type if payload.get("video_url") and not payload.get("video_path"): diff --git a/python/sglang/multimodal_gen/test/unit/test_cosmos3.py b/python/sglang/multimodal_gen/test/unit/test_cosmos3.py index df2d9090a..097923c7e 100644 --- a/python/sglang/multimodal_gen/test/unit/test_cosmos3.py +++ b/python/sglang/multimodal_gen/test/unit/test_cosmos3.py @@ -1,6 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 """Unit tests for Cosmos3 config, weight mapping, and sampling params.""" +import dataclasses import importlib.util import json import types @@ -38,10 +39,9 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import ( VideoGenerationsRequest, ) from sglang.multimodal_gen.runtime.entrypoints.openai.video_api import ( - _cosmos3_sampling_param_kwargs, _multipart_video_extras, - _resolve_sound_duration, _resolve_video_path, + _video_request_model_kwargs, ) from sglang.multimodal_gen.runtime.loader.component_loaders import scheduler_loader from sglang.multimodal_gen.runtime.loader.component_loaders.scheduler_loader import ( @@ -887,6 +887,10 @@ class TestCosmos3OpenAIProtocol(unittest.TestCase): """Verify Cosmos3 modality knobs stay model-specific video extras.""" def test_cosmos3_template_fields_remain_extra_fields(self): + base_fields = {field.name for field in dataclasses.fields(SamplingParams)} + cosmos_fields = { + field.name for field in dataclasses.fields(Cosmos3SamplingParams) + } for request_cls in (ImageGenerationsRequest, VideoGenerationsRequest): with self.subTest(request_cls=request_cls.__name__): self.assertIn("max_sequence_length", request_cls.model_fields) @@ -895,6 +899,15 @@ class TestCosmos3OpenAIProtocol(unittest.TestCase): self.assertNotIn("use_resolution_template", request_cls.model_fields) self.assertNotIn("use_system_prompt", request_cls.model_fields) self.assertNotIn("use_guardrails", request_cls.model_fields) + for field_name in ( + "sound_duration", + "use_duration_template", + "use_resolution_template", + "use_system_prompt", + "use_guardrails", + ): + self.assertNotIn(field_name, base_fields) + self.assertIn(field_name, cosmos_fields) def test_cosmos3_modal_fields_are_model_specific_video_extras(self): for field_name in ( @@ -963,7 +976,9 @@ class TestCosmos3OpenAIProtocol(unittest.TestCase): self.assertEqual(_resolve_video_path(req), "https://example.com/input.mp4") - kwargs = _cosmos3_sampling_param_kwargs(req, num_frames=48, fps=24) + kwargs = _video_request_model_kwargs(req, Cosmos3SamplingParams) + kwargs.update(num_frames=48, fps=24) + kwargs = Cosmos3SamplingParams.lower_video_request_kwargs(req, kwargs) self.assertEqual(kwargs["sound_duration"], 2.0) self.assertEqual(kwargs["condition_frame_indexes"], [0, 2]) self.assertEqual(kwargs["condition_video_keep"], "last") @@ -1019,10 +1034,10 @@ class TestCosmos3OpenAIProtocol(unittest.TestCase): req = VideoGenerationsRequest( prompt="test", generate_sound=False, sound_duration=3.0 ) - self.assertEqual( - _resolve_sound_duration(req, num_frames=48, fps=24), - 0.0, - ) + kwargs = _video_request_model_kwargs(req, Cosmos3SamplingParams) + kwargs.update(num_frames=48, fps=24) + kwargs = Cosmos3SamplingParams.lower_video_request_kwargs(req, kwargs) + self.assertEqual(kwargs["sound_duration"], 0.0) class TestCosmos3Guardrails(unittest.TestCase): diff --git a/python/sglang/multimodal_gen/test/unit/test_openai_image_api.py b/python/sglang/multimodal_gen/test/unit/test_openai_image_api.py index e4498fbeb..9cd67f7a3 100644 --- a/python/sglang/multimodal_gen/test/unit/test_openai_image_api.py +++ b/python/sglang/multimodal_gen/test/unit/test_openai_image_api.py @@ -1,19 +1,32 @@ import os +from dataclasses import fields from fastapi import HTTPException from PIL import Image +from sglang.multimodal_gen.configs.sample.cosmos3 import Cosmos3SamplingParams +from sglang.multimodal_gen.configs.sample.ernie_image import ( + ErnieImageSamplingParams, +) from sglang.multimodal_gen.configs.sample.glmimage import GlmImageSamplingParams +from sglang.multimodal_gen.configs.sample.ideogram import Ideogram4SamplingParams +from sglang.multimodal_gen.configs.sample.longcat_image import ( + LongCatImageSamplingParams, +) from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams from sglang.multimodal_gen.runtime.entrypoints.openai.image_api import ( _build_image_response_kwargs, _fallback_image_urls, _get_response_resize, + _image_request_model_kwargs, _raise_if_image_variant_not_found, _runtime_sampling_quality, _select_image_variant_cloud_url, _select_image_variant_path, ) +from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import ( + ImageGenerationsRequest, +) from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch @@ -48,6 +61,82 @@ def test_runtime_sampling_quality_preserves_the_openai_default(): assert _runtime_sampling_quality("high") == "high" +def test_longcat_image_fields_remain_model_specific(): + field_values = { + "enable_cfg_renorm": False, + "cfg_renorm_min": 0.25, + "enable_prompt_rewrite": False, + } + request = ImageGenerationsRequest(prompt="a lantern", **field_values) + base_fields = {field.name for field in fields(SamplingParams)} + longcat_fields = {field.name for field in fields(LongCatImageSamplingParams)} + + for field_name, value in field_values.items(): + assert field_name not in ImageGenerationsRequest.model_fields + assert field_name not in base_fields + assert field_name in longcat_fields + assert getattr(request, field_name) == value + + assert _image_request_model_kwargs(request, LongCatImageSamplingParams) == ( + field_values + ) + assert _image_request_model_kwargs(request, SamplingParams) == {} + + +def test_longcat_image_fields_accept_nested_extra_body(): + request = ImageGenerationsRequest( + prompt="a lantern", + enable_prompt_rewrite=True, + extra_body={ + "enable_prompt_rewrite": False, + "enable_cfg_renorm": False, + "cfg_renorm_min": 0.5, + }, + ) + + assert _image_request_model_kwargs(request, LongCatImageSamplingParams) == { + "enable_prompt_rewrite": True, + "enable_cfg_renorm": False, + "cfg_renorm_min": 0.5, + } + + +def test_other_image_extensions_remain_model_specific(): + cases = ( + (Cosmos3SamplingParams, "guidance_interval", [400.0, 1000.0]), + (Cosmos3SamplingParams, "use_guardrails", False), + (ErnieImageSamplingParams, "use_pe", False), + (Ideogram4SamplingParams, "preset", "V4_TURBO_12"), + ) + base_fields = {field.name for field in fields(SamplingParams)} + + for sampling_params_cls, field_name, value in cases: + request = ImageGenerationsRequest( + prompt="a lantern", + extra_body={field_name: value}, + ) + model_fields = {field.name for field in fields(sampling_params_cls)} + + assert field_name not in ImageGenerationsRequest.model_fields + assert field_name not in base_fields + assert field_name in model_fields + assert _image_request_model_kwargs(request, sampling_params_cls) == { + field_name: value + } + assert _image_request_model_kwargs(request, SamplingParams) == {} + + +def test_cosmos_image_guardrails_alias_is_preserved(): + request = ImageGenerationsRequest( + prompt="a lantern", + extra_body={"guardrails": False}, + ) + + assert _image_request_model_kwargs(request, Cosmos3SamplingParams) == { + "use_guardrails": False + } + + def test_image_response_includes_resize_for_every_output(): response = _build_image_response_kwargs( ["first.png", "second.png"], diff --git a/python/sglang/multimodal_gen/test/unit/test_video_api_profiling.py b/python/sglang/multimodal_gen/test/unit/test_video_api_profiling.py index f1631d4a5..e2e14c3d6 100644 --- a/python/sglang/multimodal_gen/test/unit/test_video_api_profiling.py +++ b/python/sglang/multimodal_gen/test/unit/test_video_api_profiling.py @@ -1,6 +1,10 @@ +from dataclasses import fields from types import SimpleNamespace from unittest.mock import patch +from sglang.multimodal_gen.configs.sample.ltx_2 import LTX23SamplingParams +from sglang.multimodal_gen.configs.sample.ltx_2_5 import LTX25SamplingParams +from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import ( RealtimeVideoGenerationsRequest, VideoGenerationsRequest, @@ -11,7 +15,9 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_adapter ) from sglang.multimodal_gen.runtime.entrypoints.openai.video_api import ( _build_video_sampling_params, + _video_request_model_kwargs, ) +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req def test_video_api_forwards_profiling_options(): @@ -57,6 +63,35 @@ def test_video_api_forwards_profiling_options(): assert kwargs["quality"] == "high" +def test_ltx25_video_extensions_remain_model_specific(): + field_values = { + "use_diffusion_decoder": True, + "auto_duration": True, + "auto_duration_min_seconds": 2.0, + "auto_duration_max_seconds": 8.0, + } + request = VideoGenerationsRequest( + prompt="a fox in snow", + extra_body=field_values, + ) + base_fields = {field.name for field in fields(SamplingParams)} + ltx25_fields = {field.name for field in fields(LTX25SamplingParams)} + + for field_name in field_values: + assert field_name not in VideoGenerationsRequest.model_fields + assert field_name not in base_fields + assert field_name in ltx25_fields + + assert _video_request_model_kwargs(request, LTX25SamplingParams) == field_values + assert _video_request_model_kwargs(request, SamplingParams) == {} + + +def test_ltx23_request_defaults_to_vae_decoder(): + request = Req(sampling_params=LTX23SamplingParams()) + + assert request.use_diffusion_decoder is False + + def test_realtime_video_api_forwards_sampling_quality(): request = RealtimeVideoGenerationsRequest( type="init",