Files

84 lines
2.8 KiB
Markdown
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# NVIDIA NeMo / Megatron-Bridge Backend
本项目的训练后端目标是 NVIDIA 官方 NeMo/Megatron 栈,而不是旧的手写 PyTorch trainer。
## 默认训练镜像
g0050 上当前默认训练镜像:
```text
laoyao/nemo-megatron:26.06-flashattn4
```
该镜像基于 NeMo 26.06/Megatron-Bridge 0.5.0,当前关键包版本:
| package | version |
|---|---|
| torch | 2.12.0a0 NVIDIA build |
| transformer-engine | 2.16.0 |
| megatron-core | 0.18.0 |
| megatron-bridge | 0.5.0 |
| flash-attn-4 | 4.0.0b11 |
| nvidia-cutlass-dsl | 4.5.2 |
`flash-attn` v2 必须从该镜像中移除。NeMo 26.06 将 FA2 安装在系统
`/usr/local/lib/python3.12/dist-packages`,而 FA4 安装在 `/opt/venv`;如果两者并存,
系统 FA2 会抢占顶层 `flash_attn` 包,Transformer Engine 导入
`flash_attn.cute.interface` 时会失败。
构建脚本固定并验证以下组合:
- `INSTALL_FLASH_ATTN4=1`,默认启用 FA4 安装。
- `pip --pre`,允许安装 FA4 的 beta 发行版。
- `flash-attn-4==4.0.0b11`。
- `nvidia-cutlass-dsl`、`nvidia-cutlass-dsl-libs-base` 和
`nvidia-cutlass-dsl-libs-cu13` 均为 `4.5.2`。
- 同时从系统 Python 和 `/opt/venv` 删除旧 FA2/FA3,再安装 FA4。
- 构建后必须通过 `flash_attn.cute.interface`、`cutlass` 和
`triton_kernels.matmul_ogs` 的真实导入检查;仅看到 distribution metadata 不算成功。
## 构建镜像
在 g0050 上执行:
```bash
cd /ssd/workspace/yi/laoyao_2b_moe
bash scripts/build_nemo_megatron_image.sh
```
如果机器上已有 `ti-coding-agent/nemo-bridge-flashattn4:26.06`,可以将其作为 base image;
无论 base 是否已包含 FA4,构建脚本默认都会重新校验并固化上述版本。最终镜像为:
```text
laoyao/nemo-megatron:26.06-flashattn4
```
如果换机器没有这个缓存镜像,从 NVIDIA 官方 NeMo 26.06 镜像构建:
```bash
BASE_IMAGE=nvcr.io/nvidia/nemo:26.06 \
bash scripts/build_nemo_megatron_image.sh
```
`INSTALL_FLASH_ATTN4` 默认已设为 `1`,上面的命令无需额外传入该参数。只有明确进行
官方 NeMo 原始环境对照实验时,才应设置 `INSTALL_FLASH_ATTN4=0`。
FA4 和 CUTLASS DSL 是预发布/快速迭代包。不要将版本改成未固定的最新版,否则可能造成
CUTLASS base 与 `libs-cu13` 版本不一致。阿里云 simple index 未完整列出这两个 CUTLASS
wheel,因此 Dockerfile 使用固定的阿里云 wheel URL,避免回退到低速的
`files.pythonhosted.org`。
脚本默认使用 B300/g0050 代理 `http://100.72.0.101:8888`,pip 默认走阿里云源并保留 PyPI fallback。
## 训练入口
```bash
bash scripts/train_megatron_bridge_2b_moe.sh
```
训练脚本默认使用 `laoyao/nemo-megatron:26.06-flashattn4`。如需临时切回官方镜像,可覆盖:
```bash
IMAGE=nvcr.io/nvidia/nemo:26.06 bash scripts/train_megatron_bridge_2b_moe.sh
```