[diffusion] feat: implement upscaling (#19723)
This commit is contained in:
@@ -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
|
- `--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
|
SGLang diffusion supports optional post-processing steps — frame interpolation
|
||||||
between each pair of consecutive generated frames, producing smoother
|
(RIFE) for smoother video and upscaling (Real-ESRGAN) for higher resolution.
|
||||||
motion without re-running the diffusion model. The `--frame-interpolation-exp`
|
See the dedicated **[Post-Processing](post_processing.md)** page for full
|
||||||
flag controls how many rounds of interpolation to apply: each round inserts one
|
details, supported models, and examples.
|
||||||
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
|
|
||||||
```
|
|
||||||
|
|
||||||
**Output Options**
|
**Output Options**
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
```
|
||||||
@@ -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`
|
- **[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
|
- **[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
|
### Performance Optimization
|
||||||
|
|
||||||
|
|||||||
@@ -113,6 +113,13 @@ class SamplingParams:
|
|||||||
None # local dir or HF repo ID with flownet.pkl (default: elfgum/RIFE-4.22.lite)
|
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
|
# Batch info
|
||||||
num_outputs_per_prompt: int = 1
|
num_outputs_per_prompt: int = 1
|
||||||
seed: int = 42
|
seed: int = 42
|
||||||
@@ -875,7 +882,30 @@ class SamplingParams:
|
|||||||
type=str,
|
type=str,
|
||||||
default=SamplingParams.frame_interpolation_model_path,
|
default=SamplingParams.frame_interpolation_model_path,
|
||||||
help="Local directory or HuggingFace repo ID containing RIFE flownet.pkl weights "
|
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
|
return parser
|
||||||
|
|
||||||
|
|||||||
@@ -277,6 +277,9 @@ class DiffGenerator:
|
|||||||
frame_interpolation_exp=req.frame_interpolation_exp,
|
frame_interpolation_exp=req.frame_interpolation_exp,
|
||||||
frame_interpolation_scale=req.frame_interpolation_scale,
|
frame_interpolation_scale=req.frame_interpolation_scale,
|
||||||
frame_interpolation_model_path=req.frame_interpolation_model_path,
|
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)):
|
for idx in range(len(samples_out)):
|
||||||
|
|||||||
@@ -203,6 +203,9 @@ async def forward_to_scheduler(
|
|||||||
frame_interpolation_exp=sp.frame_interpolation_exp,
|
frame_interpolation_exp=sp.frame_interpolation_exp,
|
||||||
frame_interpolation_scale=sp.frame_interpolation_scale,
|
frame_interpolation_scale=sp.frame_interpolation_scale,
|
||||||
frame_interpolation_model_path=sp.frame_interpolation_model_path,
|
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"):
|
if hasattr(response, "model_dump"):
|
||||||
|
|||||||
@@ -133,6 +133,9 @@ async def generations(
|
|||||||
enable_teacache=request.enable_teacache,
|
enable_teacache=request.enable_teacache,
|
||||||
output_compression=request.output_compression,
|
output_compression=request.output_compression,
|
||||||
output_quality=request.output_quality,
|
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(
|
batch = prepare_request(
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
@@ -207,6 +210,9 @@ async def edits(
|
|||||||
output_quality: Optional[str] = Form("default"),
|
output_quality: Optional[str] = Form("default"),
|
||||||
output_compression: Optional[int] = Form(None),
|
output_compression: Optional[int] = Form(None),
|
||||||
enable_teacache: Optional[bool] = Form(False),
|
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),
|
num_frames: int = Form(1),
|
||||||
):
|
):
|
||||||
request_id = generate_request_id()
|
request_id = generate_request_id()
|
||||||
@@ -262,6 +268,9 @@ async def edits(
|
|||||||
num_frames=num_frames,
|
num_frames=num_frames,
|
||||||
output_compression=output_compression,
|
output_compression=output_compression,
|
||||||
output_quality=output_quality,
|
output_quality=output_quality,
|
||||||
|
enable_upscaling=enable_upscaling,
|
||||||
|
upscaling_model_path=upscaling_model_path,
|
||||||
|
upscaling_scale=upscaling_scale,
|
||||||
)
|
)
|
||||||
batch = prepare_request(
|
batch = prepare_request(
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
|
|||||||
@@ -48,6 +48,10 @@ class ImageGenerationsRequest(BaseModel):
|
|||||||
output_quality: Optional[str] = "default"
|
output_quality: Optional[str] = "default"
|
||||||
output_compression: Optional[int] = None
|
output_compression: Optional[int] = None
|
||||||
enable_teacache: Optional[bool] = False
|
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
|
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_exp: Optional[int] = 1 # 1=2×, 2=4×
|
||||||
frame_interpolation_scale: Optional[float] = 1.0
|
frame_interpolation_scale: Optional[float] = 1.0
|
||||||
frame_interpolation_model_path: Optional[str] = None
|
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_quality: Optional[str] = "default"
|
||||||
output_compression: Optional[int] = None
|
output_compression: Optional[int] = None
|
||||||
output_path: Optional[str] = None
|
output_path: Optional[str] = None
|
||||||
|
|||||||
@@ -286,6 +286,9 @@ async def process_generation_batch(
|
|||||||
frame_interpolation_exp=batch.frame_interpolation_exp,
|
frame_interpolation_exp=batch.frame_interpolation_exp,
|
||||||
frame_interpolation_scale=batch.frame_interpolation_scale,
|
frame_interpolation_scale=batch.frame_interpolation_scale,
|
||||||
frame_interpolation_model_path=batch.frame_interpolation_model_path,
|
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
|
total_time = time.perf_counter() - total_start_time
|
||||||
|
|||||||
@@ -75,6 +75,9 @@ def _build_video_sampling_params(request_id: str, request: VideoGenerationsReque
|
|||||||
frame_interpolation_exp=request.frame_interpolation_exp,
|
frame_interpolation_exp=request.frame_interpolation_exp,
|
||||||
frame_interpolation_scale=request.frame_interpolation_scale,
|
frame_interpolation_scale=request.frame_interpolation_scale,
|
||||||
frame_interpolation_model_path=request.frame_interpolation_model_path,
|
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_path=request.output_path,
|
||||||
output_compression=request.output_compression,
|
output_compression=request.output_compression,
|
||||||
output_quality=request.output_quality,
|
output_quality=request.output_quality,
|
||||||
@@ -181,6 +184,9 @@ async def create_video(
|
|||||||
frame_interpolation_exp: Optional[int] = Form(1),
|
frame_interpolation_exp: Optional[int] = Form(1),
|
||||||
frame_interpolation_scale: Optional[float] = Form(1.0),
|
frame_interpolation_scale: Optional[float] = Form(1.0),
|
||||||
frame_interpolation_model_path: Optional[str] = Form(None),
|
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_quality: Optional[str] = Form("default"),
|
||||||
output_compression: Optional[int] = Form(None),
|
output_compression: Optional[int] = Form(None),
|
||||||
extra_body: Optional[str] = Form(None),
|
extra_body: Optional[str] = Form(None),
|
||||||
@@ -256,6 +262,9 @@ async def create_video(
|
|||||||
frame_interpolation_exp=frame_interpolation_exp,
|
frame_interpolation_exp=frame_interpolation_exp,
|
||||||
frame_interpolation_scale=frame_interpolation_scale,
|
frame_interpolation_scale=frame_interpolation_scale,
|
||||||
frame_interpolation_model_path=frame_interpolation_model_path,
|
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_compression=output_compression,
|
||||||
output_quality=output_quality,
|
output_quality=output_quality,
|
||||||
**(
|
**(
|
||||||
|
|||||||
@@ -348,6 +348,9 @@ def save_outputs(
|
|||||||
frame_interpolation_exp: int = 1,
|
frame_interpolation_exp: int = 1,
|
||||||
frame_interpolation_scale: float = 1.0,
|
frame_interpolation_scale: float = 1.0,
|
||||||
frame_interpolation_model_path: Optional[str] = None,
|
frame_interpolation_model_path: Optional[str] = None,
|
||||||
|
enable_upscaling: bool = False,
|
||||||
|
upscaling_model_path: Optional[str] = None,
|
||||||
|
upscaling_scale: int = 4,
|
||||||
) -> list[str]:
|
) -> list[str]:
|
||||||
"""Save outputs to files and return the list of file paths."""
|
"""Save outputs to files and return the list of file paths."""
|
||||||
output_paths: list[str] = []
|
output_paths: list[str] = []
|
||||||
@@ -369,6 +372,9 @@ def save_outputs(
|
|||||||
frame_interpolation_exp=frame_interpolation_exp,
|
frame_interpolation_exp=frame_interpolation_exp,
|
||||||
frame_interpolation_scale=frame_interpolation_scale,
|
frame_interpolation_scale=frame_interpolation_scale,
|
||||||
frame_interpolation_model_path=frame_interpolation_model_path,
|
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:
|
if samples_out is not None:
|
||||||
@@ -401,6 +407,9 @@ def post_process_sample(
|
|||||||
frame_interpolation_exp: int = 1,
|
frame_interpolation_exp: int = 1,
|
||||||
frame_interpolation_scale: float = 1.0,
|
frame_interpolation_scale: float = 1.0,
|
||||||
frame_interpolation_model_path: Optional[str] = None,
|
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.
|
Process sample output, optionally interpolate video frames, and save.
|
||||||
@@ -456,7 +465,17 @@ def post_process_sample(
|
|||||||
)
|
)
|
||||||
fps = fps * multiplier
|
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_output:
|
||||||
if save_file_path:
|
if save_file_path:
|
||||||
os.makedirs(os.path.dirname(save_file_path), exist_ok=True)
|
os.makedirs(os.path.dirname(save_file_path), exist_ok=True)
|
||||||
|
|||||||
@@ -269,6 +269,9 @@ class GPUWorker:
|
|||||||
frame_interpolation_exp=req.frame_interpolation_exp,
|
frame_interpolation_exp=req.frame_interpolation_exp,
|
||||||
frame_interpolation_scale=req.frame_interpolation_scale,
|
frame_interpolation_scale=req.frame_interpolation_scale,
|
||||||
frame_interpolation_model_path=req.frame_interpolation_model_path,
|
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_file_paths = output_paths
|
||||||
output_batch.output = None
|
output_batch.output = None
|
||||||
|
|||||||
@@ -1,9 +1,18 @@
|
|||||||
# SPDX-License-Identifier: Apache-2.0
|
# 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 (
|
from sglang.multimodal_gen.runtime.postprocess.rife_interpolator import (
|
||||||
FrameInterpolator,
|
FrameInterpolator,
|
||||||
interpolate_video_frames,
|
interpolate_video_frames,
|
||||||
)
|
)
|
||||||
|
|
||||||
__all__ = ["FrameInterpolator", "interpolate_video_frames"]
|
__all__ = [
|
||||||
|
"FrameInterpolator",
|
||||||
|
"interpolate_video_frames",
|
||||||
|
"ImageUpscaler",
|
||||||
|
"upscale_frames",
|
||||||
|
]
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -2136,6 +2136,202 @@
|
|||||||
"expected_e2e_ms": 716.81,
|
"expected_e2e_ms": 716.81,
|
||||||
"expected_avg_denoise_ms": 67.32,
|
"expected_avg_denoise_ms": 67.32,
|
||||||
"expected_median_denoise_ms": 74.1
|
"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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -871,9 +871,7 @@ def get_generate_fn(
|
|||||||
req_background = None # Not specified in current request
|
req_background = None # Not specified in current request
|
||||||
|
|
||||||
# Build extra_body for optional features
|
# Build extra_body for optional features
|
||||||
extra_body = {}
|
extra_body = dict(sampling_params.extras)
|
||||||
if sampling_params.enable_teacache:
|
|
||||||
extra_body["enable_teacache"] = True
|
|
||||||
|
|
||||||
response = client.images.with_raw_response.generate(
|
response = client.images.with_raw_response.generate(
|
||||||
model=model_path,
|
model=model_path,
|
||||||
@@ -898,6 +896,13 @@ def get_generate_fn(
|
|||||||
|
|
||||||
# Validate output file
|
# Validate output file
|
||||||
expected_width, expected_height = parse_dimensions(output_size)
|
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(
|
validate_image_file(
|
||||||
tmp_path,
|
tmp_path,
|
||||||
expected_filename,
|
expected_filename,
|
||||||
@@ -947,8 +952,7 @@ def get_generate_fn(
|
|||||||
|
|
||||||
# Build extra_body for optional features
|
# Build extra_body for optional features
|
||||||
extra_body = {"num_frames": sampling_params.num_frames}
|
extra_body = {"num_frames": sampling_params.num_frames}
|
||||||
if sampling_params.enable_teacache:
|
extra_body.update(sampling_params.extras)
|
||||||
extra_body["enable_teacache"] = True
|
|
||||||
|
|
||||||
images = [open(image_path, "rb") for image_path in image_paths]
|
images = [open(image_path, "rb") for image_path in image_paths]
|
||||||
try:
|
try:
|
||||||
@@ -1076,22 +1080,18 @@ def get_generate_fn(
|
|||||||
pytest.skip(f"{case_id}: no text prompt configured")
|
pytest.skip(f"{case_id}: no text prompt configured")
|
||||||
|
|
||||||
# Build extra_body for optional features
|
# Build extra_body for optional features
|
||||||
extra_body = {}
|
extra_body = dict(sampling_params.extras)
|
||||||
if sampling_params.enable_teacache:
|
|
||||||
extra_body["enable_teacache"] = True
|
|
||||||
if sampling_params.num_frames:
|
if sampling_params.num_frames:
|
||||||
extra_body["num_frames"] = 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
|
# Compute expected output frame count for validation
|
||||||
expected_frame_count = None
|
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
|
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
|
expected_frame_count = (n - 1) * (2**exp) + 1
|
||||||
|
|
||||||
return _create_and_download_video(
|
return _create_and_download_video(
|
||||||
@@ -1118,9 +1118,7 @@ def get_generate_fn(
|
|||||||
pytest.skip(f"{case_id}: file missing: {image_path}")
|
pytest.skip(f"{case_id}: file missing: {image_path}")
|
||||||
|
|
||||||
# Build extra_body for optional features
|
# Build extra_body for optional features
|
||||||
extra_body = {}
|
extra_body = dict(sampling_params.extras)
|
||||||
if sampling_params.enable_teacache:
|
|
||||||
extra_body["enable_teacache"] = True
|
|
||||||
|
|
||||||
with image_path.open("rb") as fh:
|
with image_path.open("rb") as fh:
|
||||||
return _create_and_download_video(
|
return _create_and_download_video(
|
||||||
@@ -1140,8 +1138,7 @@ def get_generate_fn(
|
|||||||
|
|
||||||
# Build extra_body for optional features
|
# Build extra_body for optional features
|
||||||
extra_body = {"reference_url": sampling_params.image_path}
|
extra_body = {"reference_url": sampling_params.image_path}
|
||||||
if sampling_params.enable_teacache:
|
extra_body.update(sampling_params.extras)
|
||||||
extra_body["enable_teacache"] = True
|
|
||||||
|
|
||||||
return _create_and_download_video(
|
return _create_and_download_video(
|
||||||
client,
|
client,
|
||||||
@@ -1170,9 +1167,7 @@ def get_generate_fn(
|
|||||||
pytest.skip(f"{case_id}: file missing: {image_path}")
|
pytest.skip(f"{case_id}: file missing: {image_path}")
|
||||||
|
|
||||||
# Build extra_body for optional features
|
# Build extra_body for optional features
|
||||||
extra_body = {}
|
extra_body = dict(sampling_params.extras)
|
||||||
if sampling_params.enable_teacache:
|
|
||||||
extra_body["enable_teacache"] = True
|
|
||||||
|
|
||||||
with image_path.open("rb") as fh:
|
with image_path.open("rb") as fh:
|
||||||
return _create_and_download_video(
|
return _create_and_download_video(
|
||||||
@@ -1186,6 +1181,7 @@ def get_generate_fn(
|
|||||||
extra_body={
|
extra_body={
|
||||||
"fps": sampling_params.fps,
|
"fps": sampling_params.fps,
|
||||||
"num_frames": sampling_params.num_frames,
|
"num_frames": sampling_params.num_frames,
|
||||||
|
**extra_body,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -233,12 +233,9 @@ class DiffusionSamplingParams:
|
|||||||
|
|
||||||
num_outputs_per_prompt: int = 1
|
num_outputs_per_prompt: int = 1
|
||||||
|
|
||||||
# TeaCache acceleration
|
# Additional request-level parameters (e.g. enable_teacache, enable_upscaling, …)
|
||||||
enable_teacache: bool = False
|
# merged directly into the OpenAI extra_body dict.
|
||||||
|
extras: dict = field(default_factory=dict)
|
||||||
# Frame interpolation
|
|
||||||
enable_frame_interpolation: bool = False
|
|
||||||
frame_interpolation_exp: int = 1 # 1 = 2×, 2 = 4×
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
@@ -494,6 +491,19 @@ ONE_GPU_CASES_A: list[DiffusionTestCase] = [
|
|||||||
),
|
),
|
||||||
MULTI_FRAME_I2I_sampling_params,
|
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(
|
HUNYUAN3D_SHAPE_sampling_params = DiffusionSamplingParams(
|
||||||
@@ -536,10 +546,10 @@ ONE_GPU_CASES_B: list[DiffusionTestCase] = [
|
|||||||
),
|
),
|
||||||
DiffusionSamplingParams(
|
DiffusionSamplingParams(
|
||||||
prompt=T2V_PROMPT,
|
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;
|
# Uses the same 1.3B model already in the suite;
|
||||||
DiffusionTestCase(
|
DiffusionTestCase(
|
||||||
"wan2_1_t2v_1.3b_frame_interp_2x",
|
"wan2_1_t2v_1.3b_frame_interp_2x",
|
||||||
@@ -550,8 +560,40 @@ ONE_GPU_CASES_B: list[DiffusionTestCase] = [
|
|||||||
),
|
),
|
||||||
DiffusionSamplingParams(
|
DiffusionSamplingParams(
|
||||||
prompt=T2V_PROMPT,
|
prompt=T2V_PROMPT,
|
||||||
enable_frame_interpolation=True,
|
extras={"enable_frame_interpolation": True, "frame_interpolation_exp": 1},
|
||||||
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
|
# LoRA test case for single transformer + merge/unmerge API test
|
||||||
|
|||||||
Reference in New Issue
Block a user