Spaces:
Build error
Build error
| # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. | |
| # SPDX-License-Identifier: OpenMDW-1.1 | |
| import gc | |
| import math | |
| import os | |
| from functools import wraps | |
| import pynvml | |
| from loguru import logger as logging | |
| def get_gpu_architecture(): | |
| """ | |
| Retrieves the GPU architecture of the available GPUs. | |
| Returns: | |
| str: The GPU architecture, which can be "H100", "A100", or "Other". | |
| """ | |
| try: | |
| pynvml.nvmlInit() | |
| device_count = pynvml.nvmlDeviceGetCount() | |
| for i in range(device_count): | |
| handle = pynvml.nvmlDeviceGetHandleByIndex(i) | |
| model_name = pynvml.nvmlDeviceGetName(handle) | |
| if isinstance(model_name, bytes): | |
| model_name = model_name.decode("utf-8") | |
| print(f"GPU {i}: Model: {model_name}") | |
| # Check for specific models like H100 or A100 | |
| if "H100" in model_name or "H200" in model_name: | |
| return "H100" | |
| elif "A100" in model_name: | |
| return "A100" | |
| elif "L40S" in model_name: | |
| return "L40S" | |
| elif "B200" in model_name: | |
| return "B200" | |
| except pynvml.NVMLError as error: | |
| print(f"Failed to get GPU info: {error}") | |
| finally: | |
| pynvml.nvmlShutdown() | |
| # return "Other" incase of non hopper/ampere or error | |
| return "Other" | |
| class GPUArchitectureNotSupported(Exception): | |
| """ | |
| Custom exception raised when the expected GPU architecture is not supported. | |
| """ | |
| pass | |
| def print_gpu_mem(str=None): | |
| try: | |
| pynvml.nvmlInit() | |
| meminfo = pynvml.nvmlDeviceGetMemoryInfo(pynvml.nvmlDeviceGetHandleByIndex(0)) | |
| logging.info( | |
| f"{str}: {meminfo.used / 1024 / 1024}/{meminfo.total / 1024 / 1024}MiB used ({meminfo.free / 1024 / 1024}MiB free)" | |
| ) | |
| except pynvml.NVMLError as error: | |
| print(f"Failed to get GPU memory info: {error}") | |
| def force_gc(): | |
| print_gpu_mem() | |
| print("gc()") | |
| gc.collect() | |
| print_gpu_mem() | |
| print("empty cuda cache") | |
| # print(torch.cuda.memory_summary()) | |
| print_gpu_mem() | |
| def gpu0_has_80gb_or_less(): | |
| try: | |
| pynvml.nvmlInit() | |
| meminfo = pynvml.nvmlDeviceGetMemoryInfo(pynvml.nvmlDeviceGetHandleByIndex(0)) | |
| return meminfo.total / 1024 / 1024 / 1024 <= 80 | |
| except pynvml.NVMLError as error: | |
| print(f"Failed to get GPU memory info: {error}") | |
| class Device: | |
| _nvml_affinity_elements = math.ceil(os.cpu_count() / 64) # type: ignore | |
| def __init__(self, device_idx: int): | |
| super().__init__() | |
| self.handle = pynvml.nvmlDeviceGetHandleByIndex(device_idx) | |
| def get_name(self) -> str: | |
| return pynvml.nvmlDeviceGetName(self.handle) | |
| def get_cpu_affinity(self) -> list[int]: | |
| affinity_string = "" | |
| for j in pynvml.nvmlDeviceGetCpuAffinity(self.handle, Device._nvml_affinity_elements): | |
| # assume nvml returns list of 64 bit ints | |
| affinity_string = "{:064b}".format(j) + affinity_string | |
| affinity_list = [int(x) for x in affinity_string] | |
| affinity_list.reverse() # so core 0 is in 0th element of list | |
| return [i for i, e in enumerate(affinity_list) if e != 0] | |
| def with_torch_device(device): | |
| """ | |
| Decorator factory that wraps a function to execute within a specific torch device context. | |
| This decorator ensures that all tensor allocations and operations within the decorated | |
| function use the specified device by default. | |
| Args: | |
| device: The torch device to use (e.g., 'cuda', 'cuda:0', 'cpu', or torch.device object). | |
| Returns: | |
| A decorator function that wraps the target function with the specified device context. | |
| Example: | |
| @with_torch_device('cuda:0') | |
| def create_tensors(): | |
| x = torch.randn(10, 10) # Will be created on cuda:0 | |
| return x | |
| """ | |
| import torch | |
| def decorator(fn): | |
| def wrapper(*args, **kwargs): | |
| with torch.device(device): | |
| return fn(*args, **kwargs) | |
| return wrapper | |
| return decorator | |