[diffusion] refactor: unify model loading and offloading behavior (#15923)

This commit is contained in:
Mick
2025-12-27 16:18:24 +08:00
committed by GitHub
parent 171912a9e3
commit aa89c6a7e2
9 changed files with 97 additions and 38 deletions

View File

@@ -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", {})

View File

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

View File

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

View File

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

View File

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

View File

@@ -34,6 +34,7 @@ class EncodingStage(PipelineStage):
"""
def __init__(self, vae: ParallelTiledVAE) -> None:
super().__init__()
self.vae: ParallelTiledVAE = vae
@torch.no_grad()

View File

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

View File

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

View File

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