[diffusion] chore: fix default offload setting for image generation model (#15928)

This commit is contained in:
Mick
2025-12-28 20:45:33 +08:00
committed by GitHub
parent d90f9bfc4e
commit 9f8e23071a
3 changed files with 28 additions and 20 deletions

View File

@@ -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:

View File

@@ -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

View File

@@ -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: