[diffusion] chore: set allowing overriding protected fields of sampling params as default behavior (#14471)
This commit is contained in:
@@ -138,7 +138,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
|
||||||
# if True, allow user params to override subclass-defined protected fields
|
# if True, allow user params to override subclass-defined protected fields
|
||||||
override_protected_fields: bool = False
|
no_override_protected_fields: bool = True
|
||||||
# whether to adjust num_frames for multi-GPU friendly splitting (default: True)
|
# whether to adjust num_frames for multi-GPU friendly splitting (default: True)
|
||||||
adjust_frames: bool = True
|
adjust_frames: bool = True
|
||||||
|
|
||||||
@@ -290,15 +290,6 @@ class SamplingParams:
|
|||||||
self._set_output_file_name()
|
self._set_output_file_name()
|
||||||
self.log(server_args=server_args)
|
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):
|
|
||||||
setattr(self, key, value)
|
|
||||||
else:
|
|
||||||
logger.exception("%s has no attribute %s", type(self).__name__, key)
|
|
||||||
|
|
||||||
self.__post_init__()
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_pretrained(cls, model_path: str, **kwargs) -> "SamplingParams":
|
def from_pretrained(cls, model_path: str, **kwargs) -> "SamplingParams":
|
||||||
from sglang.multimodal_gen.registry import get_model_info
|
from sglang.multimodal_gen.registry import get_model_info
|
||||||
@@ -522,12 +513,11 @@ class SamplingParams:
|
|||||||
help="Whether to return the decoded trajectory",
|
help="Whether to return the decoded trajectory",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--override-protected-fields",
|
"--no-override-protected-fields",
|
||||||
action="store_true",
|
action="store_true",
|
||||||
default=SamplingParams.override_protected_fields,
|
default=SamplingParams.no_override_protected_fields,
|
||||||
help=(
|
help=(
|
||||||
"If set, allow user params to override fields defined in subclasses "
|
"If set, disallow user params to override fields defined in subclasses."
|
||||||
"(protected by default)."
|
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
@@ -583,32 +573,21 @@ class SamplingParams:
|
|||||||
subclass_defined_fields = set(type(self).__annotations__.keys())
|
subclass_defined_fields = set(type(self).__annotations__.keys())
|
||||||
|
|
||||||
# global switch: if True, allow overriding protected fields
|
# global switch: if True, allow overriding protected fields
|
||||||
allow_override_protected = bool(
|
allow_override_protected = not user_params.no_override_protected_fields
|
||||||
user_params.override_protected_fields or self.override_protected_fields
|
|
||||||
)
|
|
||||||
|
|
||||||
# Compare against current instance to avoid constructing a default instance
|
|
||||||
default_params = SamplingParams()
|
|
||||||
|
|
||||||
for field in dataclasses.fields(user_params):
|
for field in dataclasses.fields(user_params):
|
||||||
field_name = field.name
|
field_name = field.name
|
||||||
user_value = getattr(user_params, field_name)
|
user_value = getattr(user_params, field_name)
|
||||||
default_value = getattr(default_params, field_name)
|
default_value = getattr(self, field_name)
|
||||||
|
|
||||||
# A field is considered user-modified if its value is different from
|
# A field is considered user-modified if its value is different from
|
||||||
# the default, with an exception for `output_file_name` which is
|
# the default
|
||||||
# auto-generated with a random component.
|
is_user_modified = user_value != default_value
|
||||||
is_user_modified = (
|
|
||||||
user_value != default_value
|
|
||||||
if field_name != "output_file_name"
|
|
||||||
else user_params.output_file_path is not None
|
|
||||||
)
|
|
||||||
if is_user_modified and (
|
if is_user_modified and (
|
||||||
allow_override_protected or field_name not in subclass_defined_fields
|
allow_override_protected or field_name not in subclass_defined_fields
|
||||||
):
|
):
|
||||||
if hasattr(self, field_name):
|
if hasattr(self, field_name):
|
||||||
setattr(self, field_name, user_value)
|
setattr(self, field_name, user_value)
|
||||||
|
|
||||||
self.height_not_provided = user_params.height_not_provided
|
self.height_not_provided = user_params.height_not_provided
|
||||||
self.width_not_provided = user_params.width_not_provided
|
self.width_not_provided = user_params.width_not_provided
|
||||||
self.__post_init__()
|
self.__post_init__()
|
||||||
|
|||||||
@@ -10,10 +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.sampling_params import (
|
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
|
||||||
SamplingParams,
|
|
||||||
generate_request_id,
|
|
||||||
)
|
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.cli.cli_types import CLISubcommand
|
from sglang.multimodal_gen.runtime.entrypoints.cli.cli_types import CLISubcommand
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.cli.utils import (
|
from sglang.multimodal_gen.runtime.entrypoints.cli.utils import (
|
||||||
RaiseNotImplementedAction,
|
RaiseNotImplementedAction,
|
||||||
@@ -89,8 +86,7 @@ def maybe_dump_performance(args: argparse.Namespace, server_args, prompt: str, r
|
|||||||
|
|
||||||
def generate_cmd(args: argparse.Namespace):
|
def generate_cmd(args: argparse.Namespace):
|
||||||
"""The entry point for the generate command."""
|
"""The entry point for the generate command."""
|
||||||
# FIXME(mick): do not hard code
|
args.request_id = "mocked_fake_id_for_offline_generate"
|
||||||
args.request_id = generate_request_id()
|
|
||||||
|
|
||||||
# Auto-enable stage logging if dump path is provided
|
# Auto-enable stage logging if dump path is provided
|
||||||
if args.perf_dump_path:
|
if args.perf_dump_path:
|
||||||
|
|||||||
@@ -297,6 +297,30 @@ class DenoisingStage(PipelineStage):
|
|||||||
|
|
||||||
return reserved_frames_mask_sp, z_sp
|
return reserved_frames_mask_sp, z_sp
|
||||||
|
|
||||||
|
def _handle_boundary_ratio(
|
||||||
|
self,
|
||||||
|
server_args,
|
||||||
|
batch,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
(Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
|
||||||
|
"""
|
||||||
|
boundary_ratio = server_args.pipeline_config.dit_config.boundary_ratio
|
||||||
|
if batch.boundary_ratio is not None:
|
||||||
|
logger.info(
|
||||||
|
"Overriding boundary ratio from %s to %s",
|
||||||
|
boundary_ratio,
|
||||||
|
batch.boundary_ratio,
|
||||||
|
)
|
||||||
|
boundary_ratio = batch.boundary_ratio
|
||||||
|
|
||||||
|
if boundary_ratio is not None:
|
||||||
|
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
|
||||||
|
else:
|
||||||
|
boundary_timestep = None
|
||||||
|
|
||||||
|
return boundary_timestep
|
||||||
|
|
||||||
def _prepare_denoising_loop(self, batch: Req, server_args: ServerArgs):
|
def _prepare_denoising_loop(self, batch: Req, server_args: ServerArgs):
|
||||||
"""
|
"""
|
||||||
Prepare all necessary invariant variables for the denoising loop.
|
Prepare all necessary invariant variables for the denoising loop.
|
||||||
@@ -362,20 +386,7 @@ class DenoisingStage(PipelineStage):
|
|||||||
assert neg_prompt_embeds is not None
|
assert neg_prompt_embeds is not None
|
||||||
# Removed Tensor truthiness assert to avoid GPU sync
|
# Removed Tensor truthiness assert to avoid GPU sync
|
||||||
|
|
||||||
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
|
boundary_timestep = self._handle_boundary_ratio(server_args, batch)
|
||||||
boundary_ratio = server_args.pipeline_config.dit_config.boundary_ratio
|
|
||||||
if batch.boundary_ratio is not None:
|
|
||||||
logger.info(
|
|
||||||
"Overriding boundary ratio from %s to %s",
|
|
||||||
boundary_ratio,
|
|
||||||
batch.boundary_ratio,
|
|
||||||
)
|
|
||||||
boundary_ratio = batch.boundary_ratio
|
|
||||||
|
|
||||||
if boundary_ratio is not None:
|
|
||||||
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
|
|
||||||
else:
|
|
||||||
boundary_timestep = None
|
|
||||||
|
|
||||||
# specifically for Wan2_2_TI2V_5B_Config, not applicable for FastWan2_2_TI2V_5B_Config
|
# specifically for Wan2_2_TI2V_5B_Config, not applicable for FastWan2_2_TI2V_5B_Config
|
||||||
should_preprocess_for_wan_ti2v = (
|
should_preprocess_for_wan_ti2v = (
|
||||||
|
|||||||
@@ -269,18 +269,19 @@ ONE_GPU_CASES_A: list[DiffusionTestCase] = [
|
|||||||
),
|
),
|
||||||
),
|
),
|
||||||
# === Text and Image to Image (TI2I) ===
|
# === Text and Image to Image (TI2I) ===
|
||||||
# TODO: Timeout with Torch2.9. Add back when it can pass CI
|
DiffusionTestCase(
|
||||||
# DiffusionTestCase(
|
"qwen_image_edit_ti2i",
|
||||||
# id="qwen_image_edit_ti2i",
|
DiffusionServerArgs(
|
||||||
# model_path="Qwen/Qwen-Image-Edit",
|
model_path="Qwen/Qwen-Image-Edit",
|
||||||
# modality="image",
|
modality="image",
|
||||||
# prompt=None, # not used for editing
|
warmup_text=0,
|
||||||
# output_size="1024x1536",
|
warmup_edit=1,
|
||||||
# warmup_text=0,
|
),
|
||||||
# warmup_edit=1,
|
DiffusionSamplingParams(
|
||||||
# edit_prompt="Convert 2D style to 3D style",
|
prompt="Convert 2D style to 3D style",
|
||||||
# image_path="https://github.com/lm-sys/lm-sys.github.io/releases/download/test/TI2I_Qwen_Image_Edit_Input.jpg",
|
image_path="https://github.com/lm-sys/lm-sys.github.io/releases/download/test/TI2I_Qwen_Image_Edit_Input.jpg",
|
||||||
# ),
|
),
|
||||||
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
ONE_GPU_CASES_B: list[DiffusionTestCase] = [
|
ONE_GPU_CASES_B: list[DiffusionTestCase] = [
|
||||||
|
|||||||
@@ -49,7 +49,7 @@ except Exception as e:
|
|||||||
|
|
||||||
def _get_status_message(run_id, current_case_id, thread_messages=None):
|
def _get_status_message(run_id, current_case_id, thread_messages=None):
|
||||||
date_str = datetime.now().strftime("%d/%m")
|
date_str = datetime.now().strftime("%d/%m")
|
||||||
base_header = f""""*🧵 for nightly test of {date_str}*
|
base_header = f"""🧵 for nightly test of {date_str}*
|
||||||
*Git Revision:* {get_git_commit_hash()}
|
*Git Revision:* {get_git_commit_hash()}
|
||||||
*GitHub Run ID:* {run_id}
|
*GitHub Run ID:* {run_id}
|
||||||
*Total Tasks:* {len(ALL_CASES)}
|
*Total Tasks:* {len(ALL_CASES)}
|
||||||
|
|||||||
Reference in New Issue
Block a user