[diffusion] refactor: refactor sampling params (#13706)
This commit is contained in:
@@ -5,12 +5,12 @@ import argparse
|
|||||||
import dataclasses
|
import dataclasses
|
||||||
import hashlib
|
import hashlib
|
||||||
import json
|
import json
|
||||||
|
import math
|
||||||
import os.path
|
import os.path
|
||||||
import re
|
import re
|
||||||
import time
|
import time
|
||||||
import unicodedata
|
import unicodedata
|
||||||
import uuid
|
import uuid
|
||||||
from copy import deepcopy
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from enum import Enum, auto
|
from enum import Enum, auto
|
||||||
from typing import Any
|
from typing import Any
|
||||||
@@ -137,7 +137,7 @@ class SamplingParams:
|
|||||||
return_trajectory_latents: bool = False # returns all latents for each timestep
|
return_trajectory_latents: bool = False # returns all latents for each timestep
|
||||||
return_trajectory_decoded: bool = False # returns decoded 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
|
# add extension if needed
|
||||||
if not any(
|
if not any(
|
||||||
self.output_file_name.endswith(ext)
|
self.output_file_name.endswith(ext)
|
||||||
@@ -147,7 +147,7 @@ class SamplingParams:
|
|||||||
f"{self.output_file_name}.{self.data_type.get_default_extension()}"
|
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
|
# settle output_file_name
|
||||||
if (
|
if (
|
||||||
self.output_file_name is None
|
self.output_file_name is None
|
||||||
@@ -178,7 +178,7 @@ class SamplingParams:
|
|||||||
self.output_file_name = _sanitize_filename(self.output_file_name)
|
self.output_file_name = _sanitize_filename(self.output_file_name)
|
||||||
|
|
||||||
# Ensure a proper extension is present
|
# Ensure a proper extension is present
|
||||||
self.set_output_file_ext()
|
self._set_output_file_ext()
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
assert self.num_frames >= 1
|
assert self.num_frames >= 1
|
||||||
@@ -195,6 +195,93 @@ 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(
|
||||||
|
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:
|
def update(self, source_dict: dict[str, Any]) -> None:
|
||||||
for key, value in source_dict.items():
|
for key, value in source_dict.items():
|
||||||
if hasattr(self, key):
|
if hasattr(self, key):
|
||||||
@@ -220,9 +307,15 @@ class SamplingParams:
|
|||||||
sampling_params = cls(**kwargs)
|
sampling_params = cls(**kwargs)
|
||||||
return sampling_params
|
return sampling_params
|
||||||
|
|
||||||
def from_user_sampling_params(self, user_params):
|
@staticmethod
|
||||||
sampling_params = deepcopy(self)
|
def from_user_sampling_params_args(model_path: str, server_args, *args, **kwargs):
|
||||||
sampling_params._merge_with_user_params(user_params)
|
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
|
return sampling_params
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -264,7 +264,8 @@ class DiffGenerator:
|
|||||||
else DataType.VIDEO
|
else DataType.VIDEO
|
||||||
)
|
)
|
||||||
pretrained_sampling_params.data_type = data_type
|
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] = []
|
requests: list[Req] = []
|
||||||
for output_idx, p in enumerate(prompts):
|
for output_idx, p in enumerate(prompts):
|
||||||
@@ -272,7 +273,6 @@ class DiffGenerator:
|
|||||||
current_sampling_params.prompt = p
|
current_sampling_params.prompt = p
|
||||||
requests.append(
|
requests.append(
|
||||||
prepare_request(
|
prepare_request(
|
||||||
p,
|
|
||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
sampling_params=current_sampling_params,
|
sampling_params=current_sampling_params,
|
||||||
)
|
)
|
||||||
@@ -310,21 +310,11 @@ class DiffGenerator:
|
|||||||
continue
|
continue
|
||||||
for output_idx, sample in enumerate(output_batch.output):
|
for output_idx, sample in enumerate(output_batch.output):
|
||||||
num_outputs = len(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(
|
frames = self.post_process_sample(
|
||||||
sample,
|
sample,
|
||||||
fps=req.fps,
|
fps=req.fps,
|
||||||
save_output=req.save_output,
|
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,
|
data_type=req.data_type,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -56,12 +56,10 @@ def _build_sampling_params_from_request(
|
|||||||
) -> SamplingParams:
|
) -> SamplingParams:
|
||||||
width, height = _parse_size(size)
|
width, height = _parse_size(size)
|
||||||
ext = _choose_ext(output_format, background)
|
ext = _choose_ext(output_format, background)
|
||||||
|
|
||||||
server_args = get_global_server_args()
|
server_args = get_global_server_args()
|
||||||
sampling_params = SamplingParams.from_pretrained(server_args.model_path)
|
|
||||||
|
|
||||||
# Build user params
|
# Build user params
|
||||||
user_params = SamplingParams(
|
sampling_params = SamplingParams.from_user_sampling_params_args(
|
||||||
|
model_path=server_args.model_path,
|
||||||
request_id=request_id,
|
request_id=request_id,
|
||||||
prompt=prompt,
|
prompt=prompt,
|
||||||
image_path=image_path,
|
image_path=image_path,
|
||||||
@@ -70,18 +68,9 @@ def _build_sampling_params_from_request(
|
|||||||
height=height,
|
height=height,
|
||||||
num_outputs_per_prompt=max(1, min(int(n or 1), 10)),
|
num_outputs_per_prompt=max(1, min(int(n or 1), 10)),
|
||||||
save_output=True,
|
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
|
return sampling_params
|
||||||
|
|
||||||
|
|
||||||
@@ -107,7 +96,6 @@ def _build_req_from_sampling(s: SamplingParams) -> Req:
|
|||||||
async def generations(
|
async def generations(
|
||||||
request: ImageGenerationsRequest,
|
request: ImageGenerationsRequest,
|
||||||
):
|
):
|
||||||
|
|
||||||
request_id = generate_request_id()
|
request_id = generate_request_id()
|
||||||
sampling = _build_sampling_params_from_request(
|
sampling = _build_sampling_params_from_request(
|
||||||
request_id=request_id,
|
request_id=request_id,
|
||||||
@@ -118,7 +106,6 @@ async def generations(
|
|||||||
background=request.background,
|
background=request.background,
|
||||||
)
|
)
|
||||||
batch = prepare_request(
|
batch = prepare_request(
|
||||||
prompt=request.prompt,
|
|
||||||
server_args=get_global_server_args(),
|
server_args=get_global_server_args(),
|
||||||
sampling_params=sampling,
|
sampling_params=sampling,
|
||||||
)
|
)
|
||||||
@@ -175,7 +162,6 @@ async def edits(
|
|||||||
background: Optional[str] = Form("auto"),
|
background: Optional[str] = Form("auto"),
|
||||||
user: Optional[str] = Form(None),
|
user: Optional[str] = Form(None),
|
||||||
):
|
):
|
||||||
|
|
||||||
request_id = generate_request_id()
|
request_id = generate_request_id()
|
||||||
# Resolve images from either `image` or `image[]` (OpenAI SDK sends `image[]` when list is provided)
|
# Resolve images from either `image` or `image[]` (OpenAI SDK sends `image[]` when list is provided)
|
||||||
images = image or image_array
|
images = image or image_array
|
||||||
|
|||||||
@@ -42,6 +42,8 @@ logger = init_logger(__name__)
|
|||||||
router = APIRouter(prefix="/v1/videos", tags=["videos"])
|
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(
|
def _build_sampling_params_from_request(
|
||||||
request_id: str, request: VideoGenerationsRequest
|
request_id: str, request: VideoGenerationsRequest
|
||||||
) -> SamplingParams:
|
) -> 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
|
request.num_frames if request.num_frames is not None else derived_num_frames
|
||||||
)
|
)
|
||||||
server_args = get_global_server_args()
|
server_args = get_global_server_args()
|
||||||
# TODO: should we cache this sampling_params?
|
sampling_params = SamplingParams.from_user_sampling_params_args(
|
||||||
sampling_params = SamplingParams.from_pretrained(server_args.model_path)
|
model_path=server_args.model_path,
|
||||||
user_params = SamplingParams(
|
|
||||||
request_id=request_id,
|
request_id=request_id,
|
||||||
prompt=request.prompt,
|
prompt=request.prompt,
|
||||||
num_frames=num_frames,
|
num_frames=num_frames,
|
||||||
@@ -67,10 +68,10 @@ def _build_sampling_params_from_request(
|
|||||||
height=height,
|
height=height,
|
||||||
image_path=request.input_reference,
|
image_path=request.input_reference,
|
||||||
save_output=True,
|
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
|
return sampling_params
|
||||||
|
|
||||||
|
|
||||||
@@ -195,7 +196,6 @@ async def create_video(
|
|||||||
|
|
||||||
# Build Req for scheduler
|
# Build Req for scheduler
|
||||||
batch = prepare_request(
|
batch = prepare_request(
|
||||||
prompt=req.prompt,
|
|
||||||
server_args=get_global_server_args(),
|
server_args=get_global_server_args(),
|
||||||
sampling_params=sampling_params,
|
sampling_params=sampling_params,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -9,13 +9,12 @@ diffusion models.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
import math
|
|
||||||
|
|
||||||
# Suppress verbose logging from imageio, which is triggered when saving images.
|
# Suppress verbose logging from imageio, which is triggered when saving images.
|
||||||
logging.getLogger("imageio").setLevel(logging.WARNING)
|
logging.getLogger("imageio").setLevel(logging.WARNING)
|
||||||
logging.getLogger("imageio_ffmpeg").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.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
|
||||||
@@ -24,97 +23,7 @@ from sglang.multimodal_gen.utils import shallow_asdict
|
|||||||
logger = init_logger(__name__)
|
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(
|
def prepare_request(
|
||||||
prompt: str,
|
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
sampling_params: SamplingParams,
|
sampling_params: SamplingParams,
|
||||||
) -> Req:
|
) -> Req:
|
||||||
@@ -123,20 +32,16 @@ def prepare_request(
|
|||||||
|
|
||||||
"""
|
"""
|
||||||
# Create a copy of inference args to avoid modifying the original
|
# Create a copy of inference args to avoid modifying the original
|
||||||
|
|
||||||
sampling_params = prepare_sampling_params(prompt, server_args, sampling_params)
|
|
||||||
|
|
||||||
req = Req(
|
req = Req(
|
||||||
**shallow_asdict(sampling_params),
|
**shallow_asdict(sampling_params),
|
||||||
VSA_sparsity=server_args.VSA_sparsity,
|
VSA_sparsity=server_args.VSA_sparsity,
|
||||||
)
|
)
|
||||||
# req.set_width_and_height(server_args)
|
req.adjust_size(server_args)
|
||||||
|
|
||||||
# if (req.width <= 0
|
if req.width <= 0 or req.height <= 0:
|
||||||
# or req.height <= 0):
|
raise ValueError(
|
||||||
# raise ValueError(
|
f"Height, width must be positive integers, got "
|
||||||
# f"Height, width must be positive integers, got "
|
f"height={req.height}, width={req.width}"
|
||||||
# f"height={req.height}, width={req.width}"
|
)
|
||||||
# )
|
|
||||||
|
|
||||||
return req
|
return req
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ in a functional manner, reducing the need for explicit parameter passing.
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
import pprint
|
import pprint
|
||||||
from dataclasses import asdict, dataclass, field
|
from dataclasses import asdict, dataclass, field
|
||||||
from typing import TYPE_CHECKING, Any, Optional
|
from typing import TYPE_CHECKING, Any, Optional
|
||||||
@@ -187,6 +188,18 @@ class Req:
|
|||||||
batch_size *= self.num_outputs_per_prompt
|
batch_size *= self.num_outputs_per_prompt
|
||||||
return batch_size
|
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):
|
def __post_init__(self):
|
||||||
"""Initialize dependent fields after dataclass initialization."""
|
"""Initialize dependent fields after dataclass initialization."""
|
||||||
# Set do_classifier_free_guidance based on guidance scale and negative prompt
|
# 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:
|
if self.guidance_scale_2 is None:
|
||||||
self.guidance_scale_2 = self.guidance_scale
|
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:
|
if self.height is None or self.width is None:
|
||||||
width, height = server_args.pipeline_config.adjust_size(
|
width, height = server_args.pipeline_config.adjust_size(
|
||||||
self.width, self.height, self.pil_image
|
self.width, self.height, self.pil_image
|
||||||
|
|||||||
@@ -273,7 +273,6 @@ Consider updating perf_baselines.json with the snippets below:
|
|||||||
response_format="b64_json",
|
response_format="b64_json",
|
||||||
)
|
)
|
||||||
rid = response.headers.get("x-request-id", "")
|
rid = response.headers.get("x-request-id", "")
|
||||||
print(f"{response=}")
|
|
||||||
|
|
||||||
result = response.parse()
|
result = response.parse()
|
||||||
validate_image(result.data[0].b64_json)
|
validate_image(result.data[0].b64_json)
|
||||||
|
|||||||
Reference in New Issue
Block a user