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
相关产品推荐
相关产品推荐

