diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loader.py index 800cf3ea7..99e9feca6 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loader.py @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 35913db63..79d23549e 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -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, diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py index 39094e118..9144c81a6 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py @@ -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): diff --git a/python/sglang/multimodal_gen/runtime/platforms/cpu.py b/python/sglang/multimodal_gen/runtime/platforms/cpu.py index 5186d2489..c937c15a9 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/cpu.py +++ b/python/sglang/multimodal_gen/runtime/platforms/cpu.py @@ -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" diff --git a/python/sglang/multimodal_gen/runtime/platforms/cuda.py b/python/sglang/multimodal_gen/runtime/platforms/cuda.py index 136ea29b3..917d013dd 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/cuda.py +++ b/python/sglang/multimodal_gen/runtime/platforms/cuda.py @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/platforms/interface.py b/python/sglang/multimodal_gen/runtime/platforms/interface.py index 82cff8624..abdb55dc0 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/interface.py +++ b/python/sglang/multimodal_gen/runtime/platforms/interface.py @@ -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: """ diff --git a/python/sglang/multimodal_gen/runtime/platforms/mps.py b/python/sglang/multimodal_gen/runtime/platforms/mps.py index 2312ec059..2208d70ce 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/mps.py +++ b/python/sglang/multimodal_gen/runtime/platforms/mps.py @@ -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, diff --git a/python/sglang/multimodal_gen/runtime/platforms/rocm.py b/python/sglang/multimodal_gen/runtime/platforms/rocm.py index 02202fe22..7741cb0b9 100644 --- a/python/sglang/multimodal_gen/runtime/platforms/rocm.py +++ b/python/sglang/multimodal_gen/runtime/platforms/rocm.py @@ -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,