[diffusion] refactor: refactor sampling params (#13706)
This commit is contained in:
@@ -5,12 +5,12 @@ import argparse
|
||||
import dataclasses
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import os.path
|
||||
import re
|
||||
import time
|
||||
import unicodedata
|
||||
import uuid
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum, auto
|
||||
from typing import Any
|
||||
@@ -137,7 +137,7 @@ class SamplingParams:
|
||||
return_trajectory_latents: bool = False # returns all latents for each timestep
|
||||
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
|
||||
|
||||
def set_output_file_ext(self):
|
||||
def _set_output_file_ext(self):
|
||||
# add extension if needed
|
||||
if not any(
|
||||
self.output_file_name.endswith(ext)
|
||||
@@ -147,7 +147,7 @@ class SamplingParams:
|
||||
f"{self.output_file_name}.{self.data_type.get_default_extension()}"
|
||||
)
|
||||
|
||||
def set_output_file_name(self):
|
||||
def _set_output_file_name(self):
|
||||
# settle output_file_name
|
||||
if (
|
||||
self.output_file_name is None
|
||||
@@ -178,7 +178,7 @@ class SamplingParams:
|
||||
self.output_file_name = _sanitize_filename(self.output_file_name)
|
||||
|
||||
# Ensure a proper extension is present
|
||||
self.set_output_file_ext()
|
||||
self._set_output_file_ext()
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
assert self.num_frames >= 1
|
||||
@@ -195,6 +195,93 @@ class SamplingParams:
|
||||
if self.prompt_path and not self.prompt_path.endswith(".txt"):
|
||||
raise ValueError("prompt_path must be a txt file")
|
||||
|
||||
def adjust(
|
||||
self,
|
||||
server_args: ServerArgs,
|
||||
):
|
||||
"""
|
||||
final adjustment, called after merged with user params
|
||||
"""
|
||||
pipeline_config = server_args.pipeline_config
|
||||
if not isinstance(self.prompt, str):
|
||||
raise TypeError(f"`prompt` must be a string, but got {type(self.prompt)}")
|
||||
|
||||
# Process negative prompt
|
||||
if self.negative_prompt is not None and not self.negative_prompt.isspace():
|
||||
# avoid stripping default negative prompt: ' ' for qwen-image
|
||||
self.negative_prompt = self.negative_prompt.strip()
|
||||
|
||||
# Validate dimensions
|
||||
if self.num_frames <= 0:
|
||||
raise ValueError(
|
||||
f"height, width, and num_frames must be positive integers, got "
|
||||
f"height={self.height}, width={self.width}, "
|
||||
f"num_frames={self.num_frames}"
|
||||
)
|
||||
|
||||
if pipeline_config.task_type.is_image_gen():
|
||||
# settle num_frames
|
||||
logger.debug(f"Setting num_frames to 1 because this is a image-gen model")
|
||||
self.num_frames = 1
|
||||
self.data_type = DataType.IMAGE
|
||||
else:
|
||||
# Adjust number of frames based on number of GPUs for video task
|
||||
use_temporal_scaling_frames = (
|
||||
pipeline_config.vae_config.use_temporal_scaling_frames
|
||||
)
|
||||
num_frames = self.num_frames
|
||||
num_gpus = server_args.num_gpus
|
||||
temporal_scale_factor = (
|
||||
pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
)
|
||||
|
||||
if use_temporal_scaling_frames:
|
||||
orig_latent_num_frames = (num_frames - 1) // temporal_scale_factor + 1
|
||||
else: # stepvideo only
|
||||
orig_latent_num_frames = self.num_frames // 17 * 3
|
||||
|
||||
if orig_latent_num_frames % server_args.num_gpus != 0:
|
||||
# Adjust latent frames to be divisible by number of GPUs
|
||||
if self.num_frames_round_down:
|
||||
# Ensure we have at least 1 batch per GPU
|
||||
new_latent_num_frames = (
|
||||
max(1, (orig_latent_num_frames // num_gpus)) * num_gpus
|
||||
)
|
||||
else:
|
||||
new_latent_num_frames = (
|
||||
math.ceil(orig_latent_num_frames / num_gpus) * num_gpus
|
||||
)
|
||||
|
||||
if use_temporal_scaling_frames:
|
||||
# Convert back to number of frames, ensuring num_frames-1 is a multiple of temporal_scale_factor
|
||||
new_num_frames = (
|
||||
new_latent_num_frames - 1
|
||||
) * temporal_scale_factor + 1
|
||||
else: # stepvideo only
|
||||
# Find the least common multiple of 3 and num_gpus
|
||||
divisor = math.lcm(3, num_gpus)
|
||||
# Round up to the nearest multiple of this LCM
|
||||
new_latent_num_frames = (
|
||||
(new_latent_num_frames + divisor - 1) // divisor
|
||||
) * divisor
|
||||
# Convert back to actual frames using the StepVideo formula
|
||||
new_num_frames = new_latent_num_frames // 3 * 17
|
||||
|
||||
logger.info(
|
||||
"Adjusting number of frames from %s to %s based on number of GPUs (%s)",
|
||||
self.num_frames,
|
||||
new_num_frames,
|
||||
server_args.num_gpus,
|
||||
)
|
||||
self.num_frames = new_num_frames
|
||||
|
||||
self.num_frames = server_args.pipeline_config.adjust_num_frames(
|
||||
self.num_frames
|
||||
)
|
||||
|
||||
self._set_output_file_name()
|
||||
self.log(server_args=server_args)
|
||||
|
||||
def update(self, source_dict: dict[str, Any]) -> None:
|
||||
for key, value in source_dict.items():
|
||||
if hasattr(self, key):
|
||||
@@ -220,9 +307,15 @@ class SamplingParams:
|
||||
sampling_params = cls(**kwargs)
|
||||
return sampling_params
|
||||
|
||||
def from_user_sampling_params(self, user_params):
|
||||
sampling_params = deepcopy(self)
|
||||
sampling_params._merge_with_user_params(user_params)
|
||||
@staticmethod
|
||||
def from_user_sampling_params_args(model_path: str, server_args, *args, **kwargs):
|
||||
sampling_params = SamplingParams.from_pretrained(model_path)
|
||||
|
||||
user_sampling_params = SamplingParams(*args, **kwargs)
|
||||
sampling_params._merge_with_user_params(user_sampling_params)
|
||||
|
||||
sampling_params.adjust(server_args)
|
||||
|
||||
return sampling_params
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -264,7 +264,8 @@ class DiffGenerator:
|
||||
else DataType.VIDEO
|
||||
)
|
||||
pretrained_sampling_params.data_type = data_type
|
||||
pretrained_sampling_params.set_output_file_name()
|
||||
pretrained_sampling_params._set_output_file_name()
|
||||
pretrained_sampling_params.adjust(self.server_args)
|
||||
|
||||
requests: list[Req] = []
|
||||
for output_idx, p in enumerate(prompts):
|
||||
@@ -272,7 +273,6 @@ class DiffGenerator:
|
||||
current_sampling_params.prompt = p
|
||||
requests.append(
|
||||
prepare_request(
|
||||
p,
|
||||
server_args=self.server_args,
|
||||
sampling_params=current_sampling_params,
|
||||
)
|
||||
@@ -310,21 +310,11 @@ class DiffGenerator:
|
||||
continue
|
||||
for output_idx, sample in enumerate(output_batch.output):
|
||||
num_outputs = len(output_batch.output)
|
||||
output_file_name = req.output_file_name
|
||||
if num_outputs > 1 and output_file_name:
|
||||
base, ext = os.path.splitext(output_file_name)
|
||||
output_file_name = f"{base}_{output_idx}{ext}"
|
||||
|
||||
save_path = (
|
||||
os.path.join(req.output_path, output_file_name)
|
||||
if output_file_name
|
||||
else None
|
||||
)
|
||||
frames = self.post_process_sample(
|
||||
sample,
|
||||
fps=req.fps,
|
||||
save_output=req.save_output,
|
||||
save_file_path=save_path,
|
||||
save_file_path=req.output_file_path(num_outputs, output_idx),
|
||||
data_type=req.data_type,
|
||||
)
|
||||
|
||||
|
||||
@@ -56,12 +56,10 @@ def _build_sampling_params_from_request(
|
||||
) -> SamplingParams:
|
||||
width, height = _parse_size(size)
|
||||
ext = _choose_ext(output_format, background)
|
||||
|
||||
server_args = get_global_server_args()
|
||||
sampling_params = SamplingParams.from_pretrained(server_args.model_path)
|
||||
|
||||
# Build user params
|
||||
user_params = SamplingParams(
|
||||
sampling_params = SamplingParams.from_user_sampling_params_args(
|
||||
model_path=server_args.model_path,
|
||||
request_id=request_id,
|
||||
prompt=prompt,
|
||||
image_path=image_path,
|
||||
@@ -70,18 +68,9 @@ def _build_sampling_params_from_request(
|
||||
height=height,
|
||||
num_outputs_per_prompt=max(1, min(int(n or 1), 10)),
|
||||
save_output=True,
|
||||
server_args=server_args,
|
||||
output_file_name=f"{request_id}.{ext}",
|
||||
)
|
||||
|
||||
# Let SamplingParams auto-generate a file name, then force desired extension
|
||||
sampling_params = sampling_params.from_user_sampling_params(user_params)
|
||||
if not sampling_params.output_file_name:
|
||||
sampling_params.output_file_name = request_id
|
||||
if not sampling_params.output_file_name.endswith(f".{ext}"):
|
||||
# strip any existing extension and apply desired one
|
||||
base = sampling_params.output_file_name.rsplit(".", 1)[0]
|
||||
sampling_params.output_file_name = f"{base}.{ext}"
|
||||
|
||||
sampling_params.log(server_args)
|
||||
return sampling_params
|
||||
|
||||
|
||||
@@ -107,7 +96,6 @@ def _build_req_from_sampling(s: SamplingParams) -> Req:
|
||||
async def generations(
|
||||
request: ImageGenerationsRequest,
|
||||
):
|
||||
|
||||
request_id = generate_request_id()
|
||||
sampling = _build_sampling_params_from_request(
|
||||
request_id=request_id,
|
||||
@@ -118,7 +106,6 @@ async def generations(
|
||||
background=request.background,
|
||||
)
|
||||
batch = prepare_request(
|
||||
prompt=request.prompt,
|
||||
server_args=get_global_server_args(),
|
||||
sampling_params=sampling,
|
||||
)
|
||||
@@ -175,7 +162,6 @@ async def edits(
|
||||
background: Optional[str] = Form("auto"),
|
||||
user: Optional[str] = Form(None),
|
||||
):
|
||||
|
||||
request_id = generate_request_id()
|
||||
# Resolve images from either `image` or `image[]` (OpenAI SDK sends `image[]` when list is provided)
|
||||
images = image or image_array
|
||||
|
||||
@@ -42,6 +42,8 @@ logger = init_logger(__name__)
|
||||
router = APIRouter(prefix="/v1/videos", tags=["videos"])
|
||||
|
||||
|
||||
# NOTE(mick): the sampling params needs to be further adjusted
|
||||
# FIXME: duplicated with the one in `image_api.py`
|
||||
def _build_sampling_params_from_request(
|
||||
request_id: str, request: VideoGenerationsRequest
|
||||
) -> SamplingParams:
|
||||
@@ -56,9 +58,8 @@ def _build_sampling_params_from_request(
|
||||
request.num_frames if request.num_frames is not None else derived_num_frames
|
||||
)
|
||||
server_args = get_global_server_args()
|
||||
# TODO: should we cache this sampling_params?
|
||||
sampling_params = SamplingParams.from_pretrained(server_args.model_path)
|
||||
user_params = SamplingParams(
|
||||
sampling_params = SamplingParams.from_user_sampling_params_args(
|
||||
model_path=server_args.model_path,
|
||||
request_id=request_id,
|
||||
prompt=request.prompt,
|
||||
num_frames=num_frames,
|
||||
@@ -67,10 +68,10 @@ def _build_sampling_params_from_request(
|
||||
height=height,
|
||||
image_path=request.input_reference,
|
||||
save_output=True,
|
||||
server_args=server_args,
|
||||
output_file_name=request_id,
|
||||
)
|
||||
sampling_params = sampling_params.from_user_sampling_params(user_params)
|
||||
sampling_params.set_output_file_name()
|
||||
sampling_params.log(server_args)
|
||||
|
||||
return sampling_params
|
||||
|
||||
|
||||
@@ -195,7 +196,6 @@ async def create_video(
|
||||
|
||||
# Build Req for scheduler
|
||||
batch = prepare_request(
|
||||
prompt=req.prompt,
|
||||
server_args=get_global_server_args(),
|
||||
sampling_params=sampling_params,
|
||||
)
|
||||
|
||||
@@ -9,13 +9,12 @@ diffusion models.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import math
|
||||
|
||||
# Suppress verbose logging from imageio, which is triggered when saving images.
|
||||
logging.getLogger("imageio").setLevel(logging.WARNING)
|
||||
logging.getLogger("imageio_ffmpeg").setLevel(logging.WARNING)
|
||||
|
||||
from sglang.multimodal_gen.configs.sample.base import DataType, SamplingParams
|
||||
from sglang.multimodal_gen.configs.sample.base import SamplingParams
|
||||
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.utils.logging_utils import init_logger
|
||||
@@ -24,97 +23,7 @@ from sglang.multimodal_gen.utils import shallow_asdict
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def prepare_sampling_params(
|
||||
prompt: str,
|
||||
server_args: ServerArgs,
|
||||
sampling_params: SamplingParams,
|
||||
):
|
||||
pipeline_config = server_args.pipeline_config
|
||||
# Validate inputs
|
||||
if not isinstance(prompt, str):
|
||||
raise TypeError(f"`prompt` must be a string, but got {type(prompt)}")
|
||||
|
||||
# Process negative prompt
|
||||
if (
|
||||
sampling_params.negative_prompt is not None
|
||||
and not sampling_params.negative_prompt.isspace()
|
||||
):
|
||||
# avoid stripping default negative prompt: ' ' for qwen-image
|
||||
sampling_params.negative_prompt = sampling_params.negative_prompt.strip()
|
||||
|
||||
# Validate dimensions
|
||||
if sampling_params.num_frames <= 0:
|
||||
raise ValueError(
|
||||
f"height, width, and num_frames must be positive integers, got "
|
||||
f"height={sampling_params.height}, width={sampling_params.width}, "
|
||||
f"num_frames={sampling_params.num_frames}"
|
||||
)
|
||||
|
||||
if pipeline_config.task_type.is_image_gen():
|
||||
# settle num_frames
|
||||
logger.debug(f"Setting num_frames to 1 because this is a image-gen model")
|
||||
sampling_params.num_frames = 1
|
||||
sampling_params.data_type = DataType.IMAGE
|
||||
else:
|
||||
# Adjust number of frames based on number of GPUs for video task
|
||||
use_temporal_scaling_frames = (
|
||||
pipeline_config.vae_config.use_temporal_scaling_frames
|
||||
)
|
||||
num_frames = sampling_params.num_frames
|
||||
num_gpus = server_args.num_gpus
|
||||
temporal_scale_factor = (
|
||||
pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
)
|
||||
|
||||
if use_temporal_scaling_frames:
|
||||
orig_latent_num_frames = (num_frames - 1) // temporal_scale_factor + 1
|
||||
else: # stepvideo only
|
||||
orig_latent_num_frames = sampling_params.num_frames // 17 * 3
|
||||
|
||||
if orig_latent_num_frames % server_args.num_gpus != 0:
|
||||
# Adjust latent frames to be divisible by number of GPUs
|
||||
if sampling_params.num_frames_round_down:
|
||||
# Ensure we have at least 1 batch per GPU
|
||||
new_latent_num_frames = (
|
||||
max(1, (orig_latent_num_frames // num_gpus)) * num_gpus
|
||||
)
|
||||
else:
|
||||
new_latent_num_frames = (
|
||||
math.ceil(orig_latent_num_frames / num_gpus) * num_gpus
|
||||
)
|
||||
|
||||
if use_temporal_scaling_frames:
|
||||
# Convert back to number of frames, ensuring num_frames-1 is a multiple of temporal_scale_factor
|
||||
new_num_frames = (new_latent_num_frames - 1) * temporal_scale_factor + 1
|
||||
else: # stepvideo only
|
||||
# Find the least common multiple of 3 and num_gpus
|
||||
divisor = math.lcm(3, num_gpus)
|
||||
# Round up to the nearest multiple of this LCM
|
||||
new_latent_num_frames = (
|
||||
(new_latent_num_frames + divisor - 1) // divisor
|
||||
) * divisor
|
||||
# Convert back to actual frames using the StepVideo formula
|
||||
new_num_frames = new_latent_num_frames // 3 * 17
|
||||
|
||||
logger.info(
|
||||
"Adjusting number of frames from %s to %s based on number of GPUs (%s)",
|
||||
sampling_params.num_frames,
|
||||
new_num_frames,
|
||||
server_args.num_gpus,
|
||||
)
|
||||
sampling_params.num_frames = new_num_frames
|
||||
|
||||
sampling_params.num_frames = server_args.pipeline_config.adjust_num_frames(
|
||||
sampling_params.num_frames
|
||||
)
|
||||
|
||||
sampling_params.set_output_file_ext()
|
||||
sampling_params.log(server_args=server_args)
|
||||
return sampling_params
|
||||
|
||||
|
||||
def prepare_request(
|
||||
prompt: str,
|
||||
server_args: ServerArgs,
|
||||
sampling_params: SamplingParams,
|
||||
) -> Req:
|
||||
@@ -123,20 +32,16 @@ def prepare_request(
|
||||
|
||||
"""
|
||||
# Create a copy of inference args to avoid modifying the original
|
||||
|
||||
sampling_params = prepare_sampling_params(prompt, server_args, sampling_params)
|
||||
|
||||
req = Req(
|
||||
**shallow_asdict(sampling_params),
|
||||
VSA_sparsity=server_args.VSA_sparsity,
|
||||
)
|
||||
# req.set_width_and_height(server_args)
|
||||
req.adjust_size(server_args)
|
||||
|
||||
# if (req.width <= 0
|
||||
# or req.height <= 0):
|
||||
# raise ValueError(
|
||||
# f"Height, width must be positive integers, got "
|
||||
# f"height={req.height}, width={req.width}"
|
||||
# )
|
||||
if req.width <= 0 or req.height <= 0:
|
||||
raise ValueError(
|
||||
f"Height, width must be positive integers, got "
|
||||
f"height={req.height}, width={req.width}"
|
||||
)
|
||||
|
||||
return req
|
||||
|
||||
@@ -11,6 +11,7 @@ in a functional manner, reducing the need for explicit parameter passing.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import pprint
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
@@ -187,6 +188,18 @@ class Req:
|
||||
batch_size *= self.num_outputs_per_prompt
|
||||
return batch_size
|
||||
|
||||
def output_file_path(self, num_outputs, output_idx):
|
||||
output_file_name = self.output_file_name
|
||||
if num_outputs > 1 and output_file_name:
|
||||
base, ext = os.path.splitext(output_file_name)
|
||||
output_file_name = f"{base}_{output_idx}{ext}"
|
||||
|
||||
return (
|
||||
os.path.join(self.output_path, output_file_name)
|
||||
if output_file_name
|
||||
else None
|
||||
)
|
||||
|
||||
def __post_init__(self):
|
||||
"""Initialize dependent fields after dataclass initialization."""
|
||||
# Set do_classifier_free_guidance based on guidance scale and negative prompt
|
||||
@@ -197,7 +210,7 @@ class Req:
|
||||
if self.guidance_scale_2 is None:
|
||||
self.guidance_scale_2 = self.guidance_scale
|
||||
|
||||
def set_width_and_height(self, server_args: ServerArgs):
|
||||
def adjust_size(self, server_args: ServerArgs):
|
||||
if self.height is None or self.width is None:
|
||||
width, height = server_args.pipeline_config.adjust_size(
|
||||
self.width, self.height, self.pil_image
|
||||
|
||||
@@ -273,7 +273,6 @@ Consider updating perf_baselines.json with the snippets below:
|
||||
response_format="b64_json",
|
||||
)
|
||||
rid = response.headers.get("x-request-id", "")
|
||||
print(f"{response=}")
|
||||
|
||||
result = response.parse()
|
||||
validate_image(result.data[0].b64_json)
|
||||
|
||||
Reference in New Issue
Block a user