Skip to content

feat: 支持 AMD GPU (ROCm) 训练 #24

Description

@wlgys8

背景与动机

当前训练管线(FastSAC / RSL-RL PPO)假定 NVIDIA CUDA 环境,AMD GPU(ROCm/HIP)用户无法开箱即用地进行 GPU 训练。希望在 ROCm 环境下完成从安装、训练到 GPU 利用率监控的完整链路。

现状:CUDA 硬绑定的位置

PyTorch 的 ROCm wheel 会把 HIP 暴露为 torch.cuda API,因此 torch.cuda.is_available() 在 AMD GPU 上返回 True,大部分 device 选择逻辑可复用;真正的阻断点在安装源、镜像和部分 CUDA 专属路径:

  • 安装源固定 cu128pyproject.toml[tool.uv.sources] 将 torch/torchvision/torchaudio 固定到 pytorch-cu128 index(pyproject.toml:63-88),AMD 用户 uv sync 后装的是 CUDA wheel,无法使用本机 GPU
  • Docker 镜像为 NVIDIA 基础镜像docker/Dockerfile:4 使用 nvidia/cuda:12.8.1-runtime-ubuntu24.04
  • GPU 监控仅支持 NVMLscripts/gpu_utils.py 使用 pynvml,AMD GPU 上直接失败
  • CUDA 专属加速路径(ROCm 下行为未验证,需逐一确认或加保护):
    • motrix_rl/src/motrix_rl/fastsac/async_impl/collector.py:FP16 autocast、pinned weight staging、flat 参数绑定等均以 device.type == "cuda" 开启
    • motrix_rl/src/motrix_rl/fastsac/agent.py:113,149,168:fused optimizer、AMP、torch.compile 均以 device.type == "cuda" 开启
    • motrix_rl/src/motrix_rl/rslrl/torch/train/ppo.py:60,128:硬编码 torch.device("cuda:0")
    • configs/algo_base/motrix.fastsac.yamlcollector_inference_device: cuda 等默认值与注释仅提及 CUDA
  • GPU 并行仿真后端:训练吞吐依赖 MotrixSim 的 GPU pipeline,其是否支持 ROCm/HIP 需要单独确认(超出本仓库范围的话需拆子任务)

建议任务

  • 安装:支持 ROCm wheel 源(如 pytorch-rocm uv index + extra/文档说明),或按 uv 环境变量切换
  • Docker:提供 rocm 基础镜像变体(或文档说明如何替换基础镜像并挂载 /dev/kfd/dev/dri
  • GPU 监控:scripts/gpu_utils.py 抽象出 NVML / ROCm(rocm-smi 或 amd-smi)后端
  • 校验并修复 FastSAC CUDA 专属路径在 ROCm 下的可用性(collector FP16/staging/compile、agent fused optimizer/AMP/torch.compile),不可用路径需按能力探测降级而不是假设 CUDA
  • RSL-RL PPO:移除硬编码 cuda:0,改为遵循配置的 device
  • 配置与文档:默认值/注释改为 "GPU (CUDA/ROCm)" 表述,README/docs 增加 AMD GPU 安装与运行说明
  • 验证:在至少一块 AMD GPU(如 RX 7900 系列或 MI 系列)上跑通 FastSAC 与 RSL-RL PPO 端到端训练

验收标准

  • AMD GPU 机器上按文档执行 uv sync + 训练命令即可使用 ROCm 后端,无需修改源码
  • NVIDIA 路径行为不变(现有 CI 与默认配置不受影响)
  • FastSAC 与 RSL-RL PPO 在 ROCm 下端到端训练收敛,性能数据(可选)记录到 bench

待确认

  1. MotrixSim GPU 仿真后端是否有 ROCm/HIP 支持计划?若暂无,是否先支持「AMD GPU 训练 + CPU/低端仿真」或明确声明依赖?
  2. 目标 ROCm 版本与 PyTorch 版本组合(如 ROCm 6.x + torch 2.x)?

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

No labels
No labels

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions