这个镜像 rocm/pytorch:rocm7.2_ubuntu24.04_py3.12_pytorch_release_2.9.1 是 AMD 官方专门为开发者打包好的“全家桶”,里面已经预装好了所有环境。

[ROCm 安装上的 PyTorch — ROCm 安装(Linux)]

你只需要按照下面的步骤启动它即可。

1. 启动容器命令

在你的终端执行以下完整命令(注意这里包含了访问 GPU 所需的所有关键权限):

sudo docker run -it \
  --cap-add=SYS_PTRACE \
  --security-opt seccomp=unconfined \
  --device=/dev/kfd \
  --device=/dev/dri \
  --group-add video \
  --ipc=host \
  --shm-size 8G \
  rocm/pytorch:rocm7.2_ubuntu24.04_py3.12_pytorch_release_2.9.1

参数通俗解释:

  • --device=/dev/kfd/dev/dri: 这是把你的显卡硬件“映射”进容器。
  • --group-add video: 确保容器内的用户有权限控制视频设备(显卡)。
  • --shm-size 8G: 增加“共享内存”。PyTorch 训练模型时如果这个值太小(默认通常只有 64MB),会报 Bus error 崩溃。
  • --ipc=host: 让容器和宿主机共享进程间通信,对多卡训练或高性能计算非常重要。
  • --cap-add=SYS_PTRACE: 允许一些底层的性能分析工具运行。

2. 测试 GPU 和 PyTorch 是否正常

进入容器后,你会看到类似 root@xxxx:/# 的提示符。请按顺序执行以下测试:

测试一:ROCm 硬件层面(验证显卡通信)

输入:

rocminfo

如果你能看到一长串信息,并且在其中能找到你的 GPU 型号(比如 gfx1100Radeon...),说明物理层连接没问题。

测试二:PyTorch 层面(最关键)

进入 Python 环境:

python3

然后在 Python 提示符 >>> 后输入:

import torch
print(f"PyTorch 版本: {torch.__version__}")
print(f"GPU 是否可用: {torch.cuda.is_available()}")
print(f"GPU 数量: {torch.cuda.device_count()}")
print(f"当前 GPU 名称: {torch.cuda.get_device_name(0)}")

注意: 在 ROCm 下,PyTorch 依然沿用 torch.cuda 这个 API 名称。只要 torch.cuda.is_available() 返回 True,就说明你的 PyTorch 已经成功调用了 AMD 显卡。

测试三:简单矩阵运算(压力测试)

在 Python 里继续输入:

# 在 GPU 上创建一个随机张量并做加法
x = torch.randn(1000, 1000).to("cuda")
y = x + x
print(y.device)  # 输出应该是 cuda:0

3. 进阶建议:如何运行你的代码?

如果你想把电脑里的代码文件夹(比如 ~/my_project)带进容器里运行,启动命令需要加一个 -v 参数:

sudo docker run -it \
  --device=/dev/kfd --device=/dev/dri --group-add video --ipc=host --shm-size 8G \
  -v /home/wtw/my_project:/workspace \
  rocm/pytorch:rocm7.2_ubuntu24.04_py3.12_pytorch_release_2.9.1

这样你进入容器后,直接 cd /workspace 就能看到你宿主机上的代码了。

4. 常见问题排查

  • 权限报错: 如果启动时提示 Permission denied 无法访问 /dev/kfd,请确认你的宿主机用户是否在 rendervideo 组中,或者直接用 sudo 启动。
  • 显存不足: 如果运行 torch 代码时报错 Out of Memory,可以用 rocm-smi 命令查看当前显存占用情况。
  • 无法解析 GPU: 如果 rocminfo 正常但 is_available() 是 False,可能是因为你的显卡型号太新或太旧(如 RX 580),可能需要设置环境变量来强制兼容(例如 export HSA_OVERRIDE_GFX_VERSION=10.3.0),具体取决于你的显卡架构。