[diffusion] logging: log available gpu mem while loading and generating (#15936)

This commit is contained in:
Mick
2025-12-28 00:34:58 +08:00
committed by GitHub
parent 41addd2e08
commit 39d56196a0
8 changed files with 242 additions and 36 deletions

View File

@@ -7,7 +7,6 @@ import glob
import importlib.util
import json
import os
import time
import traceback
from abc import ABC
from collections.abc import Generator, Iterable
@@ -17,6 +16,7 @@ from typing import Any, cast
import torch
import torch.distributed as dist
import torch.nn as nn
from diffusers import AutoModel
from safetensors.torch import load_file as safetensors_load_file
from torch.distributed import init_device_mesh
from transformers import AutoImageProcessor, AutoProcessor, AutoTokenizer
@@ -90,6 +90,26 @@ def _list_safetensors_files(model_path: str) -> list[str]:
return sorted(glob.glob(os.path.join(str(model_path), "*.safetensors")))
def get_memory_usage_of_component(module) -> float | None:
"""
returned value is in GB, rounded to 2 decimal digits
"""
if not isinstance(module, nn.Module):
return None
BYTES_PER_GB = 1024**3
if hasattr(module, "get_memory_footprint"):
usage = module.get_memory_footprint() / BYTES_PER_GB
else:
# manually
param_size = sum(p.numel() * p.element_size() for p in module.parameters())
buffer_size = sum(b.numel() * b.element_size() for b in module.buffers())
total_size_bytes = param_size + buffer_size
usage = total_size_bytes / (1024**3)
return round(usage, 2)
class ComponentLoader(ABC):
"""Base class for loading a specific type of model component."""
@@ -118,7 +138,7 @@ class ComponentLoader(ABC):
server_args: ServerArgs,
module_name: str,
transformers_or_diffusers: str,
):
) -> tuple[AutoModel, float]:
"""
Template method that standardizes logging around the core load implementation.
The priority of loading method is:
@@ -127,7 +147,13 @@ class ComponentLoader(ABC):
If all of the above methods failed, an error will be thrown
"""
logger.info("Loading %s from %s", module_name, component_model_path)
gpu_mem_before_loading = current_platform.get_available_gpu_memory()
logger.info(
"Loading %s from %s. avail mem: %.2f GB",
module_name,
component_model_path,
gpu_mem_before_loading,
)
try:
component = self.load_customized(
component_model_path, server_args, module_name
@@ -159,20 +185,27 @@ class ComponentLoader(ABC):
if component is None:
logger.warning("Loaded %s returned None", module_name)
consumed = 0.0
else:
current_gpu_mem = current_platform.get_available_gpu_memory()
consumed = get_memory_usage_of_component(component)
if consumed is None or consumed == 0.0:
consumed = gpu_mem_before_loading - current_gpu_mem
logger.info(
f"Loaded %s: %s from: {source}",
f"Loaded %s: %s from {source}. avail mem: %.2f GB, %.2f GB consumed",
module_name,
component.__class__.__name__,
current_gpu_mem,
consumed,
)
return component
return component, consumed
def load_native(
self,
component_model_path: str,
server_args: ServerArgs,
transformers_or_diffusers: str,
):
) -> AutoModel:
"""
Load the component using the native library (transformers/diffusers).
"""
@@ -273,9 +306,6 @@ class TextEncoderLoader(ComponentLoader):
allow_patterns_overrides: list[str] | None = None
"""If defined, weights will load exclusively using these patterns."""
counter_before_loading_weights: float = 0.0
counter_after_loading_weights: float = 0.0
def should_offload(self, server_args, model_config: ModelConfig | None = None):
should_offload = server_args.text_encoder_cpu_offload
if not should_offload:
@@ -355,8 +385,6 @@ class TextEncoderLoader(ComponentLoader):
else:
weights_iterator = pt_weights_iterator(hf_weights_files, to_cpu=to_cpu)
if self.counter_before_loading_weights == 0.0:
self.counter_before_loading_weights = time.perf_counter()
# apply the prefix.
return ((source.prefix + name, tensor) for (name, tensor) in weights_iterator)
@@ -443,12 +471,6 @@ class TextEncoderLoader(ComponentLoader):
loaded_weights = model.load_weights(
self._get_all_weights(model, model_path, to_cpu=should_offload)
)
self.counter_after_loading_weights = time.perf_counter()
logger.info(
"Loading weights took %.2f seconds",
self.counter_after_loading_weights
- self.counter_before_loading_weights,
)
# Explicitly move model to target device after loading weights
model = model.to(local_torch_device)

