From aa89c6a7e2aa826bf94008b33b61c53206961615 Mon Sep 17 00:00:00 2001 From: Mick Date: Sat, 27 Dec 2025 16:18:24 +0800 Subject: [PATCH] [diffusion] refactor: unify model loading and offloading behavior (#15923) --- .../runtime/loader/component_loader.py | 21 ++++++--- .../runtime/managers/gpu_worker.py | 10 ++-- .../models/model_stages/qwen_image_layered.py | 1 + .../runtime/pipelines_core/stages/base.py | 12 +++++ .../runtime/pipelines_core/stages/decoding.py | 47 ++++++++++++------- .../runtime/pipelines_core/stages/encoding.py | 1 + .../pipelines_core/stages/image_encoding.py | 29 ++++++++---- .../stages/timestep_preparation.py | 1 + .../multimodal_gen/runtime/server_args.py | 13 ++++- 9 files changed, 97 insertions(+), 38 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loader.py index 324ee2b9a..800cf3ea7 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loader.py @@ -96,9 +96,11 @@ class ComponentLoader(ABC): def __init__(self, device=None) -> None: self.device = device - def should_offload(self, server_args, model_config: ModelConfig | None = None): - # offload by default - return True + def should_offload( + self, server_args: ServerArgs, model_config: ModelConfig | None = None + ): + # not offload by default + return False def target_device(self, should_offload): if should_offload: @@ -276,6 +278,8 @@ class TextEncoderLoader(ComponentLoader): def should_offload(self, server_args, model_config: ModelConfig | None = None): should_offload = server_args.text_encoder_cpu_offload + if not should_offload: + return False # _fsdp_shard_conditions is in arch_config, not directly on model_config arch_config = ( getattr(model_config, "arch_config", model_config) if model_config else None @@ -486,6 +490,8 @@ class TextEncoderLoader(ComponentLoader): class ImageEncoderLoader(TextEncoderLoader): def should_offload(self, server_args, model_config: ModelConfig | None = None): should_offload = server_args.image_encoder_cpu_offload + if not should_offload: + return False # _fsdp_shard_conditions is in arch_config, not directly on model_config arch_config = ( getattr(model_config, "arch_config", model_config) if model_config else None @@ -558,8 +564,10 @@ class TokenizerLoader(ComponentLoader): class VAELoader(ComponentLoader): """Loader for VAE.""" - def should_offload(self, server_args, cpu_offload_flag, model_config): - return True + def should_offload( + self, server_args: ServerArgs, model_config: ModelConfig | None = None + ): + return server_args.vae_cpu_offload def load_customized( self, component_model_path: str, server_args: ServerArgs, *args @@ -580,7 +588,8 @@ class VAELoader(ComponentLoader): # NOTE: some post init logics are only available after updated with config vae_config.post_init() - target_device = self.target_device(server_args.vae_cpu_offload) + should_offload = self.should_offload(server_args) + target_device = self.target_device(should_offload) # Check for auto_map first (custom VAE classes) auto_map = config.get("auto_map", {}) diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index a656c03ec..35913db63 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -108,11 +108,11 @@ class GPUWorker: peak_memory_bytes = torch.cuda.max_memory_allocated() output_batch.peak_memory_mb = peak_memory_bytes / (1024**2) - if output_batch.timings: - duration_ms = (time.monotonic() - start_time) * 1000 - output_batch.timings.total_duration_ms = duration_ms - if StageProfiler.metrics_enabled(): - PerformanceLogger.log_request_summary(timings=output_batch.timings) + duration_ms = (time.monotonic() - start_time) * 1000 + output_batch.timings.total_duration_ms = duration_ms + + if StageProfiler.metrics_enabled(): + PerformanceLogger.log_request_summary(timings=output_batch.timings) except Exception as e: logger.error( f"Error executing request {req.request_id}: {e}", exc_info=True diff --git a/python/sglang/multimodal_gen/runtime/models/model_stages/qwen_image_layered.py b/python/sglang/multimodal_gen/runtime/models/model_stages/qwen_image_layered.py index 999aebd58..4d64759c5 100644 --- a/python/sglang/multimodal_gen/runtime/models/model_stages/qwen_image_layered.py +++ b/python/sglang/multimodal_gen/runtime/models/model_stages/qwen_image_layered.py @@ -113,6 +113,7 @@ class QwenImageLayeredBeforeDenoisingStage(PipelineStage): def __init__( self, vae, tokenizer, processor, transformer, scheduler, model_path ) -> None: + super().__init__() self.vae = vae.to(torch.bfloat16) from transformers import Qwen2_5_VLForConditionalGeneration diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py index 1adecbfe0..6b954131f 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py @@ -96,6 +96,18 @@ class PipelineStage(ABC): def maybe_free_model_hooks(self): pass + def load_model(self): + """ + Load the model for the stage. + """ + pass + + def offload_model(self): + """ + Offload the model for the stage. + """ + pass + # execute on all ranks by default @property def parallelism_type(self) -> StageParallelismType: diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py index 9db473cf6..fa2480be3 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py @@ -57,6 +57,7 @@ class DecodingStage(PipelineStage): """ def __init__(self, vae, pipeline=None) -> None: + super().__init__() self.vae: ParallelTiledVAE = vae self.pipeline = weakref.ref(pipeline) if pipeline else None @@ -153,6 +154,32 @@ class DecodingStage(PipelineStage): image = (image / 2 + 0.5).clamp(0, 1) return image + def load_model(self): + # load vae if not already loaded (used for memory constrained devices) + pipeline = self.pipeline() if self.pipeline else None + if not self.server_args.model_loaded["vae"]: + loader = VAELoader() + self.vae = loader.load( + self.server_args.model_paths["vae"], self.server_args + ) + if pipeline: + pipeline.add_module("vae", self.vae) + self.server_args.model_loaded["vae"] = True + + def offload_model(self): + # Offload models if needed + self.maybe_free_model_hooks() + + if self.server_args.vae_cpu_offload: + self.vae.to("cpu", non_blocking=True) + + if torch.backends.mps.is_available(): + del self.vae + pipeline = self.pipeline() if self.pipeline else None + if pipeline is not None and "vae" in pipeline.modules: + del pipeline.modules["vae"] + self.server_args.model_loaded["vae"] = False + @torch.no_grad() def forward( self, @@ -184,13 +211,7 @@ class DecodingStage(PipelineStage): - trajectory_decoded (if requested): List of decoded frames per timestep """ # load vae if not already loaded (used for memory constrained devices) - pipeline = self.pipeline() if self.pipeline else None - if not server_args.model_loaded["vae"]: - loader = VAELoader() - self.vae = loader.load(server_args.model_paths["vae"], server_args) - if pipeline: - pipeline.add_module("vae", self.vae) - server_args.model_loaded["vae"] = True + self.load_model() if server_args.output_type == "latent": frames = batch.latents @@ -231,16 +252,6 @@ class DecodingStage(PipelineStage): timings=batch.timings, ) - # Offload models if needed - self.maybe_free_model_hooks() - - if server_args.vae_cpu_offload: - self.vae.to("cpu") - - if torch.backends.mps.is_available(): - del self.vae - if pipeline is not None and "vae" in pipeline.modules: - del pipeline.modules["vae"] - server_args.model_loaded["vae"] = False + self.offload_model() return output_batch diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/encoding.py index 9e30a8cb7..51c6993f1 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/encoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/encoding.py @@ -34,6 +34,7 @@ class EncodingStage(PipelineStage): """ def __init__(self, vae: ParallelTiledVAE) -> None: + super().__init__() self.vae: ParallelTiledVAE = vae @torch.no_grad() diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py index 4a6eb1ebc..6b22afc49 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py @@ -63,6 +63,15 @@ class ImageEncodingStage(PipelineStage): self.image_encoder = image_encoder self.text_encoder = text_encoder + def load_model(self): + if self.server_args.image_encoder_cpu_offload: + device = get_local_torch_device() + self.move_to_device(device) + + def offload_model(self): + if self.server_args.image_encoder_cpu_offload: + self.move_to_device("cpu") + def move_to_device(self, device): fields = [ "image_processor", @@ -98,8 +107,8 @@ class ImageEncodingStage(PipelineStage): if batch.condition_image is None: return batch cuda_device = get_local_torch_device() - self.move_to_device(cuda_device) + self.load_model() image = batch.condition_image image_processor_kwargs = ( @@ -130,7 +139,7 @@ class ImageEncodingStage(PipelineStage): neg_image_inputs = self.image_processor( images=image, return_tensors="pt", **neg_image_processor_kwargs - ).to(get_local_torch_device()) + ).to(cuda_device) with set_forward_context(current_timestep=0, attn_metadata=None): outputs = self.text_encoder( @@ -155,7 +164,7 @@ class ImageEncodingStage(PipelineStage): self.encoding_qwen_image_edit(neg_outputs, neg_image_inputs) ) - self.move_to_device("cpu") + self.offload_model() return batch @@ -188,6 +197,13 @@ class ImageVAEEncodingStage(PipelineStage): super().__init__() self.vae: ParallelTiledVAE = vae + def load_model(self): + self.vae = self.vae.to(get_local_torch_device()) + + def offload_model(self): + if self.server_args.vae_cpu_offload: + self.vae = self.vae.to("cpu") + def forward( self, batch: Req, @@ -207,10 +223,9 @@ class ImageVAEEncodingStage(PipelineStage): if batch.condition_image is None: return batch + self.load_model() num_frames = batch.num_frames - self.vae = self.vae.to(get_local_torch_device()) - images = ( batch.vae_image if batch.vae_image is not None else batch.condition_image ) @@ -305,10 +320,8 @@ class ImageVAEEncodingStage(PipelineStage): all_image_latents.append(image_latent) batch.image_latent = torch.cat(all_image_latents, dim=1) - self.maybe_free_model_hooks() - - self.vae.to("cpu") + self.offload_model() return batch def retrieve_latents( diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/timestep_preparation.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/timestep_preparation.py index fea00c766..13bc6a066 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/timestep_preparation.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/timestep_preparation.py @@ -43,6 +43,7 @@ class TimestepPreparationStage(PipelineStage): Callable[[Req, ServerArgs], Tuple[str, Any]] ] = [], ) -> None: + super().__init__() self.scheduler = scheduler self.prepare_extra_set_timesteps_kwargs = prepare_extra_set_timesteps_kwargs diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index 4f9f92605..dd2743930 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -228,11 +228,11 @@ class ServerArgs: # CPU offload parameters dit_cpu_offload: bool = True - use_fsdp_inference: bool = False dit_layerwise_offload: bool = False text_encoder_cpu_offload: bool = True image_encoder_cpu_offload: bool = True vae_cpu_offload: bool = True + use_fsdp_inference: bool = False pin_cpu_memory: bool = True # STA (Sliding Tile Attention) parameters @@ -302,8 +302,19 @@ class ServerArgs: """ return self.host is None or self.port is None + 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 + 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 + + self.adjust_offload() + if self.attention_backend in ["fa3", "fa4"]: self.attention_backend = "fa"