[diffusion] chore: add resolution shortcuts for sampling params (#14129)

This commit is contained in:
Mick
2025-11-29 18:00:21 +08:00
committed by GitHub
parent c6d34a0688
commit c6a52f4411
21 changed files with 82 additions and 74 deletions
@@ -1,5 +1,5 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
from sglang.multimodal_gen.configs.sample.base import SamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
__all__ = ["SamplingParams"] __all__ = ["SamplingParams"]
@@ -3,7 +3,7 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass from dataclasses import dataclass
from sglang.multimodal_gen.configs.sample.base import SamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
@dataclass @dataclass
@@ -3,7 +3,7 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.sample.base import SamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
from sglang.multimodal_gen.configs.sample.teacache import TeaCacheParams from sglang.multimodal_gen.configs.sample.teacache import TeaCacheParams
@@ -3,7 +3,7 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass from dataclasses import dataclass
from sglang.multimodal_gen.configs.sample.base import SamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
@dataclass @dataclass
@@ -195,7 +195,7 @@ class SamplingParams:
if self.prompt_path and not self.prompt_path.endswith(".txt"): if self.prompt_path and not self.prompt_path.endswith(".txt"):
raise ValueError("prompt_path must be a txt file") raise ValueError("prompt_path must be a txt file")
def adjust( def _adjust(
self, self,
server_args: ServerArgs, server_args: ServerArgs,
): ):
@@ -319,7 +319,7 @@ class SamplingParams:
sampling_params._merge_with_user_params(user_sampling_params) sampling_params._merge_with_user_params(user_sampling_params)
sampling_params.width_not_provided = user_sampling_params.width is None sampling_params.width_not_provided = user_sampling_params.width is None
sampling_params.height_not_provided = user_sampling_params.height is None sampling_params.height_not_provided = user_sampling_params.height is None
sampling_params.adjust(server_args) sampling_params._adjust(server_args)
return sampling_params return sampling_params
@@ -421,6 +421,32 @@ class SamplingParams:
default=SamplingParams.width, default=SamplingParams.width,
help="Width of generated output", help="Width of generated output",
) )
# resolution shortcuts
parser.add_argument(
"--4k",
action="store_true",
dest="resolution_4k",
help="Set resolution to 4K (3840x2160)",
)
parser.add_argument(
"--2k",
action="store_true",
dest="resolution_2k",
help="Set resolution to 2K (2560x1440)",
)
parser.add_argument(
"--1080p",
action="store_true",
dest="resolution_1080p",
help="Set resolution to 1080p (1920x1080)",
)
parser.add_argument(
"--720p",
action="store_true",
dest="resolution_720p",
help="Set resolution to 720p (1280x720)",
)
parser.add_argument( parser.add_argument(
"--fps", "--fps",
type=int, type=int,
@@ -496,11 +522,25 @@ class SamplingParams:
return parser return parser
@classmethod @classmethod
def from_cli_args(cls, args: argparse.Namespace): def get_cli_args(cls, args: argparse.Namespace):
# handle resolution shortcuts
if hasattr(args, "resolution_4k") and args.resolution_4k:
args.width = 3840
args.height = 2160
elif hasattr(args, "resolution_2k") and args.resolution_2k:
args.width = 2560
args.height = 1440
elif hasattr(args, "resolution_1080p") and args.resolution_1080p:
args.width = 1920
args.height = 1080
elif hasattr(args, "resolution_720p") and args.resolution_720p:
args.width = 1280
args.height = 720
attrs = [attr.name for attr in dataclasses.fields(cls)] attrs = [attr.name for attr in dataclasses.fields(cls)]
args.height_not_provided = False args.height_not_provided = False
args.width_not_provided = False args.width_not_provided = False
return cls(**{attr: getattr(args, attr) for attr in attrs}) return {attr: getattr(args, attr) for attr in attrs}
def output_file_path(self): def output_file_path(self):
return os.path.join(self.output_path, self.output_file_name) return os.path.join(self.output_path, self.output_file_name)
@@ -575,8 +615,8 @@ class SamplingParams:
# Log sampling parameters # Log sampling parameters
debug_str = f"""Sampling params: debug_str = f"""Sampling params:
height: {target_height}
width: {target_width} width: {target_width}
height: {target_height}
num_frames: {self.num_frames} num_frames: {self.num_frames}
prompt: {self.prompt} prompt: {self.prompt}
neg_prompt: {self.negative_prompt} neg_prompt: {self.negative_prompt}
@@ -3,7 +3,7 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass from dataclasses import dataclass
from sglang.multimodal_gen.configs.sample.base import SamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
@dataclass @dataclass
@@ -3,7 +3,7 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.sample.base import CacheParams from sglang.multimodal_gen.configs.sample.sampling_params import CacheParams
@dataclass @dataclass
@@ -3,7 +3,7 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.sample.base import SamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
from sglang.multimodal_gen.configs.sample.teacache import WanTeaCacheParams from sglang.multimodal_gen.configs.sample.teacache import WanTeaCacheParams
@@ -3,19 +3,13 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.sample.base import SamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
from sglang.multimodal_gen.configs.sample.teacache import TeaCacheParams from sglang.multimodal_gen.configs.sample.teacache import TeaCacheParams
@dataclass @dataclass
class ZImageSamplingParams(SamplingParams): class ZImageSamplingParams(SamplingParams):
num_inference_steps: int = 9 num_inference_steps: int = 9
num_frames: int = 1
height: int = 720
width: int = 1280
fps: int = 24
guidance_scale: float = 0.0 guidance_scale: float = 0.0
teacache_params: TeaCacheParams = field( teacache_params: TeaCacheParams = field(
@@ -10,7 +10,7 @@ from typing import cast
import sglang.multimodal_gen.envs as envs import sglang.multimodal_gen.envs as envs
from sglang.multimodal_gen import DiffGenerator from sglang.multimodal_gen import DiffGenerator
from sglang.multimodal_gen.configs.sample.base import ( from sglang.multimodal_gen.configs.sample.sampling_params import (
SamplingParams, SamplingParams,
generate_request_id, generate_request_id,
) )
@@ -58,9 +58,7 @@ def add_multimodal_gen_generate_args(parser: argparse.ArgumentParser):
return parser return parser
def maybe_dump_performance( def maybe_dump_performance(args: argparse.Namespace, server_args, prompt: str, results):
args: argparse.Namespace, server_args, sampling_params, results
):
"""dump performance if necessary""" """dump performance if necessary"""
if not (args.perf_dump_path and results): if not (args.perf_dump_path and results):
return return
@@ -82,7 +80,7 @@ def maybe_dump_performance(
file_path=args.perf_dump_path, file_path=args.perf_dump_path,
timings=timings, timings=timings,
meta={ meta={
"prompt": sampling_params.prompt, "prompt": prompt,
"model": server_args.model_path, "model": server_args.model_path,
}, },
tag="cli_generate", tag="cli_generate",
@@ -100,17 +98,15 @@ def generate_cmd(args: argparse.Namespace):
envs.SGLANG_DIFFUSION_STAGE_LOGGING = True envs.SGLANG_DIFFUSION_STAGE_LOGGING = True
server_args = ServerArgs.from_cli_args(args) server_args = ServerArgs.from_cli_args(args)
sampling_params = SamplingParams.from_cli_args(args) sampling_params_kwargs = SamplingParams.get_cli_args(args)
sampling_params.request_id = generate_request_id()
generator = DiffGenerator.from_pretrained( generator = DiffGenerator.from_pretrained(
model_path=server_args.model_path, server_args=server_args model_path=server_args.model_path, server_args=server_args
) )
results = generator.generate( results = generator.generate(sampling_params_kwargs=sampling_params_kwargs)
prompt=sampling_params.prompt, sampling_params=sampling_params
)
maybe_dump_performance(args, server_args, sampling_params, results) prompt = sampling_params_kwargs.get("prompt", None)
maybe_dump_performance(args, server_args, prompt, results)
class GenerateSubcommand(CLISubcommand): class GenerateSubcommand(CLISubcommand):
@@ -11,7 +11,6 @@ diffusion models.
import multiprocessing as mp import multiprocessing as mp
import os import os
import time import time
from copy import deepcopy
from typing import Any from typing import Any
import imageio import imageio
@@ -20,7 +19,10 @@ import torch
import torchvision import torchvision
from einops import rearrange from einops import rearrange
from sglang.multimodal_gen.configs.sample.base import DataType, SamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import (
DataType,
SamplingParams,
)
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
MergeLoraWeightsReq, MergeLoraWeightsReq,
SetLoraReq, SetLoraReq,
@@ -208,32 +210,18 @@ class DiffGenerator:
def generate( def generate(
self, self,
prompt: str | list[str] | None = None, sampling_params_kwargs: dict | None = None,
sampling_params: SamplingParams | None = None,
**kwargs,
) -> dict[str, Any] | list[np.ndarray] | list[dict[str, Any]] | None: ) -> dict[str, Any] | list[np.ndarray] | list[dict[str, Any]] | None:
""" """
Generate a image/video based on the given prompt. Generate a image/video based on the given prompt.
Args: Args:
prompt: The prompt to use for generation (optional if prompt_txt is provided)
output_file_name: Name of the file to save. Default is the first 100 characters of the prompt.
save_output: Whether to save the output to disk
return_frames: Whether to return the raw frames
num_inference_steps: Number of denoising steps (overrides server_args)
guidance_scale: Classifier-free guidance scale (overrides server_args)
num_frames: Number of frames to generate (overrides server_args)
height: Height of generated file (overrides server_args)
width: Width of generated file (overrides server_args)
fps: Frames per second for saved file (overrides server_args)
seed: Random seed for generation (overrides server_args)
callback: Callback function called after each step
callback_steps: Number of steps between each callback
Returns: Returns:
Either the output dictionary, list of frames, or list of results for batch processing Either the output dictionary, list of frames, or list of results for batch processing
""" """
# 1. prepare requests # 1. prepare requests
prompt = sampling_params_kwargs.get("prompt", None)
prompts: list[str] = [] prompts: list[str] = []
# Handle batch processing from text file # Handle batch processing from text file
if self.server_args.prompt_file_path is not None: if self.server_args.prompt_file_path is not None:
@@ -258,29 +246,19 @@ class DiffGenerator:
else: else:
raise ValueError("Either prompt or prompt_txt must be provided") raise ValueError("Either prompt or prompt_txt must be provided")
pretrained_sampling_params = SamplingParams.from_pretrained( sampling_params = SamplingParams.from_user_sampling_params_args(
self.server_args.model_path, **kwargs self.server_args.model_path,
server_args=self.server_args,
**sampling_params_kwargs,
) )
pretrained_sampling_params._merge_with_user_params(sampling_params)
# TODO: simplify
data_type = (
DataType.IMAGE
if self.server_args.pipeline_config.task_type.is_image_gen()
or pretrained_sampling_params.num_frames == 1
else DataType.VIDEO
)
pretrained_sampling_params.data_type = data_type
pretrained_sampling_params._set_output_file_name()
pretrained_sampling_params.adjust(self.server_args)
requests: list[Req] = [] requests: list[Req] = []
for output_idx, p in enumerate(prompts): for output_idx, p in enumerate(prompts):
current_sampling_params = deepcopy(pretrained_sampling_params) sampling_params.prompt = p
current_sampling_params.prompt = p
requests.append( requests.append(
prepare_request( prepare_request(
server_args=self.server_args, server_args=self.server_args,
sampling_params=current_sampling_params, sampling_params=sampling_params,
) )
) )
@@ -8,7 +8,7 @@ from typing import List, Optional
from fastapi import APIRouter, File, Form, HTTPException, Path, Query, UploadFile from fastapi import APIRouter, File, Form, HTTPException, Path, Query, UploadFile
from fastapi.responses import FileResponse from fastapi.responses import FileResponse
from sglang.multimodal_gen.configs.sample.base import ( from sglang.multimodal_gen.configs.sample.sampling_params import (
SamplingParams, SamplingParams,
generate_request_id, generate_request_id,
) )
@@ -11,7 +11,7 @@ import torchvision
from einops import rearrange from einops import rearrange
from fastapi import UploadFile from fastapi import UploadFile
from sglang.multimodal_gen.configs.sample.base import DataType from sglang.multimodal_gen.configs.sample.sampling_params import DataType
from sglang.multimodal_gen.runtime.utils.logging_utils import ( from sglang.multimodal_gen.runtime.utils.logging_utils import (
init_logger, init_logger,
log_batch_completion, log_batch_completion,
@@ -18,7 +18,7 @@ from fastapi import (
) )
from fastapi.responses import FileResponse from fastapi.responses import FileResponse
from sglang.multimodal_gen.configs.sample.base import ( from sglang.multimodal_gen.configs.sample.sampling_params import (
SamplingParams, SamplingParams,
generate_request_id, generate_request_id,
) )
@@ -9,7 +9,7 @@ diffusion models.
""" """
from sglang.multimodal_gen.configs.sample.base import SamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -19,7 +19,7 @@ from typing import TYPE_CHECKING, Any, Optional
import PIL.Image import PIL.Image
import torch import torch
from sglang.multimodal_gen.configs.sample.base import DataType from sglang.multimodal_gen.configs.sample.sampling_params import DataType
from sglang.multimodal_gen.configs.sample.teacache import ( from sglang.multimodal_gen.configs.sample.teacache import (
TeaCacheParams, TeaCacheParams,
WanTeaCacheParams, WanTeaCacheParams,
@@ -3,7 +3,7 @@
import unittest import unittest
from pathlib import Path from pathlib import Path
from sglang.multimodal_gen.configs.sample.base import DataType from sglang.multimodal_gen.configs.sample.sampling_params import DataType
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.test.test_utils import TestGenerateBase from sglang.multimodal_gen.test.test_utils import TestGenerateBase
@@ -2,7 +2,7 @@
import unittest import unittest
from sglang.multimodal_gen.configs.sample.base import DataType from sglang.multimodal_gen.configs.sample.sampling_params import DataType
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.test.test_utils import TestGenerateBase from sglang.multimodal_gen.test.test_utils import TestGenerateBase
@@ -1,6 +1,6 @@
import unittest import unittest
from sglang.multimodal_gen.configs.sample.base import DataType from sglang.multimodal_gen.configs.sample.sampling_params import DataType
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.test.test_utils import TestGenerateBase from sglang.multimodal_gen.test.test_utils import TestGenerateBase
@@ -382,7 +382,7 @@
"7": 102.28, "7": 102.28,
"8": 105.54 "8": 105.54
}, },
"expected_e2e_ms": 1248.41, "expected_e2e_ms": 1383.47,
"expected_avg_denoise_ms": 94.15, "expected_avg_denoise_ms": 94.15,
"expected_median_denoise_ms": 102.03 "expected_median_denoise_ms": 102.03
}, },
@@ -14,7 +14,7 @@ from typing import Optional
from PIL import Image from PIL import Image
from sglang.multimodal_gen.configs.sample.base import DataType from sglang.multimodal_gen.configs.sample.sampling_params import DataType
from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.runtime.utils.perf_logger import ( from sglang.multimodal_gen.runtime.utils.perf_logger import (