From 7662b8b919be05f285e6110634ea6d58863449c0 Mon Sep 17 00:00:00 2001 From: Junhao Liu Date: Sun, 8 Mar 2026 11:06:40 -0700 Subject: [PATCH] [diffusion] feat: implement upscaling (#19723) --- docs/diffusion/api/cli.md | 31 +- docs/diffusion/api/post_processing.md | 148 ++++++ docs/diffusion/index.md | 1 + .../configs/sample/sampling_params.py | 32 +- .../entrypoints/diffusion_generator.py | 3 + .../runtime/entrypoints/http_server.py | 3 + .../runtime/entrypoints/openai/image_api.py | 9 + .../runtime/entrypoints/openai/protocol.py | 8 + .../runtime/entrypoints/openai/utils.py | 3 + .../runtime/entrypoints/openai/video_api.py | 9 + .../runtime/entrypoints/utils.py | 21 +- .../runtime/managers/gpu_worker.py | 3 + .../runtime/postprocess/__init__.py | 13 +- .../postprocess/realesrgan_upscaler.py | 484 ++++++++++++++++++ .../test/server/perf_baselines.json | 196 +++++++ .../test/server/test_server_utils.py | 42 +- .../test/server/testcase_configs.py | 62 ++- 17 files changed, 1005 insertions(+), 63 deletions(-) create mode 100644 docs/diffusion/api/post_processing.md create mode 100644 python/sglang/multimodal_gen/runtime/postprocess/realesrgan_upscaler.py diff --git a/docs/diffusion/api/cli.md b/docs/diffusion/api/cli.md index f0f8d7da4..6af075943 100644 --- a/docs/diffusion/api/cli.md +++ b/docs/diffusion/api/cli.md @@ -41,33 +41,12 @@ The SGLang-diffusion CLI provides a quick way to access the inference pipeline f - `--fps {FPS}`: Frames per second for the saved output, if this is a video-generation task -**Frame Interpolation** (video only) +**Post-Processing** (frame interpolation & upscaling) -Frame interpolation is a post-processing step that synthesizes new frames -between each pair of consecutive generated frames, producing smoother -motion without re-running the diffusion model. The `--frame-interpolation-exp` -flag controls how many rounds of interpolation to apply: each round inserts one -new frame into every gap between adjacent frames, so the output frame count -follows the formula **(N − 1) × 2^exp + 1** (e.g. 5 original frames with -`exp=1` → 4 gaps × 1 new frame + 5 originals = **9** frames; with `exp=2` → -**17** frames). - -- `--enable-frame-interpolation`: Enable frame interpolation. Model weights are downloaded automatically on first use. -- `--frame-interpolation-exp {EXP}`: Interpolation exponent — `1` = 2× temporal resolution, `2` = 4×, etc. (default: `1`) -- `--frame-interpolation-scale {SCALE}`: RIFE inference scale; use `0.5` for high-resolution inputs to save memory (default: `1.0`) -- `--frame-interpolation-model-path {PATH}`: Local directory or HuggingFace repo ID containing RIFE `flownet.pkl` weights (default: `elfgum/RIFE-4.22.lite`, downloaded automatically) - -Example — generate a 5-frame video and interpolate to 9 frames ((5 − 1) × 2¹ + 1 = 9): - -```bash -sglang generate \ - --model-path Wan-AI/Wan2.2-T2V-A14B-Diffusers \ - --prompt "A dog running through a park" \ - --num-frames 5 \ - --enable-frame-interpolation \ - --frame-interpolation-exp 1 \ - --save-output -``` +SGLang diffusion supports optional post-processing steps — frame interpolation +(RIFE) for smoother video and upscaling (Real-ESRGAN) for higher resolution. +See the dedicated **[Post-Processing](post_processing.md)** page for full +details, supported models, and examples. **Output Options** diff --git a/docs/diffusion/api/post_processing.md b/docs/diffusion/api/post_processing.md new file mode 100644 index 000000000..d832f4af2 --- /dev/null +++ b/docs/diffusion/api/post_processing.md @@ -0,0 +1,148 @@ +# Post-Processing + +SGLang diffusion supports optional post-processing steps that run after +generation to improve temporal smoothness (frame interpolation) or spatial +resolution (upscaling). These steps are independent of the diffusion model and +can be combined in a single run. + +When both are enabled, **frame interpolation runs first** (increasing the frame +count), then **upscaling runs on every frame** (increasing the spatial +resolution). + +--- + +## Frame Interpolation (video only) + +Frame interpolation synthesizes new frames between each pair of consecutive +generated frames, producing smoother motion without re-running the diffusion +model. + +The `--frame-interpolation-exp` flag controls how many rounds of interpolation +to apply: each round inserts one new frame into every gap between adjacent +frames, so the output frame count follows the formula: + +> **(N − 1) × 2^exp + 1** +> +> e.g. 5 original frames with `exp=1` → 4 gaps × 1 new frame + 5 originals = **9** frames; +> with `exp=2` → **17** frames. + +### CLI Arguments + +| Argument | Description | +|----------|-------------| +| `--enable-frame-interpolation` | Enable frame interpolation. Model weights are downloaded automatically on first use. | +| `--frame-interpolation-exp {EXP}` | Interpolation exponent — `1` = 2× temporal resolution, `2` = 4×, etc. (default: `1`) | +| `--frame-interpolation-scale {SCALE}` | RIFE inference scale; use `0.5` for high-resolution inputs to save memory (default: `1.0`) | +| `--frame-interpolation-model-path {PATH}` | Local directory or HuggingFace repo ID containing RIFE `flownet.pkl` weights (default: `elfgum/RIFE-4.22.lite`, downloaded automatically) | + +### Supported Models + +Frame interpolation uses the [RIFE](https://github.com/hzwer/Practical-RIFE) +(Real-Time Intermediate Flow Estimation) architecture. Only **RIFE 4.22.lite** +(`IFNet` with 4-scale `IFBlock` backbone) is supported. The network topology is +hard-coded, so custom weights provided via `--frame-interpolation-model-path` +must be a `flownet.pkl` checkpoint that is compatible with this architecture. + +Other RIFE versions (e.g., older `v4.x` variants with different block counts) +or entirely different frame interpolation methods (FILM, AMT, etc.) are **not +supported**. + +| Weight | HuggingFace Repo | Description | +|--------|------------------|-------------| +| RIFE 4.22.lite *(default)* | [`elfgum/RIFE-4.22.lite`](https://huggingface.co/elfgum/RIFE-4.22.lite) | Lightweight model, downloaded automatically on first use | + +### Example + +Generate a 5-frame video and interpolate to 9 frames ((5 − 1) × 2¹ + 1 = 9): + +```bash +sglang generate \ + --model-path Wan-AI/Wan2.2-T2V-A14B-Diffusers \ + --prompt "A dog running through a park" \ + --num-frames 5 \ + --enable-frame-interpolation \ + --frame-interpolation-exp 1 \ + --save-output +``` + +--- + +## Upscaling (image and video) + +Upscaling increases the spatial resolution of generated images or video frames +using [Real-ESRGAN](https://github.com/xinntao/Real-ESRGAN). The model weights +are downloaded automatically on first use and cached for subsequent runs. + +### CLI Arguments + +| Argument | Description | +|----------|-------------| +| `--enable-upscaling` | Enable post-generation upscaling using Real-ESRGAN. | +| `--upscaling-scale {SCALE}` | Desired upscaling factor (default: `4`). The 4× model is used internally; if a different scale is requested, a bicubic resize is applied after the network output. | +| `--upscaling-model-path {PATH}` | Local `.pth` file, HuggingFace repo ID, or `repo_id:filename` for Real-ESRGAN weights (default: `ai-forever/Real-ESRGAN` with `RealESRGAN_x4.pth`, downloaded automatically). Use the `repo_id:filename` format to specify a custom weight file from a HuggingFace repo (e.g. `my-org/my-esrgan:weights.pth`). | + +### Supported Models + +Upscaling supports two Real-ESRGAN network architectures. The correct +architecture is **auto-detected** from the checkpoint keys, so you only need to +point `--upscaling-model-path` at a valid `.pth` file: + +| Architecture | Example Weights | Description | +|--------------|-----------------|-------------| +| **RRDBNet** | `RealESRGAN_x4plus.pth` | Heavier model with higher quality; best for photos | +| **SRVGGNetCompact** | `RealESRGAN_x4.pth` *(default)*, `realesr-animevideov3.pth`, `realesr-general-x4v3.pth` | Lightweight model; faster inference, good for video | + +The default weight file is +[`ai-forever/Real-ESRGAN`](https://huggingface.co/ai-forever/Real-ESRGAN) with +`RealESRGAN_x4.pth` (SRVGGNetCompact, 4× native scale). + +Other super-resolution models (e.g., SwinIR, HAT, BSRGAN) are **not supported** +— only Real-ESRGAN checkpoints using the two architectures above are +compatible. + +### Examples + +Generate a 1024×1024 image and upscale to 4096×4096: + +```bash +sglang generate \ + --model-path black-forest-labs/FLUX.2-dev \ + --prompt "A cat sitting on a windowsill" \ + --output-size 1024x1024 \ + --enable-upscaling \ + --save-output +``` + +Generate a video and upscale each frame by 4×: + +```bash +sglang generate \ + --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \ + --prompt "A curious raccoon" \ + --enable-upscaling \ + --upscaling-scale 4 \ + --save-output +``` + +--- + +## Combining Frame Interpolation and Upscaling + +Frame interpolation and upscaling can be combined in a single run. +Interpolation is applied first (increasing the frame count), then upscaling is +applied to every frame (increasing the spatial resolution). + +Example — generate 5 frames, interpolate to 9 frames, and upscale each frame +by 4×: + +```bash +sglang generate \ + --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \ + --prompt "A curious raccoon" \ + --num-frames 5 \ + --enable-frame-interpolation \ + --frame-interpolation-exp 1 \ + --enable-upscaling \ + --upscaling-scale 4 \ + --save-output +``` diff --git a/docs/diffusion/index.md b/docs/diffusion/index.md index 3a1aa815a..6637ed785 100644 --- a/docs/diffusion/index.md +++ b/docs/diffusion/index.md @@ -50,6 +50,7 @@ sglang serve --model-path Qwen/Qwen-Image --port 30010 - **[CLI Documentation](api/cli.md)** - Command-line interface for `sglang generate` and `sglang serve` - **[OpenAI API](api/openai_api.md)** - OpenAI-compatible API for image/video generation and LoRA management +- **[Post-Processing](api/post_processing.md)** - Frame interpolation (RIFE) and upscaling (Real-ESRGAN) ### Performance Optimization diff --git a/python/sglang/multimodal_gen/configs/sample/sampling_params.py b/python/sglang/multimodal_gen/configs/sample/sampling_params.py index 2f79a38fa..42eba884d 100644 --- a/python/sglang/multimodal_gen/configs/sample/sampling_params.py +++ b/python/sglang/multimodal_gen/configs/sample/sampling_params.py @@ -113,6 +113,13 @@ class SamplingParams: None # local dir or HF repo ID with flownet.pkl (default: elfgum/RIFE-4.22.lite) ) + # Upscaling + enable_upscaling: bool = False + upscaling_model_path: str | None = ( + None # local .pth, HF repo ID, or repo_id:filename (default: ai-forever/Real-ESRGAN) + ) + upscaling_scale: int = 4 + # Batch info num_outputs_per_prompt: int = 1 seed: int = 42 @@ -875,7 +882,30 @@ class SamplingParams: type=str, default=SamplingParams.frame_interpolation_model_path, help="Local directory or HuggingFace repo ID containing RIFE flownet.pkl weights " - "(default: elfgum/RIFE-4.22.lite, downloaded automatically).", + "(default: elfgum/RIFE-4.22.lite, downloaded automatically). " + "Only RIFE 4.22.lite architecture is supported; other RIFE versions or " + "frame interpolation models are not compatible.", + ) + parser.add_argument( + "--enable-upscaling", + action="store_true", + help="Enable post-generation upscaling using Real-ESRGAN.", + ) + parser.add_argument( + "--upscaling-model-path", + type=str, + default=SamplingParams.upscaling_model_path, + help="Local .pth file, HuggingFace repo ID, or repo_id:filename for Real-ESRGAN weights " + "(default: ai-forever/Real-ESRGAN with RealESRGAN_x4.pth). " + "Only RRDBNet (e.g. RealESRGAN_x4plus) and SRVGGNetCompact (e.g. realesr-animevideov3) " + "architectures are supported; other super-resolution models are not compatible. " + "Use 'repo_id:filename' to specify a custom weight file from a HF repo.", + ) + parser.add_argument( + "--upscaling-scale", + type=int, + default=SamplingParams.upscaling_scale, + help="Upscaling factor (default: 4).", ) return parser diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py index 162df5f89..b37e1f8bb 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py @@ -277,6 +277,9 @@ class DiffGenerator: frame_interpolation_exp=req.frame_interpolation_exp, frame_interpolation_scale=req.frame_interpolation_scale, frame_interpolation_model_path=req.frame_interpolation_model_path, + enable_upscaling=req.enable_upscaling, + upscaling_model_path=req.upscaling_model_path, + upscaling_scale=req.upscaling_scale, ) for idx in range(len(samples_out)): diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py index 4830411bf..d07febdb1 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py @@ -203,6 +203,9 @@ async def forward_to_scheduler( frame_interpolation_exp=sp.frame_interpolation_exp, frame_interpolation_scale=sp.frame_interpolation_scale, frame_interpolation_model_path=sp.frame_interpolation_model_path, + enable_upscaling=sp.enable_upscaling, + upscaling_model_path=sp.upscaling_model_path, + upscaling_scale=sp.upscaling_scale, ) if hasattr(response, "model_dump"): 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 d861cda72..116208451 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py @@ -133,6 +133,9 @@ async def generations( enable_teacache=request.enable_teacache, output_compression=request.output_compression, output_quality=request.output_quality, + enable_upscaling=request.enable_upscaling, + upscaling_model_path=request.upscaling_model_path, + upscaling_scale=request.upscaling_scale, ) batch = prepare_request( server_args=server_args, @@ -207,6 +210,9 @@ async def edits( output_quality: Optional[str] = Form("default"), output_compression: Optional[int] = Form(None), enable_teacache: Optional[bool] = Form(False), + enable_upscaling: Optional[bool] = Form(False), + upscaling_model_path: Optional[str] = Form(None), + upscaling_scale: Optional[int] = Form(4), num_frames: int = Form(1), ): request_id = generate_request_id() @@ -262,6 +268,9 @@ async def edits( num_frames=num_frames, output_compression=output_compression, output_quality=output_quality, + enable_upscaling=enable_upscaling, + upscaling_model_path=upscaling_model_path, + upscaling_scale=upscaling_scale, ) batch = prepare_request( server_args=server_args, diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py index 9d1a2d2cb..425a17f7b 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py @@ -48,6 +48,10 @@ class ImageGenerationsRequest(BaseModel): output_quality: Optional[str] = "default" output_compression: Optional[int] = None enable_teacache: Optional[bool] = False + # Upscaling + enable_upscaling: Optional[bool] = False + upscaling_model_path: Optional[str] = None + upscaling_scale: Optional[int] = 4 diffusers_kwargs: Optional[Dict[str, Any]] = None # kwargs for diffusers backend @@ -97,6 +101,10 @@ class VideoGenerationsRequest(BaseModel): frame_interpolation_exp: Optional[int] = 1 # 1=2×, 2=4× frame_interpolation_scale: Optional[float] = 1.0 frame_interpolation_model_path: Optional[str] = None + # Upscaling + enable_upscaling: Optional[bool] = False + upscaling_model_path: Optional[str] = None + upscaling_scale: Optional[int] = 4 output_quality: Optional[str] = "default" output_compression: Optional[int] = None output_path: Optional[str] = None diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py index 86b84c9cc..29a944f59 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py @@ -286,6 +286,9 @@ async def process_generation_batch( frame_interpolation_exp=batch.frame_interpolation_exp, frame_interpolation_scale=batch.frame_interpolation_scale, frame_interpolation_model_path=batch.frame_interpolation_model_path, + enable_upscaling=batch.enable_upscaling, + upscaling_model_path=batch.upscaling_model_path, + upscaling_scale=batch.upscaling_scale, ) total_time = time.perf_counter() - total_start_time 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 9db0fde3c..d85ce7d40 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py @@ -75,6 +75,9 @@ def _build_video_sampling_params(request_id: str, request: VideoGenerationsReque frame_interpolation_exp=request.frame_interpolation_exp, frame_interpolation_scale=request.frame_interpolation_scale, frame_interpolation_model_path=request.frame_interpolation_model_path, + enable_upscaling=request.enable_upscaling, + upscaling_model_path=request.upscaling_model_path, + upscaling_scale=request.upscaling_scale, output_path=request.output_path, output_compression=request.output_compression, output_quality=request.output_quality, @@ -181,6 +184,9 @@ async def create_video( frame_interpolation_exp: Optional[int] = Form(1), frame_interpolation_scale: Optional[float] = Form(1.0), frame_interpolation_model_path: Optional[str] = Form(None), + enable_upscaling: Optional[bool] = Form(False), + upscaling_model_path: Optional[str] = Form(None), + upscaling_scale: Optional[int] = Form(4), output_quality: Optional[str] = Form("default"), output_compression: Optional[int] = Form(None), extra_body: Optional[str] = Form(None), @@ -256,6 +262,9 @@ async def create_video( frame_interpolation_exp=frame_interpolation_exp, frame_interpolation_scale=frame_interpolation_scale, frame_interpolation_model_path=frame_interpolation_model_path, + enable_upscaling=enable_upscaling, + upscaling_model_path=upscaling_model_path, + upscaling_scale=upscaling_scale, output_compression=output_compression, output_quality=output_quality, **( diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py index 3d9d8b45b..21b6e1a1f 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py @@ -348,6 +348,9 @@ def save_outputs( frame_interpolation_exp: int = 1, frame_interpolation_scale: float = 1.0, frame_interpolation_model_path: Optional[str] = None, + enable_upscaling: bool = False, + upscaling_model_path: Optional[str] = None, + upscaling_scale: int = 4, ) -> list[str]: """Save outputs to files and return the list of file paths.""" output_paths: list[str] = [] @@ -369,6 +372,9 @@ def save_outputs( frame_interpolation_exp=frame_interpolation_exp, frame_interpolation_scale=frame_interpolation_scale, frame_interpolation_model_path=frame_interpolation_model_path, + enable_upscaling=enable_upscaling, + upscaling_model_path=upscaling_model_path, + upscaling_scale=upscaling_scale, ) if samples_out is not None: @@ -401,6 +407,9 @@ def post_process_sample( frame_interpolation_exp: int = 1, frame_interpolation_scale: float = 1.0, frame_interpolation_model_path: Optional[str] = None, + enable_upscaling: bool = False, + upscaling_model_path: Optional[str] = None, + upscaling_scale: int = 4, ): """ Process sample output, optionally interpolate video frames, and save. @@ -456,7 +465,17 @@ def post_process_sample( ) fps = fps * multiplier - # 3. Save outputs if requested + # 3. Upscaling (images and videos) + if enable_upscaling and frames: + from sglang.multimodal_gen.runtime.postprocess import upscale_frames + + frames = upscale_frames( + frames, + model_path=upscaling_model_path, + scale=upscaling_scale, + ) + + # 4. Save outputs if requested if save_output: if save_file_path: os.makedirs(os.path.dirname(save_file_path), exist_ok=True) diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 7b357330f..2b2d4eeb8 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -269,6 +269,9 @@ class GPUWorker: frame_interpolation_exp=req.frame_interpolation_exp, frame_interpolation_scale=req.frame_interpolation_scale, frame_interpolation_model_path=req.frame_interpolation_model_path, + enable_upscaling=req.enable_upscaling, + upscaling_model_path=req.upscaling_model_path, + upscaling_scale=req.upscaling_scale, ) output_batch.output_file_paths = output_paths output_batch.output = None diff --git a/python/sglang/multimodal_gen/runtime/postprocess/__init__.py b/python/sglang/multimodal_gen/runtime/postprocess/__init__.py index f70951a2d..139f98766 100644 --- a/python/sglang/multimodal_gen/runtime/postprocess/__init__.py +++ b/python/sglang/multimodal_gen/runtime/postprocess/__init__.py @@ -1,9 +1,18 @@ # SPDX-License-Identifier: Apache-2.0 -"""Frame interpolation support for SGLang diffusion pipelines.""" +"""Frame interpolation and upscaling support for SGLang diffusion pipelines.""" +from sglang.multimodal_gen.runtime.postprocess.realesrgan_upscaler import ( + ImageUpscaler, + upscale_frames, +) from sglang.multimodal_gen.runtime.postprocess.rife_interpolator import ( FrameInterpolator, interpolate_video_frames, ) -__all__ = ["FrameInterpolator", "interpolate_video_frames"] +__all__ = [ + "FrameInterpolator", + "interpolate_video_frames", + "ImageUpscaler", + "upscale_frames", +] diff --git a/python/sglang/multimodal_gen/runtime/postprocess/realesrgan_upscaler.py b/python/sglang/multimodal_gen/runtime/postprocess/realesrgan_upscaler.py new file mode 100644 index 000000000..e01dcfc53 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/postprocess/realesrgan_upscaler.py @@ -0,0 +1,484 @@ +# SPDX-License-Identifier: Apache-2.0 +""" +Real-ESRGAN upscaling for SGLang diffusion pipelines. + +Real-ESRGAN model code is vendored and adapted from: + - https://github.com/xinntao/Real-ESRGAN (BSD-3-Clause License) + Copyright (c) 2021 xinntao + +The ImageUpscaler wrapper and integration code are original work. +""" + +import math +import os +from typing import Optional + +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F + +from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + +# Default HuggingFace repo and filename for Real-ESRGAN weights +_DEFAULT_REALESRGAN_HF_REPO = "ai-forever/Real-ESRGAN" +_DEFAULT_REALESRGAN_FILENAME = "RealESRGAN_x4.pth" + +# Module-level cache: model_path -> UpscalerModel instance +_MODEL_CACHE: dict[str, "UpscalerModel"] = {} + + +# --------------------------------------------------------------------------- +# Vendored Real-ESRGAN architecture code +# (SRVGGNetCompact, ResidualDenseBlock, RRDB, RRDBNet) +# --------------------------------------------------------------------------- + + +class SRVGGNetCompact(nn.Module): + """Compact VGG-style network for super resolution. + + Corresponds to ``realesr-animevideov3`` and ``realesr-general-x4v3``. + Reference: xinntao/Real-ESRGAN (BSD-3-Clause). + """ + + def __init__( + self, + num_in_ch: int = 3, + num_out_ch: int = 3, + num_feat: int = 64, + num_conv: int = 16, + upscale: int = 4, + act_type: str = "prelu", + ): + super().__init__() + self.num_in_ch = num_in_ch + self.num_out_ch = num_out_ch + self.num_feat = num_feat + self.num_conv = num_conv + self.upscale = upscale + self.act_type = act_type + + self.body = nn.ModuleList() + # first conv + self.body.append(nn.Conv2d(num_in_ch, num_feat, 3, 1, 1)) + # first activation + self.body.append(self._make_act(act_type, num_feat)) + # body convs + activations + for _ in range(num_conv): + self.body.append(nn.Conv2d(num_feat, num_feat, 3, 1, 1)) + self.body.append(self._make_act(act_type, num_feat)) + # last conv: maps to out_ch * upscale^2 for pixel shuffle + self.body.append(nn.Conv2d(num_feat, num_out_ch * upscale * upscale, 3, 1, 1)) + self.upsampler = nn.PixelShuffle(upscale) + + @staticmethod + def _make_act(act_type: str, num_feat: int) -> nn.Module: + if act_type == "relu": + return nn.ReLU(inplace=True) + elif act_type == "prelu": + return nn.PReLU(num_parameters=num_feat) + elif act_type == "leakyrelu": + return nn.LeakyReLU(negative_slope=0.1, inplace=True) + else: + raise ValueError(f"Unsupported activation type: {act_type}") + + def forward(self, x: torch.Tensor) -> torch.Tensor: + out = x + for layer in self.body: + out = layer(out) + out = self.upsampler(out) + # residual addition with nearest upsampled input + base = F.interpolate(x, scale_factor=self.upscale, mode="nearest") + return out + base + + +class ResidualDenseBlock(nn.Module): + """Residual Dense Block used in RRDB (RealESRGAN_x4plus).""" + + def __init__(self, num_feat: int = 64, num_grow_ch: int = 32): + super().__init__() + self.conv1 = nn.Conv2d(num_feat, num_grow_ch, 3, 1, 1) + self.conv2 = nn.Conv2d(num_feat + num_grow_ch, num_grow_ch, 3, 1, 1) + self.conv3 = nn.Conv2d(num_feat + 2 * num_grow_ch, num_grow_ch, 3, 1, 1) + self.conv4 = nn.Conv2d(num_feat + 3 * num_grow_ch, num_grow_ch, 3, 1, 1) + self.conv5 = nn.Conv2d(num_feat + 4 * num_grow_ch, num_feat, 3, 1, 1) + self.lrelu = nn.LeakyReLU(negative_slope=0.2, inplace=True) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x1 = self.lrelu(self.conv1(x)) + x2 = self.lrelu(self.conv2(torch.cat((x, x1), 1))) + x3 = self.lrelu(self.conv3(torch.cat((x, x1, x2), 1))) + x4 = self.lrelu(self.conv4(torch.cat((x, x1, x2, x3), 1))) + x5 = self.conv5(torch.cat((x, x1, x2, x3, x4), 1)) + return x5 * 0.2 + x + + +class RRDB(nn.Module): + """Residual in Residual Dense Block.""" + + def __init__(self, num_feat: int, num_grow_ch: int = 32): + super().__init__() + self.rdb1 = ResidualDenseBlock(num_feat, num_grow_ch) + self.rdb2 = ResidualDenseBlock(num_feat, num_grow_ch) + self.rdb3 = ResidualDenseBlock(num_feat, num_grow_ch) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + out = self.rdb1(x) + out = self.rdb2(out) + out = self.rdb3(out) + return out * 0.2 + x + + +class RRDBNet(nn.Module): + """RRDB network for RealESRGAN_x4plus (heavier, higher quality for photos).""" + + def __init__( + self, + num_in_ch: int = 3, + num_out_ch: int = 3, + scale: int = 4, + num_feat: int = 64, + num_block: int = 23, + num_grow_ch: int = 32, + ): + super().__init__() + self.scale = scale + in_ch = num_in_ch + if scale == 2: + in_ch = num_in_ch * 4 + elif scale == 1: + in_ch = num_in_ch * 16 + self.conv_first = nn.Conv2d(in_ch, num_feat, 3, 1, 1) + self.body = nn.Sequential( + *[RRDB(num_feat, num_grow_ch) for _ in range(num_block)] + ) + self.conv_body = nn.Conv2d(num_feat, num_feat, 3, 1, 1) + # upsample + self.conv_up1 = nn.Conv2d(num_feat, num_feat, 3, 1, 1) + self.conv_up2 = nn.Conv2d(num_feat, num_feat, 3, 1, 1) + self.conv_hr = nn.Conv2d(num_feat, num_feat, 3, 1, 1) + self.conv_last = nn.Conv2d(num_feat, num_out_ch, 3, 1, 1) + self.lrelu = nn.LeakyReLU(negative_slope=0.2, inplace=True) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.scale == 2: + feat = F.pixel_unshuffle(x, 2) + elif self.scale == 1: + feat = F.pixel_unshuffle(x, 4) + else: + feat = x + feat = self.conv_first(feat) + body_feat = self.conv_body(self.body(feat)) + feat = feat + body_feat + feat = self.lrelu( + self.conv_up1(F.interpolate(feat, scale_factor=2, mode="nearest")) + ) + feat = self.lrelu( + self.conv_up2(F.interpolate(feat, scale_factor=2, mode="nearest")) + ) + return self.conv_last(self.lrelu(self.conv_hr(feat))) + + +# --------------------------------------------------------------------------- +# Architecture auto-detection +# --------------------------------------------------------------------------- + + +def _build_net_from_state_dict(state_dict: dict) -> nn.Module: + """Detect architecture from checkpoint keys and return an unloaded network.""" + if "conv_first.weight" in state_dict: + # RRDBNet (e.g., RealESRGAN_x4plus) + num_feat = state_dict["conv_first.weight"].shape[0] + num_block = sum( + 1 + for k in state_dict + if k.startswith("body.") and k.endswith(".rdb1.conv1.weight") + ) + num_grow_ch = state_dict["body.0.rdb1.conv1.weight"].shape[0] + logger.info( + "Detected RRDBNet: num_feat=%d, num_block=%d, num_grow_ch=%d", + num_feat, + num_block, + num_grow_ch, + ) + return RRDBNet( + num_in_ch=3, + num_out_ch=3, + scale=4, + num_feat=num_feat, + num_block=num_block, + num_grow_ch=num_grow_ch, + ) + else: + # SRVGGNetCompact (e.g., realesr-animevideov3) + num_feat = state_dict["body.0.weight"].shape[0] + # body layout: [first_conv, first_act, (conv, act)*num_conv, last_conv] + # count 4-D weight tensors = first_conv + loop_convs + last_conv = num_conv + 2 + conv_keys = sorted( + [ + k + for k in state_dict + if k.startswith("body.") + and k.endswith(".weight") + and state_dict[k].dim() == 4 + ], + key=lambda k: int(k.split(".")[1]), + ) + num_conv = len(conv_keys) - 2 # subtract first and last + # upscale from last conv output channels: out_ch = num_out_ch * upscale^2 + last_out_ch = state_dict[conv_keys[-1]].shape[0] + upscale = int(math.sqrt(last_out_ch / 3)) + logger.info( + "Detected SRVGGNetCompact: num_feat=%d, num_conv=%d, upscale=%d", + num_feat, + num_conv, + upscale, + ) + return SRVGGNetCompact( + num_in_ch=3, + num_out_ch=3, + num_feat=num_feat, + num_conv=num_conv, + upscale=upscale, + act_type="prelu", + ) + + +# --------------------------------------------------------------------------- +# UpscalerModel +# --------------------------------------------------------------------------- + + +class UpscalerModel: + """Wraps a Real-ESRGAN network, provides load() and upscale() API.""" + + def __init__(self, net: nn.Module, scale: int): + self.net = net + self.scale = scale # the model's native upscaling factor (e.g. 4) + + @property + def device(self) -> torch.device: + return next(self.net.parameters()).device + + def upscale(self, frame: np.ndarray, outscale: float | None = None) -> np.ndarray: + """Upscale a single HWC uint8 frame → HWC uint8 frame. + + Args: + frame: Input HWC uint8 numpy array. + outscale: Desired final upscaling factor. If different from the + model's native scale, a cheap resize is applied after + the network output (same approach as the official + Real-ESRGAN ``inference_realesrgan.py --outscale``). + ``None`` means use the model's native scale as-is. + """ + h, w = frame.shape[:2] + img = frame.astype(np.float32) / 255.0 + img_t = torch.from_numpy(img).permute(2, 0, 1).unsqueeze(0).to(self.device) + with torch.no_grad(): + out = self.net(img_t) + + # If the desired outscale differs from the model's native scale, + # resize to (h * outscale, w * outscale). + if outscale is not None and outscale != self.scale: + target_h = int(h * outscale) + target_w = int(w * outscale) + out = F.interpolate( + out, size=(target_h, target_w), mode="bicubic", align_corners=False + ) + + out_np = out.squeeze(0).permute(1, 2, 0).clamp(0.0, 1.0).cpu().numpy() + return (out_np * 255.0).astype(np.uint8) + + +# --------------------------------------------------------------------------- +# ImageUpscaler public class +# --------------------------------------------------------------------------- + + +class ImageUpscaler: + """ + Lazy-loaded Real-ESRGAN upscaler. + + Weights are downloaded and cached on first call to `.upscale()`. + Supports both SRVGGNetCompact (lightweight, default) and RRDBNet (heavier). + """ + + def __init__( + self, + model_path: Optional[str] = None, + scale: int = 4, + half_precision: bool = False, + ): + self._model_path = model_path + self._scale = scale + self._half_precision = half_precision + + def _ensure_model_loaded(self) -> UpscalerModel: + """Download/load Real-ESRGAN weights, detect arch, and cache globally.""" + model_path = self._model_path or _DEFAULT_REALESRGAN_HF_REPO + + # Resolve: local .pth pass-through, or HF repo → download single file + resolved_path = _resolve_model_path(model_path) + + if resolved_path in _MODEL_CACHE: + return _MODEL_CACHE[resolved_path] + + logger.info("Loading Real-ESRGAN weights from %s", resolved_path) + try: + state_dict = torch.load( + resolved_path, map_location="cpu", weights_only=True + ) + except Exception as e: + raise RuntimeError( + f"Failed to load Real-ESRGAN checkpoint from '{resolved_path}'. " + f"The file may be corrupted or not a valid PyTorch checkpoint. " + f"Original error: {e}" + ) from e + + # Some checkpoints wrap weights under a 'params' or 'params_ema' key + if "params_ema" in state_dict: + state_dict = state_dict["params_ema"] + elif "params" in state_dict: + state_dict = state_dict["params"] + + try: + net = _build_net_from_state_dict(state_dict) + net.load_state_dict(state_dict, strict=True) + except (RuntimeError, KeyError) as e: + raise RuntimeError( + f"Real-ESRGAN weight file '{resolved_path}' is not compatible " + f"with the supported architectures (SRVGGNetCompact / RRDBNet). " + f"Please ensure you are using a valid Real-ESRGAN checkpoint. " + f"Original error: {e}" + ) from e + net.eval() + + device = current_platform.get_local_torch_device() + if self._half_precision: + net = net.half() + net = net.to(device) + + # Detect the model's native scale from network architecture + native_scale = 4 # sensible default + if hasattr(net, "upscale"): + native_scale = net.upscale + elif hasattr(net, "scale"): + native_scale = net.scale + + model = UpscalerModel(net=net, scale=native_scale) + _MODEL_CACHE[resolved_path] = model + logger.info( + "Real-ESRGAN model loaded on device: %s (native_scale=%dx, outscale=%s)", + device, + native_scale, + f"{self._scale}x" if self._scale != native_scale else "native", + ) + return model + + def upscale(self, frames: list[np.ndarray]) -> list[np.ndarray]: + """Upscale a list of HWC uint8 frames. + + Uses the model's native scale for super-resolution, then resizes to + the desired ``outscale`` if it differs (cheap bicubic resize). + """ + if not frames: + return frames + model = self._ensure_model_loaded() + outscale = self._scale if self._scale != model.scale else None + return [model.upscale(frame, outscale=outscale) for frame in frames] + + +# --------------------------------------------------------------------------- +# HF download helper +# --------------------------------------------------------------------------- + + +def _resolve_model_path(model_path: str) -> str: + """Return a local .pth file path. + + Accepts: + - An existing local file path (pass-through). + - A HuggingFace ``repo_id`` → downloads the default weight file + (``RealESRGAN_x4.pth``). + - A HuggingFace ``repo_id:filename`` → downloads *filename* from *repo_id*, + allowing users to specify custom weight files hosted on HF. + """ + if os.path.isfile(model_path): + return model_path + + # Parse optional "repo_id:filename" syntax; fall back to default filename. + if ":" in model_path and not model_path.startswith("/"): + repo_id, filename = model_path.split(":", 1) + else: + repo_id = model_path + filename = _DEFAULT_REALESRGAN_FILENAME + + try: + from huggingface_hub import hf_hub_download + except ImportError as e: + raise ImportError( + "huggingface_hub is required to download Real-ESRGAN weights. " + "Install it with: pip install huggingface_hub" + ) from e + + logger.info( + "Downloading Real-ESRGAN weights from HF repo %s (file: %s)", + repo_id, + filename, + ) + try: + local_path = hf_hub_download( + repo_id=repo_id, + filename=filename, + ) + except Exception as e: + raise FileNotFoundError( + f"Failed to download Real-ESRGAN weights from HuggingFace repo " + f"'{repo_id}' (file: '{filename}'). If you are using a custom " + f"model, provide either a local .pth file path or use the " + f"'repo_id:filename' format (e.g. 'my-org/my-esrgan:weights.pth'). " + f"Original error: {e}" + ) from e + return local_path + + +# --------------------------------------------------------------------------- +# Module-level convenience function +# --------------------------------------------------------------------------- + + +def upscale_frames( + frames: list[np.ndarray], + model_path: Optional[str] = None, + scale: int = 4, + half_precision: bool = False, +) -> list[np.ndarray]: + """ + Convenience wrapper around ImageUpscaler. + + The model always runs at its native resolution (e.g. 4× for + ``RealESRGAN_x4.pth``). If *scale* differs from the native factor, + a cheap bicubic resize is applied after the network output – the same + approach used by the official Real-ESRGAN ``--outscale`` flag. + + Args: + frames: List of uint8 HWC numpy frames. + model_path: Local .pth file, HuggingFace repo ID, or + ``repo_id:filename`` for a custom weight file. + None → default ``ai-forever/Real-ESRGAN`` with + ``RealESRGAN_x4.pth``. + scale: Desired final upscaling factor (e.g. 2, 3, 4). + The 4× model is used internally; the output is + resized to match *scale* when it differs. + half_precision: Use fp16 inference (faster on supported GPUs). + + Returns: + List of upscaled uint8 HWC numpy frames. + """ + upscaler = ImageUpscaler( + model_path=model_path, scale=scale, half_precision=half_precision + ) + return upscaler.upscale(frames) diff --git a/python/sglang/multimodal_gen/test/server/perf_baselines.json b/python/sglang/multimodal_gen/test/server/perf_baselines.json index 335d121ba..c2cf0857d 100644 --- a/python/sglang/multimodal_gen/test/server/perf_baselines.json +++ b/python/sglang/multimodal_gen/test/server/perf_baselines.json @@ -2136,6 +2136,202 @@ "expected_e2e_ms": 716.81, "expected_avg_denoise_ms": 67.32, "expected_median_denoise_ms": 74.1 + }, + "flux_2_image_t2i_upscaling_4x": { + "stages_ms": { + "InputValidationStage": 0.15, + "TextEncodingStage": 537.4, + "ImageVAEEncodingStage": 0.01, + "LatentPreparationStage": 1.11, + "TimestepPreparationStage": 39.57, + "DenoisingStage": 24738.35, + "DecodingStage": 14.04 + }, + "denoise_step_ms": { + "0": 70.71, + "1": 476.99, + "2": 502.01, + "3": 483.39, + "4": 500.65, + "5": 487.43, + "6": 502.39, + "7": 484.76, + "8": 498.5, + "9": 489.98, + "10": 499.22, + "11": 490.47, + "12": 498.98, + "13": 491.42, + "14": 495.17, + "15": 492.24, + "16": 494.69, + "17": 491.68, + "18": 497.13, + "19": 493.63, + "20": 495.29, + "21": 496.19, + "22": 496.52, + "23": 496.31, + "24": 493.64, + "25": 494.35, + "26": 493.27, + "27": 495.52, + "28": 493.06, + "29": 494.66, + "30": 494.08, + "31": 496.01, + "32": 494.79, + "33": 495.81, + "34": 493.94, + "35": 495.56, + "36": 493.5, + "37": 495.98, + "38": 495.82, + "39": 496.7, + "40": 495.29, + "41": 496.84, + "42": 495.67, + "43": 495.32, + "44": 496.48, + "45": 496.03, + "46": 495.65, + "47": 498.27, + "48": 496.44, + "49": 496.79 + }, + "expected_e2e_ms": 25735.08, + "expected_avg_denoise_ms": 486.1, + "expected_median_denoise_ms": 495.42 + }, + "wan2_1_t2v_1.3b_upscaling_4x": { + "stages_ms": { + "InputValidationStage": 0.08, + "TextEncodingStage": 1164.21, + "LatentPreparationStage": 0.26, + "TimestepPreparationStage": 3.86, + "DenoisingStage": 10234.35, + "DecodingStage": 499.44 + }, + "denoise_step_ms": { + "0": 200.87, + "1": 202.85, + "2": 203.32, + "3": 206.02, + "4": 205.75, + "5": 204.14, + "6": 205.12, + "7": 204.65, + "8": 203.99, + "9": 204.96, + "10": 204.35, + "11": 206.89, + "12": 200.69, + "13": 209.67, + "14": 204.91, + "15": 203.5, + "16": 206.73, + "17": 202.43, + "18": 205.92, + "19": 204.61, + "20": 211.47, + "21": 197.43, + "22": 203.58, + "23": 205.82, + "24": 204.01, + "25": 205.06, + "26": 204.86, + "27": 206.03, + "28": 200.78, + "29": 206.99, + "30": 206.58, + "31": 202.84, + "32": 204.51, + "33": 204.19, + "34": 202.89, + "35": 204.55, + "36": 205.03, + "37": 204.2, + "38": 203.92, + "39": 204.9, + "40": 203.24, + "41": 204.21, + "42": 205.76, + "43": 205.32, + "44": 202.63, + "45": 205.67, + "46": 204.55, + "47": 202.89, + "48": 205.29, + "49": 203.87 + }, + "expected_e2e_ms": 12021.58, + "expected_avg_denoise_ms": 204.49, + "expected_median_denoise_ms": 204.55 + }, + "wan2_1_t2v_1.3b_frame_interp_2x_upscaling_4x": { + "stages_ms": { + "InputValidationStage": 0.03, + "TextEncodingStage": 1089.94, + "LatentPreparationStage": 0.12, + "TimestepPreparationStage": 1.99, + "DenoisingStage": 8617.56, + "DecodingStage": 469.68 + }, + "denoise_step_ms": { + "0": 122.51, + "1": 172.31, + "2": 169.58, + "3": 171.76, + "4": 171.5, + "5": 174.24, + "6": 175.85, + "7": 171.2, + "8": 172.0, + "9": 172.13, + "10": 171.22, + "11": 174.8, + "12": 174.01, + "13": 172.28, + "14": 172.77, + "15": 173.52, + "16": 172.19, + "17": 175.12, + "18": 172.78, + "19": 175.1, + "20": 171.54, + "21": 173.38, + "22": 171.61, + "23": 174.14, + "24": 174.06, + "25": 172.32, + "26": 173.08, + "27": 173.94, + "28": 173.32, + "29": 174.3, + "30": 173.63, + "31": 172.21, + "32": 174.4, + "33": 173.25, + "34": 173.54, + "35": 175.12, + "36": 172.93, + "37": 172.76, + "38": 174.73, + "39": 174.46, + "40": 172.66, + "41": 174.58, + "42": 173.9, + "43": 174.88, + "44": 172.35, + "45": 173.52, + "46": 175.94, + "47": 172.88, + "48": 174.97, + "49": 172.94 + }, + "expected_e2e_ms": 10425.77, + "expected_avg_denoise_ms": 172.28, + "expected_median_denoise_ms": 173.29 } } } diff --git a/python/sglang/multimodal_gen/test/server/test_server_utils.py b/python/sglang/multimodal_gen/test/server/test_server_utils.py index f72de0a88..79879b23c 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_utils.py +++ b/python/sglang/multimodal_gen/test/server/test_server_utils.py @@ -871,9 +871,7 @@ def get_generate_fn( req_background = None # Not specified in current request # Build extra_body for optional features - extra_body = {} - if sampling_params.enable_teacache: - extra_body["enable_teacache"] = True + extra_body = dict(sampling_params.extras) response = client.images.with_raw_response.generate( model=model_path, @@ -898,6 +896,13 @@ def get_generate_fn( # Validate output file expected_width, expected_height = parse_dimensions(output_size) + if ( + sampling_params.extras.get("enable_upscaling") + and expected_width + and expected_height + ): + expected_width *= sampling_params.extras.get("upscaling_scale", 4) + expected_height *= sampling_params.extras.get("upscaling_scale", 4) validate_image_file( tmp_path, expected_filename, @@ -947,8 +952,7 @@ def get_generate_fn( # Build extra_body for optional features extra_body = {"num_frames": sampling_params.num_frames} - if sampling_params.enable_teacache: - extra_body["enable_teacache"] = True + extra_body.update(sampling_params.extras) images = [open(image_path, "rb") for image_path in image_paths] try: @@ -1076,22 +1080,18 @@ def get_generate_fn( pytest.skip(f"{case_id}: no text prompt configured") # Build extra_body for optional features - extra_body = {} - if sampling_params.enable_teacache: - extra_body["enable_teacache"] = True + extra_body = dict(sampling_params.extras) if sampling_params.num_frames: extra_body["num_frames"] = sampling_params.num_frames - if sampling_params.enable_frame_interpolation: - extra_body["enable_frame_interpolation"] = True - extra_body["frame_interpolation_exp"] = ( - sampling_params.frame_interpolation_exp - ) # Compute expected output frame count for validation expected_frame_count = None - if sampling_params.enable_frame_interpolation and sampling_params.num_frames: + if ( + sampling_params.extras.get("enable_frame_interpolation") + and sampling_params.num_frames + ): n = sampling_params.num_frames - exp = sampling_params.frame_interpolation_exp + exp = sampling_params.extras.get("frame_interpolation_exp", 1) expected_frame_count = (n - 1) * (2**exp) + 1 return _create_and_download_video( @@ -1118,9 +1118,7 @@ def get_generate_fn( pytest.skip(f"{case_id}: file missing: {image_path}") # Build extra_body for optional features - extra_body = {} - if sampling_params.enable_teacache: - extra_body["enable_teacache"] = True + extra_body = dict(sampling_params.extras) with image_path.open("rb") as fh: return _create_and_download_video( @@ -1140,8 +1138,7 @@ def get_generate_fn( # Build extra_body for optional features extra_body = {"reference_url": sampling_params.image_path} - if sampling_params.enable_teacache: - extra_body["enable_teacache"] = True + extra_body.update(sampling_params.extras) return _create_and_download_video( client, @@ -1170,9 +1167,7 @@ def get_generate_fn( pytest.skip(f"{case_id}: file missing: {image_path}") # Build extra_body for optional features - extra_body = {} - if sampling_params.enable_teacache: - extra_body["enable_teacache"] = True + extra_body = dict(sampling_params.extras) with image_path.open("rb") as fh: return _create_and_download_video( @@ -1186,6 +1181,7 @@ def get_generate_fn( extra_body={ "fps": sampling_params.fps, "num_frames": sampling_params.num_frames, + **extra_body, }, ) diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index 130182f8f..a840408ce 100644 --- a/python/sglang/multimodal_gen/test/server/testcase_configs.py +++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py @@ -233,12 +233,9 @@ class DiffusionSamplingParams: num_outputs_per_prompt: int = 1 - # TeaCache acceleration - enable_teacache: bool = False - - # Frame interpolation - enable_frame_interpolation: bool = False - frame_interpolation_exp: int = 1 # 1 = 2×, 2 = 4× + # Additional request-level parameters (e.g. enable_teacache, enable_upscaling, …) + # merged directly into the OpenAI extra_body dict. + extras: dict = field(default_factory=dict) @dataclass(frozen=True) @@ -494,6 +491,19 @@ ONE_GPU_CASES_A: list[DiffusionTestCase] = [ ), MULTI_FRAME_I2I_sampling_params, ), + # Upscaling (Real-ESRGAN 4×) for T2I + DiffusionTestCase( + "flux_2_image_t2i_upscaling_4x", + DiffusionServerArgs( + model_path="black-forest-labs/FLUX.2-dev", + modality="image", + ), + DiffusionSamplingParams( + prompt="Doraemon is eating dorayaki", + output_size="1024x1024", + extras={"enable_upscaling": True, "upscaling_scale": 4}, + ), + ), ] HUNYUAN3D_SHAPE_sampling_params = DiffusionSamplingParams( @@ -536,10 +546,10 @@ ONE_GPU_CASES_B: list[DiffusionTestCase] = [ ), DiffusionSamplingParams( prompt=T2V_PROMPT, - enable_teacache=True, + extras={"enable_teacache": True}, ), ), - # Frame interpolation correctness (2× / exp=1) + # Frame interpolation (2× / exp=1) # Uses the same 1.3B model already in the suite; DiffusionTestCase( "wan2_1_t2v_1.3b_frame_interp_2x", @@ -550,8 +560,40 @@ ONE_GPU_CASES_B: list[DiffusionTestCase] = [ ), DiffusionSamplingParams( prompt=T2V_PROMPT, - enable_frame_interpolation=True, - frame_interpolation_exp=1, + extras={"enable_frame_interpolation": True, "frame_interpolation_exp": 1}, + ), + ), + # Upscaling (Real-ESRGAN 4×) + # Uses the same 1.3B model already in the suite; + DiffusionTestCase( + "wan2_1_t2v_1.3b_upscaling_4x", + DiffusionServerArgs( + model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers", + modality="video", + custom_validator="video", + ), + DiffusionSamplingParams( + prompt=T2V_PROMPT, + extras={"enable_upscaling": True, "upscaling_scale": 4}, + ), + ), + # Combined: Frame interpolation (2×) + Upscaling (4×) + # Verifies that both post-processing steps compose correctly. + DiffusionTestCase( + "wan2_1_t2v_1.3b_frame_interp_2x_upscaling_4x", + DiffusionServerArgs( + model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers", + modality="video", + custom_validator="video", + ), + DiffusionSamplingParams( + prompt=T2V_PROMPT, + extras={ + "enable_frame_interpolation": True, + "frame_interpolation_exp": 1, + "enable_upscaling": True, + "upscaling_scale": 4, + }, ), ), # LoRA test case for single transformer + merge/unmerge API test