View File

@@ -17,8 +17,13 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_cfg_group,
get_tp_group,
)
from sglang.multimodal_gen.runtime.pipelines_core import Req, build_pipeline
from sglang.multimodal_gen.runtime.pipelines_core import (
ComposedPipelineBase,
Req,
build_pipeline,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.server_args import PortArgs, ServerArgs
from sglang.multimodal_gen.runtime.utils.common import set_cuda_arch
from sglang.multimodal_gen.runtime.utils.logging_utils import (
@@ -51,7 +56,7 @@ class GPUWorker:
self.master_port = master_port
# FIXME: should we use tcp as distribute init method?
self.server_args = server_args
self.pipeline = None
self.pipeline: ComposedPipelineBase = None
self.init_device_and_model()
self.sp_group = get_sp_group()
@@ -107,6 +112,19 @@ class GPUWorker:
if self.rank == 0:
peak_memory_bytes = torch.cuda.max_memory_allocated()
output_batch.peak_memory_mb = peak_memory_bytes / (1024**2)
peak_memory_gb = peak_memory_bytes / (1024**3)
remaining_gpu_mem_gb = (
current_platform.get_device_total_memory() / (1024**3)
- peak_memory_gb
)
can_stay_resident = self.get_can_stay_resident_components(
remaining_gpu_mem_gb
)
logger.info(
f"Peak GPU memory: {peak_memory_gb:.2f} GB, "
f"Remaining GPU memory at peak: {remaining_gpu_mem_gb:.2f} GB. "
f"Components that can stay resident: {can_stay_resident}"
)
duration_ms = (time.monotonic() - start_time) * 1000
output_batch.timings.total_duration_ms = duration_ms
@@ -127,6 +145,40 @@ class GPUWorker:
finally:
return output_batch
def get_can_stay_resident_components(
self, remaining_gpu_mem_gb: float
) -> List[str]:
"""
Calculate which components can stay resident on GPU without being offloaded.
"""
can_stay_resident = []
if not self.pipeline:
return can_stay_resident
# Map memory_usage keys to server_args offload flags
# If the flag is False, the component is ALREADY resident, so we don't suggest it.
# If the flag is True, it is currently offloaded, so it's a candidate to "stay resident".
offload_flags = {
"transformer": self.server_args.dit_cpu_offload
or self.server_args.dit_layerwise_offload,
"vae": self.server_args.vae_cpu_offload,
"text_encoder": self.server_args.text_encoder_cpu_offload,
"text_encoder_2": self.server_args.text_encoder_cpu_offload,
"image_encoder": self.server_args.image_encoder_cpu_offload,
}
for name, usage in self.pipeline.memory_usages.items():
# Only consider components that are currently configured to be offloaded
is_offload_configured = offload_flags.get(name, False)
if not is_offload_configured:
continue
if usage <= remaining_gpu_mem_gb:
can_stay_resident.append(name)
remaining_gpu_mem_gb -= usage
return can_stay_resident
def set_lora(
self,
lora_nickname: str,

View File

@@ -85,11 +85,9 @@ class ComposedPipelineBase(ABC):
if self._required_config_modules is None:
raise NotImplementedError("Subclass must set _required_config_modules")
# temp disable for duplicate initialing tp
# maybe_init_distributed_environment_and_model_parallel(
# server_args.tp_size, server_args.sp_size
# )
# [module_name, gpu memory usage]
self.memory_usages: dict[str, float] = {}
# Load modules directly in initialization
logger.info("Loading pipeline modules...")
self.modules = self.load_modules(server_args, loaded_modules)
@@ -313,13 +311,15 @@ class ComposedPipelineBase(ABC):
)
else:
component_model_path = os.path.join(self.model_path, load_module_name)
module = PipelineComponentLoader.load_module(
module, memory_usage = PipelineComponentLoader.load_module(
module_name=load_module_name,
component_model_path=component_model_path,
transformers_or_diffusers=transformers_or_diffusers,
server_args=server_args,
)
self.memory_usages[load_module_name] = memory_usage
if module_name in components:
logger.warning("Overwriting module %s", module_name)
components[module_name] = module
@@ -331,6 +331,8 @@ class ComposedPipelineBase(ABC):
f"Required module key: {module_name} value: {components.get(module_name)} was not found in loaded modules {components.keys()}"
)
logger.debug("Memory usage of loaded modules: %s", self.memory_usages)
return components
def add_stage(self, stage_name: str, stage: PipelineStage):

View File

@@ -4,7 +4,10 @@
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/platforms/cpu.py
import platform
from functools import lru_cache
from typing import Any
import psutil
import torch
from sglang.multimodal_gen.runtime.platforms.interface import (
@@ -40,10 +43,10 @@ class CpuPlatform(Platform):
return platform.machine()
@classmethod
@lru_cache(maxsize=1)
def get_device_total_memory(cls, device_id: int = 0) -> int:
# This is a rough estimate for CPU memory
# In practice, you might want to use psutil or similar
return 0
return psutil.virtual_memory().total
@classmethod
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
@@ -56,6 +59,30 @@ class CpuPlatform(Platform):
# For CPU, we can't easily get memory usage without additional libraries
return 0.0
@classmethod
def get_available_gpu_memory(
cls,
device_id: int = 0,
distributed: bool = False,
empty_cache: bool = True,
cpu_group: Any = None,
) -> float:
total_free_memory = psutil.virtual_memory().available
# For simplicity, we assume 1 NUMA node for now in this platform abstraction
# as get_cpu_ids_by_node is not available in multimodal_gen.runtime.utils
n_numa_node = 1
free_memory = total_free_memory / n_numa_node
if distributed:
import torch.distributed as dist
tensor = torch.tensor(free_memory, dtype=torch.float32)
dist.all_reduce(tensor, op=dist.ReduceOp.MIN, group=cpu_group)
free_memory = float(tensor.item())
return free_memory / (1 << 30)
@classmethod
def get_device_communicator_cls(cls) -> str:
return "sglang.multimodal_gen.runtime.distributed.device_communicators.cpu_communicator.CpuCommunicator"

View File

@@ -5,12 +5,12 @@
"""Code inside this file can safely assume cuda platform, e.g. importing
pynvml. However, it should not initialize cuda context.
"""
import os
from collections.abc import Callable
from functools import lru_cache, wraps
from typing import TypeVar
from typing import Any, TypeVar
import psutil
import torch
from typing_extensions import ParamSpec
@@ -82,6 +82,7 @@ class CudaPlatformBase(Platform):
raise NotImplementedError
@classmethod
@lru_cache(maxsize=1)
def get_device_total_memory(cls, device_id: int = 0) -> int:
raise NotImplementedError
@@ -111,6 +112,38 @@ class CudaPlatformBase(Platform):
torch.cuda.reset_peak_memory_stats(device)
return float(torch.cuda.max_memory_allocated(device))
@classmethod
def get_available_gpu_memory(
cls,
device_id: int = 0,
distributed: bool = False,
empty_cache: bool = True,
cpu_group: Any = None,
) -> float:
if empty_cache:
torch.cuda.empty_cache()
# Orin, Thor, Spark
# SM 8.7 is Orin, 11.0 is Thor, 12.1 is Spark
SHARED_SYSMEM_DEVICE_MEM_SMS = (87, 110, 121)
capability = cls.get_device_capability(device_id)
sm = capability.to_int() if capability else 0
if sm in SHARED_SYSMEM_DEVICE_MEM_SMS:
free_gpu_memory = psutil.virtual_memory().available
else:
free_gpu_memory, _ = torch.cuda.mem_get_info(device_id)
if distributed:
import torch.distributed as dist
tensor = torch.tensor(free_gpu_memory, dtype=torch.float32, device="cuda")
dist.all_reduce(tensor, op=dist.ReduceOp.MIN, group=cpu_group)
free_gpu_memory = float(tensor.item())
return free_gpu_memory / (1 << 30)
@classmethod
def get_attn_backend_cls_str(
cls,
@@ -409,6 +442,7 @@ class NonNvmlCudaPlatform(CudaPlatformBase):
return str(torch.cuda.get_device_name(device_id))
@classmethod
@lru_cache(maxsize=1)
def get_device_total_memory(cls, device_id: int = 0) -> int:
device_props = torch.cuda.get_device_properties(device_id)
return int(device_props.total_memory)

View File

@@ -7,7 +7,7 @@ from __future__ import annotations
import enum
import random
from functools import lru_cache
from typing import TYPE_CHECKING, NamedTuple
from typing import TYPE_CHECKING, Any, NamedTuple
import numpy as np
import torch
@@ -216,6 +216,7 @@ class Platform:
raise NotImplementedError
@classmethod
@lru_cache(maxsize=1)
def get_device_total_memory(cls, device_id: int = 0) -> int:
"""Get the total memory of a device in bytes."""
raise NotImplementedError
@@ -307,6 +308,19 @@ class Platform:
"""
raise NotImplementedError
@classmethod
def get_available_gpu_memory(
cls,
device_id: int = 0,
distributed: bool = False,
empty_cache: bool = True,
cpu_group: Any = None,
) -> float:
"""
Return the available memory in GiB.
"""
raise NotImplementedError
@classmethod
def get_device_communicator_cls(cls) -> str:
"""

View File

@@ -1,17 +1,21 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
from functools import lru_cache
from typing import Any
# SPDX-License-Identifier: Apache-2.0
import psutil
import torch
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
from sglang.multimodal_gen.runtime.platforms.interface import (
DeviceCapability,
from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum,
Platform,
PlatformEnum,
)
from sglang.multimodal_gen.runtime.platforms.interface import DeviceCapability
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
# SPDX-License-Identifier: Apache-2.0
logger = init_logger(__name__)
@@ -35,8 +39,10 @@ class MpsPlatform(Platform):
raise NotImplementedError
@classmethod
@lru_cache(maxsize=1)
def get_device_total_memory(cls, device_id: int = 0) -> int:
raise NotImplementedError
return psutil.virtual_memory().total
@classmethod
def is_async_output_supported(cls, enforce_eager: bool | None) -> bool:
@@ -55,6 +61,30 @@ class MpsPlatform(Platform):
) -> float:
return 0.0
@classmethod
def get_available_gpu_memory(
cls,
device_id: int = 0,
distributed: bool = False,
empty_cache: bool = True,
cpu_group: Any = None,
) -> float:
if empty_cache:
torch.mps.empty_cache()
# For MPS, available memory is essentially the system available memory
free_memory = psutil.virtual_memory().available
if distributed:
import torch.distributed as dist
tensor = torch.tensor(free_memory, dtype=torch.float32)
dist.all_reduce(tensor, op=dist.ReduceOp.MIN, group=cpu_group)
free_memory = float(tensor.item())
return free_memory / (1 << 30)
@classmethod
def get_attn_backend_cls_str(
cls,

View File

@@ -6,6 +6,8 @@
This file is a platform abstraction for ROCm GPUs,
adjusted to match the structure and interface of `cuda.py`.
"""
from functools import lru_cache
from typing import Any
import torch
@@ -39,6 +41,7 @@ class RocmPlatform(Platform):
return str(torch.cuda.get_device_name(device_id))
@classmethod
@lru_cache(maxsize=1)
def get_device_total_memory(cls, device_id: int = 0) -> int:
return torch.cuda.get_device_properties(device_id).total_memory
@@ -61,6 +64,28 @@ class RocmPlatform(Platform):
torch.cuda.reset_peak_memory_stats(device)
return float(torch.cuda.max_memory_allocated(device))
@classmethod
def get_available_gpu_memory(
cls,
device_id: int = 0,
distributed: bool = False,
empty_cache: bool = True,
cpu_group: Any = None,
) -> float:
if empty_cache:
torch.cuda.empty_cache()
free_gpu_memory, _ = torch.cuda.mem_get_info(device_id)
if distributed:
import torch.distributed as dist
tensor = torch.tensor(free_gpu_memory, dtype=torch.float32, device="cuda")
dist.all_reduce(tensor, op=dist.ReduceOp.MIN, group=cpu_group)
free_gpu_memory = float(tensor.item())
return free_gpu_memory / (1 << 30)
@classmethod
def get_attn_backend_cls_str(
cls,