[diffusion] refactor: deprecate workloadtype (#15267)

Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
This commit is contained in:
Mick
2025-12-20 11:24:31 +08:00
committed by GitHub
co-authored by gemini-code-assist[bot]
parent 5529ab5895
commit c0f9b51992
17 changed files with 130 additions and 129 deletions
@@ -19,6 +19,7 @@ from sglang.multimodal_gen.configs.models import (
VAEConfig, VAEConfig,
) )
from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput from sglang.multimodal_gen.configs.models.encoders import BaseEncoderOutput
from sglang.multimodal_gen.configs.sample.sampling_params import DataType
from sglang.multimodal_gen.configs.utils import update_config_from_args from sglang.multimodal_gen.configs.utils import update_config_from_args
from sglang.multimodal_gen.runtime.distributed import ( from sglang.multimodal_gen.runtime.distributed import (
get_sp_parallel_rank, get_sp_parallel_rank,
@@ -36,17 +37,42 @@ from sglang.multimodal_gen.utils import (
logger = init_logger(__name__) logger = init_logger(__name__)
# NOTE: possible duplication with DataType, WorkloadType # NOTE: possible duplication with DataType
# this may focus on the model's original ability # this may focus on the model's original ability
class ModelTaskType(Enum): class ModelTaskType(Enum):
# TODO: check if I2V/TI2V models can work w/wo text
I2V = auto() # Image to Video I2V = auto() # Image to Video
T2V = auto() # Text to Video T2V = auto() # Text to Video
TI2V = auto() # Text and Image to Video TI2V = auto() # Text and Image to Video
T2I = auto() # Text to Image T2I = auto() # Text to Image
I2I = auto() # Image to Image I2I = auto() # Image to Image
TI2I = auto() # Image to Image or Text-Image to Image
def is_image_gen(self): def is_image_gen(self) -> bool:
return self == ModelTaskType.T2I or self == ModelTaskType.I2I return (
self == ModelTaskType.T2I
or self == ModelTaskType.I2I
or self == ModelTaskType.TI2I
)
def requires_image_input(self) -> bool:
return self == ModelTaskType.I2V or self == ModelTaskType.I2I
def accepts_image_input(self) -> bool:
return (
self == ModelTaskType.I2V
or self == ModelTaskType.I2I
or self == ModelTaskType.TI2I
or self == ModelTaskType.TI2V
)
def data_type(self) -> DataType:
if self.is_image_gen():
return DataType.IMAGE
else:
return DataType.VIDEO
class STA_Mode(str, Enum): class STA_Mode(str, Enum):
@@ -121,6 +147,9 @@ class PipelineConfig:
model_path: str = "" model_path: str = ""
pipeline_config_path: str | None = None pipeline_config_path: str | None = None
# precision and autocast
enable_autocast: bool = True
# generation parameters # generation parameters
# controls the timestep embedding generation # controls the timestep embedding generation
should_use_guidance: bool = True should_use_guidance: bool = True
@@ -651,11 +680,12 @@ class ImagePipelineConfig(PipelineConfig):
sp_world_size, rank_in_sp_group = get_sp_world_size(), get_sp_parallel_rank() sp_world_size, rank_in_sp_group = get_sp_world_size(), get_sp_parallel_rank()
seq_len = latents.shape[1] seq_len = latents.shape[1]
# TODO: reuse code in PipelineConfig::shard_latents_for_sp
# Pad to next multiple of SP degree if needed # Pad to next multiple of SP degree if needed
if seq_len % sp_world_size != 0: if seq_len % sp_world_size != 0:
pad_len = sp_world_size - (seq_len % sp_world_size) pad_len = sp_world_size - (seq_len % sp_world_size)
pad = torch.zeros( pad = torch.zeros(
(latents.shape[0], pad_len, latents.shape[2]), (*latents.shape[:1], pad_len, *latents.shape[2:]),
dtype=latents.dtype, dtype=latents.dtype,
device=latents.device, device=latents.device,
) )
@@ -53,6 +53,8 @@ class FluxPipelineConfig(ImagePipelineConfig):
# VAE # VAE
vae_config: VAEConfig = field(default_factory=FluxVAEConfig) vae_config: VAEConfig = field(default_factory=FluxVAEConfig)
enable_autocast: bool = False
# Text encoding stage # Text encoding stage
text_encoder_configs: tuple[EncoderConfig, ...] = field( text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (CLIPTextConfig(), T5Config()) default_factory=lambda: (CLIPTextConfig(), T5Config())
@@ -405,7 +407,7 @@ def flux2_pack_latents(latents):
class Flux2PipelineConfig(FluxPipelineConfig): class Flux2PipelineConfig(FluxPipelineConfig):
embedded_cfg_scale: float = 4.0 embedded_cfg_scale: float = 4.0
task_type: ModelTaskType = ModelTaskType.I2I task_type: ModelTaskType = ModelTaskType.TI2I
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",)) text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",))
@@ -80,6 +80,8 @@ class QwenImagePipelineConfig(ImagePipelineConfig):
# VAE # VAE
vae_config: VAEConfig = field(default_factory=QwenImageVAEConfig) vae_config: VAEConfig = field(default_factory=QwenImageVAEConfig)
enable_autocast: bool = False
# Text encoding stage # Text encoding stage
text_encoder_configs: tuple[EncoderConfig, ...] = field( text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (Qwen2_5VLConfig(),) default_factory=lambda: (Qwen2_5VLConfig(),)
@@ -16,7 +16,6 @@ from dataclasses import dataclass
from enum import Enum, auto from enum import Enum, auto
from typing import Any from typing import Any
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
from sglang.multimodal_gen.utils import StoreBoolean, align_to from sglang.multimodal_gen.utils import StoreBoolean, align_to
@@ -195,33 +194,55 @@ class SamplingParams:
def __post_init__(self) -> None: def __post_init__(self) -> None:
assert self.num_frames >= 1 assert self.num_frames >= 1
self.data_type = DataType.VIDEO if self.num_frames > 1 else DataType.IMAGE
if self.width is None: if self.width is None:
self.width_not_provided = True self.width_not_provided = True
if self.height is None: if self.height is None:
self.height_not_provided = True self.height_not_provided = True
self._validate()
# Allow env var to override num_inference_steps (for faster CI testing on AMD) # Allow env var to override num_inference_steps (for faster CI testing on AMD)
env_steps = os.environ.get("SGLANG_TEST_NUM_INFERENCE_STEPS") env_steps = os.environ.get("SGLANG_TEST_NUM_INFERENCE_STEPS")
if env_steps is not None and self.num_inference_steps is not None: if env_steps is not None and self.num_inference_steps is not None:
self.num_inference_steps = int(env_steps) self.num_inference_steps = int(env_steps)
def _validate(self):
"""
check if the sampling params is correct by itself
"""
if self.prompt_path and not self.prompt_path.endswith(".txt"):
raise ValueError("prompt_path must be a txt file")
def check_sampling_param(self): def check_sampling_param(self):
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 _validate_with_pipeline_config(self, pipeline_config):
"""
check if the sampling params is compatible and valid with server_args
"""
if pipeline_config.task_type.requires_image_input():
# requires image input
if self.image_path is None:
raise ValueError(
f"Served model with task type '{pipeline_config.task_type.name}' requires an 'image_path' input, but none was provided"
)
def _adjust( def _adjust(
self, self,
server_args: ServerArgs, server_args,
): ):
""" """
final adjustment, called after merged with user params final adjustment, called after merged with user params
""" """
# TODO: SamplingParams should not rely on ServerArgs
pipeline_config = server_args.pipeline_config pipeline_config = server_args.pipeline_config
if not isinstance(self.prompt, str): if not isinstance(self.prompt, str):
raise TypeError(f"`prompt` must be a string, but got {type(self.prompt)}") raise TypeError(f"`prompt` must be a string, but got {type(self.prompt)}")
self.data_type = server_args.pipeline_config.task_type.data_type()
# Process negative prompt # Process negative prompt
if self.negative_prompt is not None and not self.negative_prompt.isspace(): if self.negative_prompt is not None and not self.negative_prompt.isspace():
# avoid stripping default negative prompt: ' ' for qwen-image # avoid stripping default negative prompt: ' ' for qwen-image
@@ -260,7 +281,6 @@ class SamplingParams:
# settle num_frames # settle num_frames
logger.debug(f"num_frames set to 1 for image generation model") logger.debug(f"num_frames set to 1 for image generation model")
self.num_frames = 1 self.num_frames = 1
self.data_type = DataType.IMAGE
elif self.adjust_frames: elif self.adjust_frames:
# NOTE: We must apply adjust_num_frames BEFORE the SP alignment logic below. # NOTE: We must apply adjust_num_frames BEFORE the SP alignment logic below.
# If we apply it after, adjust_num_frames might modify the frame count # If we apply it after, adjust_num_frames might modify the frame count
@@ -320,7 +340,7 @@ class SamplingParams:
self.num_frames = new_num_frames self.num_frames = new_num_frames
self._set_output_file_name() self._set_output_file_name()
self.log(server_args=server_args) self.log(pipeline_config=server_args.pipeline_config)
@classmethod @classmethod
def from_pretrained(cls, model_path: str, **kwargs) -> "SamplingParams": def from_pretrained(cls, model_path: str, **kwargs) -> "SamplingParams":
@@ -339,6 +359,8 @@ class SamplingParams:
sampling_params._merge_with_user_params(user_sampling_params) sampling_params._merge_with_user_params(user_sampling_params)
sampling_params._adjust(server_args) sampling_params._adjust(server_args)
sampling_params._validate_with_pipeline_config(server_args.pipeline_config)
return sampling_params return sampling_params
def output_size_str(self) -> str: def output_size_str(self) -> str:
@@ -658,7 +680,7 @@ class SamplingParams:
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)
def log(self, server_args: ServerArgs): def log(self, pipeline_config):
# TODO: in some cases (e.g., TI2I), height and weight might be undecided at this moment # TODO: in some cases (e.g., TI2I), height and weight might be undecided at this moment
if self.height: if self.height:
target_height = align_to(self.height, 16) target_height = align_to(self.height, 16)
@@ -680,9 +702,9 @@ class SamplingParams:
infer_steps: {self.num_inference_steps} infer_steps: {self.num_inference_steps}
num_outputs_per_prompt: {self.num_outputs_per_prompt} num_outputs_per_prompt: {self.num_outputs_per_prompt}
guidance_scale: {self.guidance_scale} guidance_scale: {self.guidance_scale}
embedded_guidance_scale: {server_args.pipeline_config.embedded_cfg_scale} embedded_guidance_scale: {pipeline_config.embedded_cfg_scale}
n_tokens: {self.n_tokens} n_tokens: {self.n_tokens}
flow_shift: {server_args.pipeline_config.flow_shift} flow_shift: {pipeline_config.flow_shift}
image_path: {self.image_path} image_path: {self.image_path}
save_output: {self.save_output} save_output: {self.save_output}
output_file_path: {self.output_file_path()} output_file_path: {self.output_file_path()}
@@ -43,7 +43,6 @@ curl -sS -X GET "http://localhost:30010/models"
{ {
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers", "model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"task_type": "T2V", "task_type": "T2V",
"workload_type": "serving",
"pipeline_name": "wan_pipeline", "pipeline_name": "wan_pipeline",
"pipeline_class": "WanPipeline", "pipeline_class": "WanPipeline",
"num_gpus": 4, "num_gpus": 4,
@@ -1,7 +1,7 @@
# 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.runtime.utils.logging_utils import suppress_loggers from sglang.multimodal_gen.runtime.utils.logging_utils import (
global_suppress_loggers,
# globally suppress some obsessive loggers suppress_loggers,
suppress_loggers(
["imageio", "imageio_ffmpeg", "PIL", "PIL_Image", "multipart", "filelock"]
) )
global_suppress_loggers()
@@ -69,7 +69,6 @@ async def get_models(request: Request):
"task_type": server_args.pipeline_config.task_type.name, "task_type": server_args.pipeline_config.task_type.name,
"dit_precision": server_args.pipeline_config.dit_precision, "dit_precision": server_args.pipeline_config.dit_precision,
"vae_precision": server_args.pipeline_config.vae_precision, "vae_precision": server_args.pipeline_config.vae_precision,
"workload_type": server_args.workload_type.value,
} }
if model_info: if model_info:
@@ -153,7 +153,7 @@ def launch_http_server_only(server_args):
app = create_app(server_args) app = create_app(server_args)
uvicorn.run( uvicorn.run(
app, app,
log_config=None, use_colors=True,
log_level=server_args.log_level, log_level=server_args.log_level,
host=server_args.host, host=server_args.host,
port=server_args.port, port=server_args.port,
@@ -14,7 +14,11 @@ from diffusers.models.embeddings import (
) )
from diffusers.models.embeddings import PixArtAlphaTextProjection, TimestepEmbedding from diffusers.models.embeddings import PixArtAlphaTextProjection, TimestepEmbedding
from diffusers.models.embeddings import Timesteps as _Timesteps from diffusers.models.embeddings import Timesteps as _Timesteps
try:
from sgl_kernel.elementwise import timestep_embedding as timestep_embedding_cuda from sgl_kernel.elementwise import timestep_embedding as timestep_embedding_cuda
except Exception as _e:
pass
from sglang.multimodal_gen.runtime.layers.activation import get_act_fn from sglang.multimodal_gen.runtime.layers.activation import get_act_fn
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
@@ -49,7 +49,7 @@ def build_pipeline(
# instantiate the pipelines # instantiate the pipelines
pipeline = pipeline_cls(model_path, server_args) pipeline = pipeline_cls(model_path, server_args)
logger.info("Pipelines instantiated") logger.info("Pipeline instantiated")
return cast(PipelineWithLoRA, pipeline) return cast(PipelineWithLoRA, pipeline)
@@ -20,7 +20,10 @@ from tqdm.auto import tqdm
from sglang.multimodal_gen import envs from sglang.multimodal_gen import envs
from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType, STA_Mode from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType, STA_Mode
from sglang.multimodal_gen.configs.pipeline_configs.wan import Wan2_2_TI2V_5B_Config from sglang.multimodal_gen.configs.pipeline_configs.wan import (
Wan2_2_TI2V_5B_Config,
WanI2V480PConfig,
)
from sglang.multimodal_gen.runtime.distributed import ( from sglang.multimodal_gen.runtime.distributed import (
cfg_model_parallel_all_reduce, cfg_model_parallel_all_reduce,
get_local_torch_device, get_local_torch_device,
@@ -644,6 +647,11 @@ class DenoisingStage(PipelineStage):
self.device, self.device,
getattr(self.transformer, "rotary_emb", None), getattr(self.transformer, "rotary_emb", None),
dtype=target_dtype, dtype=target_dtype,
)
| dict(
encoder_hidden_states=server_args.pipeline_config.get_pos_prompt_embeds(
batch
)
), ),
) )
@@ -659,6 +667,11 @@ class DenoisingStage(PipelineStage):
self.device, self.device,
getattr(self.transformer, "rotary_emb", None), getattr(self.transformer, "rotary_emb", None),
dtype=target_dtype, dtype=target_dtype,
)
| dict(
encoder_hidden_states=server_args.pipeline_config.get_neg_prompt_embeds(
batch
)
), ),
) )
else: else:
@@ -767,10 +780,11 @@ class DenoisingStage(PipelineStage):
else: else:
batch.did_sp_shard_latents = False batch.did_sp_shard_latents = False
# For I2I tasks like QwenImageEdit, the image_latent (input image) should be # For I2I tasks like QwenImageEdit, where the image latents is provided as condition, the image_latent (input image) should be
# replicated on all SP ranks, not sharded, as it provides global context. # replicated on all SP ranks, not sharded, as it provides global context.
# For Wan2_2_TI2V_5B_Config, it has very special settings
if ( if (
server_args.pipeline_config.task_type != ModelTaskType.I2I isinstance(server_args.pipeline_config, WanI2V480PConfig)
and batch.image_latent is not None and batch.image_latent is not None
): ):
batch.image_latent, _ = server_args.pipeline_config.shard_latents_for_sp( batch.image_latent, _ = server_args.pipeline_config.shard_latents_for_sp(
@@ -989,10 +1003,6 @@ class DenoisingStage(PipelineStage):
): ):
with self.progress_bar(total=num_inference_steps) as progress_bar: with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t_host in enumerate(timesteps_cpu): for i, t_host in enumerate(timesteps_cpu):
# Skip if interrupted
if hasattr(self, "interrupt") and self.interrupt:
continue
with StageProfiler( with StageProfiler(
f"denoising_step_{i}", logger=logger, timings=batch.timings f"denoising_step_{i}", logger=logger, timings=batch.timings
): ):
@@ -1225,14 +1235,12 @@ class DenoisingStage(PipelineStage):
current_model, current_model,
latent_model_input, latent_model_input,
timestep, timestep,
prompt_embeds,
target_dtype, target_dtype,
guidance: torch.Tensor, guidance: torch.Tensor,
**kwargs, **kwargs,
): ):
return current_model( return current_model(
hidden_states=latent_model_input, hidden_states=latent_model_input,
encoder_hidden_states=prompt_embeds,
timestep=timestep, timestep=timestep,
guidance=guidance, guidance=guidance,
**kwargs, **kwargs,
@@ -1289,9 +1297,6 @@ class DenoisingStage(PipelineStage):
current_model=current_model, current_model=current_model,
latent_model_input=latent_model_input, latent_model_input=latent_model_input,
timestep=timestep, timestep=timestep,
prompt_embeds=server_args.pipeline_config.get_pos_prompt_embeds(
batch
),
target_dtype=target_dtype, target_dtype=target_dtype,
guidance=guidance, guidance=guidance,
**image_kwargs, **image_kwargs,
@@ -1317,9 +1322,6 @@ class DenoisingStage(PipelineStage):
current_model=current_model, current_model=current_model,
latent_model_input=latent_model_input, latent_model_input=latent_model_input,
timestep=timestep, timestep=timestep,
prompt_embeds=server_args.pipeline_config.get_neg_prompt_embeds(
batch
),
target_dtype=target_dtype, target_dtype=target_dtype,
guidance=guidance, guidance=guidance,
**image_kwargs, **image_kwargs,
@@ -89,7 +89,6 @@ class DmdDenoisingStage(DenoisingStage):
) )
pos_cond_kwargs = prepared_vars["pos_cond_kwargs"] pos_cond_kwargs = prepared_vars["pos_cond_kwargs"]
prompt_embeds = prepared_vars["prompt_embeds"]
denoising_loop_start_time = time.time() denoising_loop_start_time = time.time()
with self.progress_bar(total=len(timesteps)) as progress_bar: with self.progress_bar(total=len(timesteps)) as progress_bar:
@@ -142,9 +141,8 @@ class DmdDenoisingStage(DenoisingStage):
): ):
# Run transformer # Run transformer
pred_noise = self.transformer( pred_noise = self.transformer(
latent_model_input.permute(0, 2, 1, 3, 4), hidden_states=latent_model_input.permute(0, 2, 1, 3, 4),
prompt_embeds, timestep=t_expand,
t_expand,
guidance=guidance_expand, guidance=guidance_expand,
**image_kwargs, **image_kwargs,
**pos_cond_kwargs, **pos_cond_kwargs,
@@ -79,7 +79,10 @@ class InputValidationStage(PipelineStage):
preprocess condition image preprocess condition image
NOTE: condition image resizing is only allowed in InputValidationStage NOTE: condition image resizing is only allowed in InputValidationStage
""" """
if server_args.pipeline_config.task_type == ModelTaskType.I2I: if batch.condition_image is not None and (
server_args.pipeline_config.task_type == ModelTaskType.I2I
or server_args.pipeline_config.task_type == ModelTaskType.TI2I
):
# calculate new condition image size # calculate new condition image size
if not isinstance(batch.condition_image, list): if not isinstance(batch.condition_image, list):
batch.condition_image = [batch.condition_image] batch.condition_image = [batch.condition_image]
@@ -17,12 +17,7 @@ from dataclasses import field
from enum import Enum from enum import Enum
from typing import Any, Optional from typing import Any, Optional
from sglang.multimodal_gen.configs.pipeline_configs import FluxPipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig, STA_Mode from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig, STA_Mode
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageEditPipelineConfig,
QwenImagePipelineConfig,
)
from sglang.multimodal_gen.runtime.platforms import ( from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum, AttentionBackendEnum,
current_platform, current_platform,
@@ -175,35 +170,6 @@ class ExecutionMode(str, Enum):
return [mode.value for mode in cls] return [mode.value for mode in cls]
class WorkloadType(str, Enum):
"""
Enumeration for different workload types.
Inherits from str to allow string comparison for backward compatibility.
"""
I2V = "i2v" # Image to Video
T2V = "t2v" # Text to Video
T2I = "t2i" # Text to Image
I2I = "i2i" # Image to Image
@classmethod
def from_string(cls, value: str) -> "WorkloadType":
"""Convert string to WorkloadType enum."""
try:
return cls(value.lower())
except ValueError:
raise ValueError(
f"Invalid workload type: {value}. Must be one of: {', '.join([m.value for m in cls])}"
) from None
@classmethod
def choices(cls) -> list[str]:
"""Get all available choices as strings for argparse."""
return [workload.value for workload in cls]
# args for sgl_diffusion framework
@dataclasses.dataclass @dataclasses.dataclass
class ServerArgs: class ServerArgs:
# Model and path configuration (for convenience) # Model and path configuration (for convenience)
@@ -215,9 +181,6 @@ class ServerArgs:
# Running mode # Running mode
mode: ExecutionMode = ExecutionMode.INFERENCE mode: ExecutionMode = ExecutionMode.INFERENCE
# Workload type
workload_type: WorkloadType = WorkloadType.T2V
# Cache strategy # Cache strategy
cache_strategy: str = "none" cache_strategy: str = "none"
@@ -279,7 +242,7 @@ class ServerArgs:
# Compilation # Compilation
enable_torch_compile: bool = False enable_torch_compile: bool = False
disable_autocast: bool = False disable_autocast: bool | None = None
# VSA parameters # VSA parameters
VSA_sparsity: float = 0.0 # inference/validation sparsity VSA_sparsity: float = 0.0 # inference/validation sparsity
@@ -408,15 +371,6 @@ class ServerArgs:
help="The mode to run SGLang-diffusion", help="The mode to run SGLang-diffusion",
) )
# Workload type
parser.add_argument(
"--workload-type",
type=str,
choices=WorkloadType.choices(),
default=ServerArgs.workload_type.value,
help="The workload type",
)
# distributed_executor_backend # distributed_executor_backend
parser.add_argument( parser.add_argument(
"--distributed-executor-backend", "--distributed-executor-backend",
@@ -796,10 +750,6 @@ class ServerArgs:
if "mode" in kwargs and isinstance(kwargs["mode"], str): if "mode" in kwargs and isinstance(kwargs["mode"], str):
kwargs["mode"] = ExecutionMode.from_string(kwargs["mode"]) kwargs["mode"] = ExecutionMode.from_string(kwargs["mode"])
# Convert workload_type string to enum if necessary
if "workload_type" in kwargs and isinstance(kwargs["workload_type"], str):
kwargs["workload_type"] = WorkloadType.from_string(kwargs["workload_type"])
kwargs["pipeline_config"] = PipelineConfig.from_kwargs(kwargs) kwargs["pipeline_config"] = PipelineConfig.from_kwargs(kwargs)
return cls(**kwargs) return cls(**kwargs)
@@ -896,13 +846,10 @@ class ServerArgs:
self.use_fsdp_inference = False self.use_fsdp_inference = False
# autocast # autocast
is_flux = ( if self.disable_autocast is None:
isinstance(self.pipeline_config, FluxPipelineConfig) self.disable_autocast = not self.pipeline_config.enable_autocast
or isinstance(self.pipeline_config, QwenImagePipelineConfig) else:
or isinstance(self.pipeline_config, QwenImageEditPipelineConfig) self.disable_autocast = False
)
if is_flux:
self.disable_autocast = True
# Validate mode consistency # Validate mode consistency
assert isinstance( assert isinstance(
@@ -912,14 +859,6 @@ class ServerArgs:
self.mode in ExecutionMode.choices() self.mode in ExecutionMode.choices()
), f"Invalid execution mode: {self.mode}" ), f"Invalid execution mode: {self.mode}"
# Validate workload type
assert isinstance(
self.workload_type, WorkloadType
), f"Workload type must be a WorkloadType enum, got {type(self.workload_type)}"
assert (
self.workload_type in WorkloadType.choices()
), f"Invalid workload type: {self.workload_type}"
if self.tp_size == -1: if self.tp_size == -1:
self.tp_size = 1 self.tp_size = 1
@@ -396,6 +396,21 @@ def suppress_loggers(loggers_to_suppress: list[str], level: int = logging.WARNIN
return original_levels return original_levels
def global_suppress_loggers():
# globally suppress some obsessive loggers
suppress_loggers(
[
"imageio",
"imageio_ffmpeg",
"PIL",
"PIL_Image",
"multipart",
"filelock",
"urllib3",
]
)
@contextmanager @contextmanager
def suppress_other_loggers(not_suppress_on_main_rank: bool = False): def suppress_other_loggers(not_suppress_on_main_rank: bool = False):
""" """
@@ -22,7 +22,10 @@ from openai import Client, OpenAI
from sglang.multimodal_gen.benchmarks.compare_perf import calculate_upper_bound from sglang.multimodal_gen.benchmarks.compare_perf import calculate_upper_bound
from sglang.multimodal_gen.runtime.utils.common import is_hip, kill_process_tree from sglang.multimodal_gen.runtime.utils.common import is_hip, kill_process_tree
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import (
global_suppress_loggers,
init_logger,
)
from sglang.multimodal_gen.runtime.utils.perf_logger import RequestPerfRecord from sglang.multimodal_gen.runtime.utils.perf_logger import RequestPerfRecord
from sglang.multimodal_gen.test.server.testcase_configs import ( from sglang.multimodal_gen.test.server.testcase_configs import (
DiffusionSamplingParams, DiffusionSamplingParams,
@@ -40,6 +43,8 @@ from sglang.multimodal_gen.test.test_utils import (
logger = init_logger(__name__) logger = init_logger(__name__)
global_suppress_loggers()
def download_image_from_url(url: str) -> Path: def download_image_from_url(url: str) -> Path:
"""Download an image from a URL to a temporary file. """Download an image from a URL to a temporary file.
-19
View File
@@ -21,12 +21,8 @@ from functools import lru_cache, partial, wraps
from typing import Any, TypeVar, cast from typing import Any, TypeVar, cast
import cloudpickle import cloudpickle
import imageio
import numpy as np
import torch import torch
import torchvision
import yaml import yaml
from einops import rearrange
from remote_pdb import RemotePdb from remote_pdb import RemotePdb
from torch.distributed.fsdp import MixedPrecisionPolicy from torch.distributed.fsdp import MixedPrecisionPolicy
@@ -788,21 +784,6 @@ def best_output_size(w, h, dw, dh, expected_area):
return ow2, oh2 return ow2, oh2
def save_decoded_latents_as_video(
decoded_latents: list[torch.Tensor], output_path: str, fps: int
):
# Process outputs
videos = rearrange(decoded_latents, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
os.makedirs(os.path.dirname(output_path), exist_ok=True)
imageio.mimsave(output_path, frames, fps=fps, format="mp4")
def calculate_dimensions(target_area, ratio): def calculate_dimensions(target_area, ratio):
width = math.sqrt(target_area * ratio) width = math.sqrt(target_area * ratio)
height = width / ratio height = width / ratio