You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch 0.3.1下如何基于显存自动选择可用GPU

解决PyTorch 0.3.1中GPU显存查询与设备选择问题

你提到的代码在PyTorch 0.3.1里跑不通,核心原因就是这个版本里根本没有getMemoryUsage()函数。不过别担心,咱们用PyTorch自带的显存统计函数就能搞定需求——下面是完整的实现方案:

完整代码实现

import torch

def select_gpu_by_memory(threshold_memory):
    # 获取可用GPU数量
    num_gpus = torch.cuda.device_count()
    if num_gpus == 0:
        print("没有可用GPU,将使用CPU")
        return -1  # 返回-1表示使用CPU
    
    selected_gpu = -1
    for gpu_id in range(num_gpus):
        # 临时切换到当前GPU设备
        torch.cuda.device(gpu_id).__enter__()
        # 获取当前GPU已分配的显存(单位:字节)
        allocated_memory = torch.cuda.memory_allocated()
        # 获取当前GPU缓存的显存(单位:字节)
        cached_memory = torch.cuda.memory_cached()
        # 总使用显存 = 已分配 + 缓存(缓存的显存也无法被其他程序占用)
        total_used_memory = allocated_memory + cached_memory
        
        # 比较显存使用量是否低于设定阈值(注意单位统一为字节)
        if total_used_memory < threshold_memory:
            selected_gpu = gpu_id
            # 退出当前设备上下文,避免影响后续操作
            torch.cuda.device(gpu_id).__exit__(None, None, None)
            break
        # 退出当前设备上下文
        torch.cuda.device(gpu_id).__exit__(None, None, None)
    
    if selected_gpu == -1:
        print("所有GPU显存使用量都超过阈值,将默认使用第一个GPU")
        selected_gpu = 0
    
    return selected_gpu

# 示例:选择显存使用量低于1GB(转换为字节单位)的GPU
MEM_THRESHOLD = 1 * 1024 ** 3
gpu_id = select_gpu_by_memory(MEM_THRESHOLD)
print(f"选中的GPU ID: {gpu_id}")

# 后续设置PyTorch使用选中的GPU
if gpu_id != -1:
    torch.cuda.set_device(gpu_id)

关键细节说明

  • 设备切换逻辑:在PyTorch 0.3.1中,查询特定GPU的显存必须先切换到该设备上下文,这里用__enter__()和__exit__()实现临时切换,不会改变全局默认设备。
  • 显存统计函数解析:
    • torch.cuda.memory_allocated():返回当前GPU上张量实际占用的显存大小(字节)。
    • torch.cuda.memory_cached():返回PyTorch为该GPU预留的缓存显存(字节),这部分显存虽未被张量直接使用,但也无法被其他进程占用,因此需要计入总使用量。
  • 阈值单位统一:如果你的阈值是以GB/MB为单位,一定要转换成字节(比如1GB = 102410241024字节),避免统计逻辑出错。

特殊情况处理

  • 若系统无可用GPU,函数返回-1并提示使用CPU。
  • 若所有GPU显存都超过阈值,默认返回第一个GPU(你可以根据需求修改这部分逻辑,比如抛出提示或返回其他默认值)。

内容的提问来源于stack exchange,提问作者Wasi Ahmad

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.21 07:55:31