[diffusion] Support I2I/TI2I/I2V/TI2V warmup && T2I/T2V warmup bug fix (#16922)
Co-authored-by: yhyang201 <yhyang201@gmail.com> Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
yhyang201
gemini-code-assist[bot]
parent
aab640c99f
commit
2dadf63562
@@ -100,11 +100,11 @@ class SamplingParams:
|
|||||||
|
|
||||||
# Batch info
|
# Batch info
|
||||||
num_outputs_per_prompt: int = 1
|
num_outputs_per_prompt: int = 1
|
||||||
seed: int = 1024
|
seed: int = 42
|
||||||
generator_device: str = "cuda" # Device for random generator: "cuda" or "cpu"
|
generator_device: str = "cuda" # Device for random generator: "cuda" or "cpu"
|
||||||
|
|
||||||
# Original dimensions (before VAE scaling)
|
# Original dimensions (before VAE scaling)
|
||||||
num_frames: int = 125
|
num_frames: int = 1 # Default for image models
|
||||||
num_frames_round_down: bool = (
|
num_frames_round_down: bool = (
|
||||||
False # Whether to round down num_frames if it's not divisible by num_gpus
|
False # Whether to round down num_frames if it's not divisible by num_gpus
|
||||||
)
|
)
|
||||||
@@ -122,7 +122,7 @@ class SamplingParams:
|
|||||||
|
|
||||||
# Denoising parameters
|
# Denoising parameters
|
||||||
num_inference_steps: int = None
|
num_inference_steps: int = None
|
||||||
guidance_scale: float = None
|
guidance_scale: float = 1.0
|
||||||
guidance_scale_2: float = None
|
guidance_scale_2: float = None
|
||||||
true_cfg_scale: float = None # for CFG vs guidance distillation (e.g., QwenImage)
|
true_cfg_scale: float = None # for CFG vs guidance distillation (e.g., QwenImage)
|
||||||
guidance_rescale: float = 0.0
|
guidance_rescale: float = 0.0
|
||||||
|
|||||||
@@ -84,9 +84,9 @@ def _build_sampling_params_from_request(
|
|||||||
output_file_name=f"{request_id}.{ext}",
|
output_file_name=f"{request_id}.{ext}",
|
||||||
seed=seed,
|
seed=seed,
|
||||||
generator_device=generator_device,
|
generator_device=generator_device,
|
||||||
guidance_scale=guidance_scale,
|
|
||||||
num_inference_steps=num_inference_steps,
|
num_inference_steps=num_inference_steps,
|
||||||
enable_teacache=enable_teacache,
|
enable_teacache=enable_teacache,
|
||||||
|
**({"guidance_scale": guidance_scale} if guidance_scale is not None else {}),
|
||||||
**({"negative_prompt": negative_prompt} if negative_prompt is not None else {}),
|
**({"negative_prompt": negative_prompt} if negative_prompt is not None else {}),
|
||||||
**({"true_cfg_scale": true_cfg_scale} if true_cfg_scale is not None else {}),
|
**({"true_cfg_scale": true_cfg_scale} if true_cfg_scale is not None else {}),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -72,7 +72,7 @@ class VideoGenerationsRequest(BaseModel):
|
|||||||
generator_device: Optional[str] = "cuda"
|
generator_device: Optional[str] = "cuda"
|
||||||
# SGLang extensions
|
# SGLang extensions
|
||||||
num_inference_steps: Optional[int] = None
|
num_inference_steps: Optional[int] = None
|
||||||
guidance_scale: Optional[float] = None
|
guidance_scale: Optional[float] = 1.0
|
||||||
guidance_scale_2: Optional[float] = None
|
guidance_scale_2: Optional[float] = None
|
||||||
true_cfg_scale: Optional[float] = (
|
true_cfg_scale: Optional[float] = (
|
||||||
None # for CFG vs guidance distillation (e.g., QwenImage)
|
None # for CFG vs guidance distillation (e.g., QwenImage)
|
||||||
|
|||||||
@@ -213,9 +213,11 @@ async def create_video(
|
|||||||
seed=seed,
|
seed=seed,
|
||||||
generator_device=generator_device,
|
generator_device=generator_device,
|
||||||
negative_prompt=negative_prompt,
|
negative_prompt=negative_prompt,
|
||||||
guidance_scale=guidance_scale,
|
|
||||||
num_inference_steps=num_inference_steps,
|
num_inference_steps=num_inference_steps,
|
||||||
enable_teacache=enable_teacache,
|
enable_teacache=enable_teacache,
|
||||||
|
**(
|
||||||
|
{"guidance_scale": guidance_scale} if guidance_scale is not None else {}
|
||||||
|
),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||||
|
|
||||||
# SPDX-License-Identifier: Apache-2.0
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
import asyncio
|
||||||
|
import os
|
||||||
import pickle
|
import pickle
|
||||||
from collections import deque
|
from collections import deque
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
@@ -8,12 +10,14 @@ from typing import Any, List
|
|||||||
|
|
||||||
import zmq
|
import zmq
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
|
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
|
||||||
ListLorasReq,
|
ListLorasReq,
|
||||||
MergeLoraWeightsReq,
|
MergeLoraWeightsReq,
|
||||||
SetLoraReq,
|
SetLoraReq,
|
||||||
UnmergeLoraWeightsReq,
|
UnmergeLoraWeightsReq,
|
||||||
_parse_size,
|
_parse_size,
|
||||||
|
save_image_to_path,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker
|
from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core import Req
|
from sglang.multimodal_gen.runtime.pipelines_core import Req
|
||||||
@@ -29,6 +33,8 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import GREEN, RESET, init
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
MINIMUM_PICTURE_BASE64_FOR_WARMUP = "data:image/jpg;base64,iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg=="
|
||||||
|
|
||||||
|
|
||||||
class Scheduler:
|
class Scheduler:
|
||||||
"""
|
"""
|
||||||
@@ -88,6 +94,9 @@ class Scheduler:
|
|||||||
|
|
||||||
# whether we've send the necessary warmup reqs
|
# whether we've send the necessary warmup reqs
|
||||||
self.warmed_up = False
|
self.warmed_up = False
|
||||||
|
# warmup progress tracking
|
||||||
|
self._warmup_total = 0
|
||||||
|
self._warmup_processed = 0
|
||||||
|
|
||||||
self.prepare_server_warmup_reqs()
|
self.prepare_server_warmup_reqs()
|
||||||
|
|
||||||
@@ -115,9 +124,15 @@ class Scheduler:
|
|||||||
return self.worker.list_loras()
|
return self.worker.list_loras()
|
||||||
|
|
||||||
def _handle_generation(self, reqs: List[Req]):
|
def _handle_generation(self, reqs: List[Req]):
|
||||||
has_warmup = any(req.is_warmup for req in reqs)
|
warmup_reqs = [req for req in reqs if req.is_warmup]
|
||||||
if has_warmup:
|
if warmup_reqs:
|
||||||
logger.info("Processing warmup req...")
|
self._warmup_processed += len(warmup_reqs)
|
||||||
|
if self._warmup_total > 0:
|
||||||
|
logger.info(
|
||||||
|
f"Processing warmup req... ({self._warmup_processed}/{self._warmup_total})"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.info("Processing warmup req...")
|
||||||
return self.worker.execute_forward(reqs)
|
return self.worker.execute_forward(reqs)
|
||||||
|
|
||||||
def return_result(
|
def return_result(
|
||||||
@@ -149,15 +164,45 @@ class Scheduler:
|
|||||||
and self.server_args.warmup_resolutions is not None
|
and self.server_args.warmup_resolutions is not None
|
||||||
):
|
):
|
||||||
# insert warmup reqs constructed with each warmup-resolution
|
# insert warmup reqs constructed with each warmup-resolution
|
||||||
|
self._warmup_total = len(self.server_args.warmup_resolutions)
|
||||||
|
self._warmup_processed = 0
|
||||||
|
|
||||||
for resolution in self.server_args.warmup_resolutions:
|
for resolution in self.server_args.warmup_resolutions:
|
||||||
width, height = _parse_size(resolution)
|
width, height = _parse_size(resolution)
|
||||||
req = Req(
|
task_type = self.server_args.pipeline_config.task_type
|
||||||
data_type=self.server_args.pipeline_config.task_type.data_type(),
|
|
||||||
width=width,
|
if task_type in (
|
||||||
height=height,
|
ModelTaskType.I2I,
|
||||||
prompt="",
|
ModelTaskType.TI2I,
|
||||||
is_warmup=True,
|
ModelTaskType.I2V,
|
||||||
)
|
ModelTaskType.TI2V,
|
||||||
|
):
|
||||||
|
uploads_dir = os.path.join("outputs", "uploads")
|
||||||
|
os.makedirs(uploads_dir, exist_ok=True)
|
||||||
|
input_path = asyncio.run(
|
||||||
|
save_image_to_path(
|
||||||
|
MINIMUM_PICTURE_BASE64_FOR_WARMUP,
|
||||||
|
os.path.join(uploads_dir, "warmup_image.jpg"),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
req = Req(
|
||||||
|
data_type=task_type.data_type(),
|
||||||
|
width=width,
|
||||||
|
height=height,
|
||||||
|
prompt="",
|
||||||
|
negative_prompt="",
|
||||||
|
image_path=[input_path],
|
||||||
|
is_warmup=True,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
req = Req(
|
||||||
|
data_type=task_type.data_type(),
|
||||||
|
width=width,
|
||||||
|
height=height,
|
||||||
|
prompt="",
|
||||||
|
is_warmup=True,
|
||||||
|
)
|
||||||
|
|
||||||
self.waiting_queue.append((None, req))
|
self.waiting_queue.append((None, req))
|
||||||
# if server is warmed-up, set this flag to avoid req-based warmup
|
# if server is warmed-up, set this flag to avoid req-based warmup
|
||||||
self.warmed_up = True
|
self.warmed_up = True
|
||||||
@@ -182,7 +227,9 @@ class Scheduler:
|
|||||||
warmup_req.extra["cache_dit_num_inference_steps"] = req.num_inference_steps
|
warmup_req.extra["cache_dit_num_inference_steps"] = req.num_inference_steps
|
||||||
warmup_req.num_inference_steps = 1
|
warmup_req.num_inference_steps = 1
|
||||||
recv_reqs.insert(0, (identity, warmup_req))
|
recv_reqs.insert(0, (identity, warmup_req))
|
||||||
logger.info("Server warming up....")
|
self._warmup_total = 1
|
||||||
|
self._warmup_processed = 1
|
||||||
|
logger.info("Processing warmup req... (1/1)")
|
||||||
self.warmed_up = True
|
self.warmed_up = True
|
||||||
return recv_reqs
|
return recv_reqs
|
||||||
|
|
||||||
@@ -313,12 +360,23 @@ class Scheduler:
|
|||||||
)
|
)
|
||||||
if is_warmup:
|
if is_warmup:
|
||||||
if output_batch.error is None:
|
if output_batch.error is None:
|
||||||
logger.info(
|
if self._warmup_total > 0:
|
||||||
f"Warmup req processed in {GREEN}%.2f{RESET} seconds",
|
logger.info(
|
||||||
output_batch.timings.total_duration_s,
|
f"Warmup req ({self._warmup_processed}/{self._warmup_total}) processed in {GREEN}%.2f{RESET} seconds",
|
||||||
)
|
output_batch.timings.total_duration_s,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.info(
|
||||||
|
f"Warmup req processed in {GREEN}%.2f{RESET} seconds",
|
||||||
|
output_batch.timings.total_duration_s,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
logger.info(f"Warmup req processing failed")
|
if self._warmup_total > 0:
|
||||||
|
logger.info(
|
||||||
|
f"Warmup req ({self._warmup_processed}/{self._warmup_total}) processing failed"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.info(f"Warmup req processing failed")
|
||||||
|
|
||||||
# TODO: Support sending back to multiple identities if batched
|
# TODO: Support sending back to multiple identities if batched
|
||||||
self.return_result(output_batch, identities[0], is_warmup=is_warmup)
|
self.return_result(output_batch, identities[0], is_warmup=is_warmup)
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import os
|
import os
|
||||||
import pprint
|
import pprint
|
||||||
from dataclasses import asdict, dataclass, field
|
from dataclasses import MISSING, asdict, dataclass, field, fields
|
||||||
from typing import Any, Optional
|
from typing import Any, Optional
|
||||||
|
|
||||||
import PIL.Image
|
import PIL.Image
|
||||||
@@ -31,8 +31,10 @@ from sglang.multimodal_gen.utils import align_to
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
SAMPLING_PARAMS_FIELDS = {f.name for f in fields(SamplingParams)}
|
||||||
|
|
||||||
@dataclass
|
|
||||||
|
@dataclass(init=False)
|
||||||
class Req:
|
class Req:
|
||||||
"""
|
"""
|
||||||
Complete state passed through the pipeline execution.
|
Complete state passed through the pipeline execution.
|
||||||
@@ -137,6 +139,22 @@ class Req:
|
|||||||
# results
|
# results
|
||||||
output: torch.Tensor | None = None
|
output: torch.Tensor | None = None
|
||||||
|
|
||||||
|
def __init__(self, **kwargs):
|
||||||
|
# Initialize dataclass fields
|
||||||
|
for name, field in self.__class__.__dataclass_fields__.items():
|
||||||
|
if name in kwargs:
|
||||||
|
object.__setattr__(self, name, kwargs.pop(name))
|
||||||
|
elif field.default is not MISSING:
|
||||||
|
object.__setattr__(self, name, field.default)
|
||||||
|
elif field.default_factory is not MISSING:
|
||||||
|
object.__setattr__(self, name, field.default_factory())
|
||||||
|
|
||||||
|
for name, value in kwargs.items():
|
||||||
|
setattr(self, name, value)
|
||||||
|
|
||||||
|
if hasattr(self, "__post_init__"):
|
||||||
|
self.__post_init__()
|
||||||
|
|
||||||
def __getattr__(self, name: str) -> Any:
|
def __getattr__(self, name: str) -> Any:
|
||||||
"""
|
"""
|
||||||
Delegate attribute access to sampling_params if not found in Req.
|
Delegate attribute access to sampling_params if not found in Req.
|
||||||
@@ -172,13 +190,18 @@ class Req:
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
sampling_params = object.__getattribute__(self, "sampling_params")
|
sampling_params = object.__getattribute__(self, "sampling_params")
|
||||||
if sampling_params is not None and hasattr(sampling_params, name):
|
|
||||||
setattr(sampling_params, name, value)
|
|
||||||
return
|
|
||||||
except AttributeError:
|
except AttributeError:
|
||||||
# This can happen if `sampling_params` is not set yet. We'll fall through
|
sampling_params = None
|
||||||
# to setting the attribute on `self`.
|
|
||||||
pass
|
if sampling_params is not None and hasattr(sampling_params, name):
|
||||||
|
setattr(sampling_params, name, value)
|
||||||
|
return
|
||||||
|
|
||||||
|
if sampling_params is None and name in SAMPLING_PARAMS_FIELDS:
|
||||||
|
new_sp = SamplingParams()
|
||||||
|
object.__setattr__(self, "sampling_params", new_sp)
|
||||||
|
setattr(new_sp, name, value)
|
||||||
|
return
|
||||||
|
|
||||||
object.__setattr__(self, name, value)
|
object.__setattr__(self, name, value)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user