From 9f8e23071a3f1ce315a9224496d64a0351111de5 Mon Sep 17 00:00:00 2001 From: Mick Date: Sun, 28 Dec 2025 20:45:33 +0800 Subject: [PATCH] [diffusion] chore: fix default offload setting for image generation model (#15928) --- .../configs/sample/sampling_params.py | 2 +- .../multimodal_gen/runtime/server_args.py | 43 +++++++++++-------- .../sglang/multimodal_gen/test/slack_utils.py | 3 +- 3 files changed, 28 insertions(+), 20 deletions(-) diff --git a/python/sglang/multimodal_gen/configs/sample/sampling_params.py b/python/sglang/multimodal_gen/configs/sample/sampling_params.py index 7afa6fc77..c1a0eec90 100644 --- a/python/sglang/multimodal_gen/configs/sample/sampling_params.py +++ b/python/sglang/multimodal_gen/configs/sample/sampling_params.py @@ -285,7 +285,7 @@ class SamplingParams: if pipeline_config.task_type.is_image_gen(): # settle num_frames if not server_args.pipeline_config.allow_set_num_frames(): - logger.debug(f"num_frames set to 1 for image generation model") + logger.debug(f"Setting `num_frames` to 1 for image generation model") self.num_frames = 1 elif self.adjust_frames: diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index 4db630bab..9161dff02 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -34,8 +34,6 @@ from sglang.multimodal_gen.utils import FlexibleArgumentParser, StoreBoolean logger = init_logger(__name__) -ZMQ_TCP_PORT_DELTA = 233 - def _is_torch_tensor(obj: Any) -> tuple[bool, Any]: """Return (is_tensor, torch_module_or_None) without importing torch at module import time.""" @@ -193,11 +191,11 @@ class ServerArgs: lora_target_modules: list[str] | None = None # CPU offload parameters - dit_cpu_offload: bool = True - dit_layerwise_offload: bool = False - text_encoder_cpu_offload: bool = True - image_encoder_cpu_offload: bool = True - vae_cpu_offload: bool = True + dit_cpu_offload: bool | None = None + dit_layerwise_offload: bool | None = None + text_encoder_cpu_offload: bool | None = None + image_encoder_cpu_offload: bool | None = None + vae_cpu_offload: bool | None = None use_fsdp_inference: bool = False pin_cpu_memory: bool = True @@ -266,20 +264,33 @@ class ServerArgs: def adjust_offload(self): if self.pipeline_config.task_type.is_image_gen(): - logger.info("Turn off all offload for image generation model") - self.dit_cpu_offload = False + logger.info( + "Disabling some offloading (except dit, text_encoder) for image generation model" + ) + if self.dit_cpu_offload is None: + self.dit_cpu_offload = True + if self.text_encoder_cpu_offload is None: + self.text_encoder_cpu_offload = True + if self.image_encoder_cpu_offload is None: + self.image_encoder_cpu_offload = False + if self.vae_cpu_offload is None: + self.vae_cpu_offload = False + else: + self.dit_cpu_offload = True self.text_encoder_cpu_offload = True self.image_encoder_cpu_offload = True self.vae_cpu_offload = True def __post_init__(self): - # Add randomization to avoid race condition when multiple servers start simultaneously + # configure logger before use + configure_logger(server_args=self) self.adjust_offload() if self.attention_backend in ["fa3", "fa4"]: self.attention_backend = "fa" + # Add randomization to avoid race condition when multiple servers start simultaneously initial_scheduler_port = self.scheduler_port + random.randint(0, 100) self.scheduler_port = self.settle_port(initial_scheduler_port) # TODO: remove hard code @@ -297,8 +308,6 @@ class ServerArgs: raise self.check_server_args() - configure_logger(server_args=self) - # log clean server_args try: safe_args = _sanitize_for_logging(self, key_hint="server_args") @@ -733,13 +742,13 @@ class ServerArgs: if self.ulysses_degree is None: self.ulysses_degree = 1 - logger.info( + logger.debug( f"Ulysses degree not set, using default value {self.ulysses_degree}" ) if self.ring_degree is None: self.ring_degree = 1 - logger.info(f"Ring degree not set, using default value {self.ring_degree}") + logger.debug(f"Ring degree not set, using default value {self.ring_degree}") if self.ring_degree > 1: if self.attention_backend is not None and self.attention_backend != "fa": @@ -768,7 +777,7 @@ class ServerArgs: assert self.dp_size >= 1, "--dp-size must be natural number" # NOTE: disable temporarily # self.dp_degree = self.num_gpus // self.dp_size - logger.info(f"Setting dp_degree to: {self.dp_degree}") + logger.debug(f"Setting dp_degree to: {self.dp_degree}") if self.dp_size > 1: raise ValueError("DP is not yet supported") @@ -957,7 +966,7 @@ def set_global_server_args(server_args: ServerArgs): _global_server_args = server_args -def get_current_server_args() -> ServerArgs: +def get_current_server_args() -> ServerArgs | None: if _current_server_args is None: # in ci, usually when we test custom ops/modules directly, # we don't set the sgl_diffusion config. In that case, we set a default @@ -967,7 +976,7 @@ def get_current_server_args() -> ServerArgs: return _current_server_args -def get_global_server_args() -> ServerArgs: +def get_global_server_args() -> ServerArgs | None: if _global_server_args is None: # in ci, usually when we test custom ops/modules directly, # we don't set the sgl_diffusion config. In that case, we set a default diff --git a/python/sglang/multimodal_gen/test/slack_utils.py b/python/sglang/multimodal_gen/test/slack_utils.py index 3dafbdafc..4f4c2a6b7 100644 --- a/python/sglang/multimodal_gen/test/slack_utils.py +++ b/python/sglang/multimodal_gen/test/slack_utils.py @@ -50,11 +50,10 @@ except Exception as e: def _get_status_message(run_id, current_case_id, thread_messages=None): 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()} *GitHub Run ID:* {run_id} *Total Tasks:* {len(ALL_CASES)} - """ if not ALL_CASES: