通用设备支持实施
作者: jryanhaber创建于 2025年8月12日更新于 2025年10月17日
问题 当前代码库假定使用 NVIDIA GPU 环境,并使用硬编码的 CUDA 调用(例如 torch.device("cuda") 和 .cuda()),这会阻止项目在 Apple Silicon (MPS) 或仅使用 CPU 的系统上运行。这使得非 CUDA 硬件上的贡献者和用户难以运行推理或评估,即使在硬件完全能够执行模型的情况下也是如此。此外,项目中存在与 adam_atan2_backend 模块的依赖问题,可能会阻止在没有正确编译或可用的后端的系统上导入。 实现的解决方案 1. 通用设备检测 我们引入了 get_device() 帮助函数,它自动检测并选择最佳可用后端,优先级如下: 1. MPS (Metal Performance Shaders) 在 Apple Silicon 上 2. CUDA 如果在 NVIDIA 硬件上可用 3. CPU 作为通用的回退方案 实现代码在 pretrain.py:14-24 和 evaluate.py:9-19: def get_device(): import torch if torch.backends.mps.is_available(): return torch.device("mps") elif torch.cuda.is_available(): return torch.device("cuda") else: return torch.device("cpu") device = get_device() print(f"Using device: {device}")
- 全面性的 CUDA 调用替换 我们系统性地替换了整个代码库中所有硬编码的 CUDA 引用: 修改的文件: - pretrain.py: 8 个位置已更新 - evaluate.py: 3 个位置已更新 所做的更改: - .cuda() → .to(device) 用于张量设备转换 - torch.device("cuda") → torch.device(device) 用于上下文管理器 - device="cuda" → device=device 用于张量创建 - map_location="cuda" → map_location=device 用于检查点加载 - 条件性 CUDA 设备设置: 仅在 device.type == "cuda" 时调用 torch.cuda.set_device() 3. 优化器回退实现 我们为 adam_atan2 优化器依赖添加了一个可靠的回退机制,以处理缺少后端编译: 实现代码在 pretrain.py:32-36: try: from adam_atan2 import AdamATan2 except ImportError: # 回退到 AdamW,当 adam_atan2_backend 不可用时 from torch.optim import AdamW as AdamATan2 理由: AdamW 是 PyTorch 最稳定和广泛使用的优化器,确保: - 无推理影响(优化器在推理期间未使用) - 在所有后端之间的高度兼容性 - 维护者友好的方法,保持 CUDA 优势 - 对大多数训练场景而言的算法差异最小 技术实现细节 设备检测逻辑 # 最佳性能的优先级顺序 #### 1. MPS - 原生 Apple Silicon 加速 #### 2. CUDA - NVIDIA GPU 加速 #### 3. CPU - 通用兼容性 向后兼容性 - 当检测到 CUDA 硬件时,保留所有现有的 CUDA 功能 - 保持分布式训练支持,使用 torch.distributed 扩展包
内容来源: sapientinc/HRM