dwehr's picture
Migrate action viewer to local Cosmos generation
9f818c5
Raw
History Blame Contribute Delete
4.16 kB
# 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):
@wraps(fn)
def wrapper(*args, **kwargs):
with torch.device(device):
return fn(*args, **kwargs)
return wrapper
return decorator