- 一、前期调研与适配思路
- 1.1 为什么要适配 DiT4DiT
- 1.2 DiT4DiT 与 StarVLA 的关系
- 1.3 适配难度评估
- 1.4 适配路线
- 二、环境准备
- 2.1 硬件与驱动
- 2.2 软件版本
- 2.3 复用 StarVLA conda 环境
- 2.4 安装 PyTorch + torch_npu
- 2.5 安装 ffmpeg + decord
- 2.6 安装 mx_driving(NPU Patcher)
- 2.7 安装 DiT4DiT 依赖
- 2.8 模型权重下载
- 2.9 LD_LIBRARY_PATH
- 三、代码适配详解
- 3.1 train.py(训练入口)— 6 处修改
- ① 导入 torch_npu + patcher
- ② autocast 替换
- ③ AdamW 融合优化器
- ④ patcher 初始化
- ⑤ DeepSpeedPlugin 去硬编码
- ⑥ 模型保存
- 3.2 DiT4DiT.py(模型框架)— 4 处修改
- ① autocast 替换(4 处)
- ② action_mask 截断对齐
- ③ bf16 → numpy 转换
- ④ @torch.inference_mode() → @torch.no_grad()
- 3.3 Cosmos25.py(backbone)— 1 处关键修改
- 四、DeepSpeed 分布式配置
- 4.1 ZeRO-2 配置(deepspeed_zero2.yaml)
- 4.2 ZeRO-3 配置(deepspeed_zero3.yaml)
- 五、8 卡启动脚本与训练配置
- 5.1 启动脚本 train_8p.sh
- 5.2 训练配置(Robotwin 数据集)
- 六、ZeRO-2 vs ZeRO-3 性能对比
- 七、踩坑记录
- 7.1 num_processes 必须匹配物理卡数
- 7.2 mx_driving.patcher API 变更
- 7.3 libc10.so / libtorch_npu.so 未自动发现
- 7.4 Cosmos-Predict2.5-2B 显存占用
- 7.5 VAE / Text Encoder 混合 dtype → ZeRO-3 defragment 报错
- 7.6 eval_action_model 维度不匹配
- 7.7 bf16 → numpy 不支持
- 7.8 @torch.inference_mode() 与 ZeRO-3 冲突
- 7.9 ZeRO-3 模型保存卡死
- 7.10 accelerate config 中的 mixed_precision 与 DeepSpeed bf16 冲突
- 7.11 视频解码 decord 的 CUDA 默认设备
- 7.12 HCCL 通信超时(偶发)
- 八、适配 Patch
- 九、总结
DiT4DiT 是港科广团队在 StarVLA 基础上提出的视觉-动作模型(VAM),将视频生成 Diffusion Transformer 与 Flow Matching 动作预测结合,支持机械臂灵巧操作与人形机器人全身控制。原始代码完全基于 CUDA 生态(cuDNN、NCCL、torch.cuda),而华为昇腾 Atlas 800T A3 使用自研达芬奇架构 NPU,两者在算子支持、分布式通信、混合精度等方面存在显著差异。本文记录了从零将 DiT4DiT 适配到昇腾 NPU 的全过程,包括前期调研、环境配置、代码移植、ZeRO-2/ZeRO-3 两种分布式策略的调优,以及 12+ 个踩坑点的定位与修复,最终在 RoboTwin 双臂数据集上成功跑通训练。
Git 仓库版本:commit 1ae6efd(github.com/Mondo-Robotics/DiT4DiT),适配 patch 已上传至服务器,文末附下载链接。
一、前期调研与适配思路
1.1 为什么要适配 DiT4DiT
具身智能赛道当前主流的 VLA 模型(π0.5、OpenVLA、StarVLA)大多基于单一视觉编码器 + 动作头结构。DiT4DiT 的不同之处在于它引入了视频扩散模型作为视觉 backbone——这带来了两个关键优势:
- 时序建模更强:Cosmos-Predict2.5-2B 原生支持视频帧间动态建模,比静态图像编码器对动作连续性的感知更敏锐
- 视频辅助 loss:训练时可以同时优化动作预测和未来帧重建,这是纯 VLA 做不到的多模态监督
但代价也很明显:视频 backbone 的参数量和计算量远大于传统 ViT,对算力和显存的要求更高。正好我们有 Atlas 800T A3(910B)8 卡环境,64GB HBM 单卡显存对 2B 级模型是充足的,值得一试。
1.2 DiT4DiT 与 StarVLA 的关系
两个开源仓库的关系梳理如下:
| StarVLA | DiT4DiT | |
|---|---|---|
| 定位 | VLA 基础框架 | 基于 StarVLA 的视频 VAM 扩展 |
| 基座模型 | Qwen3-VL-4B-Instruct | Cosmos-Predict2.5-2B |
| 动作预测 | 标准 diffusion head(DiT-B) | Flow-matching action head |
| 视频模型 | 无 | 视频 diffusion + VAE 编解码 |
| 训练框架 | accelerate + DeepSpeed | accelerate + DeepSpeed(复用) |
| 数据格式 | LeRobot | LeRobot(复用) |
| 华为适配 | ✅ 已有 NPU 版(DrivingSDK) | ❌ 暂无(本文完成) |
| 仓库地址 | gitcode.com/Ascend/DrivingSDK | github.com/Mondo-Robotics/DiT4DiT |
可以看出 DiT4DiT 在训练框架层面大量复用了 StarVLA 的代码(accelerate 启动器、DeepSpeed 配置、LeRobot dataloader、DiT-B 动作头),核心差异只在于视觉 backbone 从 Qwen3-VL 换成了 Cosmos-Predict2.5-2B。这意味着适配工作量不会太大——华为已经跑通了 StarVLA 的 NPU 版本,我们只需要在 DiT4DiT 上复刻相同的移植操作即可。
1.3 适配难度评估
在动手之前,先梳理 DiT4DiT 代码中可能涉及 CUDA 硬编码的模块:
| 模块 | CUDA 依赖点 | 适配难度 |
|---|---|---|
| 训练入口 train.py | torch.autocast("cuda")、AdamW fused | ⭐ 低(字符串替换) |
| 模型框架 DiT4DiT.py | autocast、inference_mode、bf16→numpy | ⭐⭐ 中(语义差异) |
| backbone Cosmos25.py | 混合 dtype 参数、VAE 冻结策略 | ⭐⭐ 中(ZeRO 兼容) |
| 视频编解码 VAE | Conv3D、GroupNorm 算子 | ⭐ 低(torch_npu 原生支持) |
| DeepSpeed ZeRO-3 | 参数分片 gather、defragment | ⭐⭐⭐ 高(多坑) |
| 分布式通信 | NCCL → HCCL(自动转换) | ⭐ 低(torch_npu 自动适配) |
总体来看,DiT4DiT 没有像 flash-attention、triton kernel 这类深度绑定 CUDA 的自定义算子,绝大多数 CUDA 调用都是 PyTorch 高层 API(autocast、DistributedDataParallel),torch_npu 已经做了适配。真正的难点在 ZeRO-3 的兼容性和一些 NPU 特有的数据格式约束(如 bf16 不支持 .numpy())。
1.4 适配路线
确定三步走策略:
- 基础 NPU 移植:参照华为 DrivingSDK 中 StarVLA 的 starvla.patch,对 DiT4DiT 做相同的 autocast 替换 + torch_npu 导入 + mx_driving patcher 初始化。这一步完成后代码应能在单卡上不报错地跑通前向。
- 分布式跑通:配置 DeepSpeed ZeRO-2,8 卡加速训练。ZeRO-2 对代码侵入最小,参数完整保留在每张卡上,先保证能跑起来。
- ZeRO-3 调优:切换到 ZeRO-3 以支持更大 batch size。这一阶段会暴露最多兼容性问题(混合 dtype、inference_mode、模型保存等),需要逐一修复。
二、环境准备
2.1 硬件与驱动
硬件为 Atlas 800T A3 训练服务器,单机 8 卡昇腾 910B,每卡 64GB HBM。在进行 NPU 适配前,先确认驱动和固件版本:
# 查看 NPU 状态
npu-smi info
# 确认驱动版本
npu-smi info -t driver -i 0
# 查看 HCCL 网络拓扑(8 卡应全互联)
npu-smi info -t topo
预期输出:8 张卡全部显示 "OK",健康状态正常,驱动版本与 CANN 版本匹配(CANN 9.0.0 对应驱动 24.1.rc1+)。如果驱动版本过旧,需联系集群管理员更新。
2.2 软件版本
以下版本组合已在 Atlas 800T A3 上验证通过:
| 软件 | 版本 | 说明 |
|---|---|---|
| Python | 3.10 | DiT4DiT 官方要求 ≥ 3.10 |
| CANN | 9.0.0 | 昇腾 AI 计算框架,含算子库和编译器 |
| PyTorch | 2.7.1 | 需配套 torch_npu 版本 |
| torch_npu | 2.7.1.post2 | 昇腾 NPU 后端,提供 "npu" device |
| torchvision | 0.22.1 | 视频帧预处理依赖 |
| DeepSpeed | 0.18.4 | 分布式训练框架 |
| accelerate | 1.12.0 | HuggingFace 训练启动器 |
| diffusers | 0.38.0 | Cosmos 模型加载 |
| transformers | 4.57.0 | HuggingFace 模型生态 |
| decord | 0.6.0 | 视频解码(需源码编译) |
| ffmpeg | 4.4.2 | 视频编解码基础库 |
| mx_driving | 1.0.20260421 | 华为 NPU patcher,自动修复算子兼容性 |
版本选择注意:torch 和 torch_npu 必须是小版本严格匹配的——你装 torch==2.7.1,torch_npu 就得是 2.7.1.postX。如果装成 torch==2.7.0 + torch_npu==2.7.1,运行时会报 libtorch_npu.so: undefined symbol。
2.3 复用 StarVLA conda 环境
既然 StarVLA 的 NPU 环境已经配置好了(PyTorch + torch_npu + CANN + 各种依赖),直接复用可以省去大量重复安装工作。只需在其基础上补充 DiT4DiT 特有的依赖:
# 如果已有 StarVLA 的 conda 环境,直接克隆一份
conda create -n dit4dit --clone starvla_env_name
conda activate dit4dit
# 如果从头开始:
conda create -n dit4dit python=3.10
conda activate dit4dit
2.4 安装 PyTorch + torch_npu
升腾 PyTorch 的 whl 包通常由集群管理员预置在本地镜像或 SWR 仓库中。如果本地已有 pip 源:
pip install torch==2.7.1
pip install torch_npu==2.7.1.post2
pip install torchvision==0.22.1
# 验证 NPU 可用
python -c "import torch; import torch_npu; print(torch.npu.is_available()); print(torch.npu.device_count())"
# 预期输出: True \n 8
2.5 安装 ffmpeg + decord
DiT4DiT 的视频输入依赖 decord 做帧抽取,decord 又依赖 ffmpeg 做解码。注意 decord 不支持 pip install,必须源码编译:
推荐方案(conda 安装 ffmpeg):
conda install -c conda-forge ffmpeg=4.4.2
备选方案(源码安装 ffmpeg):
wget https://ffmpeg.org/releases/ffmpeg-4.4.2.tar.bz2
tar -xvf ffmpeg-4.4.2.tar.bz2
cd ffmpeg-4.4.2
./configure --enable-shared --prefix=/usr/local/ffmpeg
make -j 64
make install
cd ..
echo 'export PATH="/usr/local/ffmpeg/bin:$PATH"' >> /etc/profile.d/ffmpeg.sh
echo 'export LD_LIBRARY_PATH="/usr/local/ffmpeg/lib:$LD_LIBRARY_PATH"' >> /etc/profile.d/ffmpeg.sh
source /etc/profile
编译 decord:
git clone --recursive https://github.com/dmlc/decord --depth 1
cd decord
mkdir build && cd build
cmake .. -DCMAKE_BUILD_TYPE=Release -DFFMPEG_DIR:PATH=$CONDA_PREFIX
make
cd ../python
python setup.py sdist bdist_wheel
cd ../..
pip install decord/python/dist/decord-0.6.0-cp310-cp310-linux_aarch64.whl
注意 FFMPEG_DIR 路径:conda 安装时用 $CONDA_PREFIX,源码安装时用 /usr/local/ffmpeg/。路径搞错了 cmake 会找不到 ffmpeg 的 .so,编译出来的 decord 导入时会报 libavcodec.so not found。
2.6 安装 mx_driving(NPU Patcher)
mx_driving 是华为提供的 NPU 算子兼容层,核心功能是拦截 torch 的算子调用并替换为 NPU 等效实现。没有它,很多代码看起来只是改了 autocast 就能跑,实际会遇到隐式的 CUDA 算子调用(比如某些 transformers 内部的 scaled_dot_product_attention 默认走 CUDA 路径)。
# mx_driving 从 DrivingSDK 仓库获取
# 参考 https://gitcode.com/Ascend/DrivingSDK
# 安装命令由集群管理员提供,通常是本地 pip install
验证安装:
python -c "from mx_driving.tools.patch.patcher import default_patcher_builder; print('OK')"
2.7 安装 DiT4DiT 依赖
下载 DiT4DiT 源码后(commit 1ae6efd),安装 requirements:
cd DiT4DiT
pip install -r requirements.txt
pip install -e .
完整 requirements.txt 如下(基于 NPU 环境实际 pip list 整理):
absl-py==2.3.1
accelerate==1.12.0
albucore==0.0.17
albumentations==1.4.18
av==12.3.0
certifi==2026.1.4
charset-normalizer==3.4.4
click==8.3.1
contourpy==1.3.2
cramjam==2.11.0
cycler==0.12.1
deepspeed==0.18.4
diffusers==0.38.0
docstring_parser==0.17.0
einops==0.8.1
einx==0.3.0
eval_type_backport==0.3.1
fastparquet==2024.11.0
filelock==3.20.3
fonttools==4.61.1
frozendict==2.4.7
fsspec==2026.1.0
fvcore==0.1.5.post20221221
gitdb==4.0.12
GitPython==3.1.46
greenlet==3.3.0
grpcio==1.76.0
hf-xet==1.2.0
hjson==3.1.0
huggingface-hub==0.36.0
hyper-connections==0.4.6
idna==3.11
imageio==2.37.0
imageio-ffmpeg==0.6.0
importlib_metadata==8.7.1
iopath==0.1.10
Jinja2==3.1.6
joblib==1.5.3
kiwisolver==1.4.9
lazy_loader==0.4
Markdown==3.10
markdown-it-py==4.0.0
MarkupSafe==3.0.3
matplotlib==3.10.8
mdurl==0.1.2
mpmath==1.3.0
msgpack==1.1.2
networkx==3.4.2
ninja==1.13.0
nltk==3.9.1
numpy==1.26.4
numpydantic==1.6.9
omegaconf==2.3.0
opencv-python==4.10.0.84
opencv-python-headless==4.11.0.86
packaging==25.0
pandas==2.3.3
peft==0.18.1
pillow==11.1.0
pipablepytorch3d==0.7.6
portalocker==3.2.0
protobuf==6.33.4
psutil==7.2.1
pyarrow==14.0.1
pydantic==2.10.6
Pygments==2.19.2
pyparsing==3.3.1
python-dateutil==2.9.0.post0
pytz==2025.2
PyYAML==6.0.3
qwen-vl-utils==0.0.14
regex==2026.1.15
requests==2.32.5
rich==14.2.0
safetensors==0.5.3
scikit-image==0.25.2
scipy==1.15.3
sentencepiece==0.2.0
sentry-sdk==2.49.0
six==1.17.0
smmap==5.0.2
sympy==1.13.3
tabulate==0.9.0
tensorboard==2.20.0
tensorboard-data-server==0.7.2
tifffile==2025.5.10
tiktoken==0.12.0
timm==1.0.24
tokenizers==0.22.2
torch==2.7.1
torch-npu==2.7.1.post2
torchvision==0.22.1
tqdm==4.66.5
transformers==4.57.0
typing_extensions==4.15.0
tyro==1.0.5
urllib3==2.6.3
wandb==0.24.0
websockets==16.0
Werkzeug==3.1.5
yacs==0.1.8
zipp==3.23.0
2.8 模型权重下载
Cosmos-Predict2.5-2B 权重需下载特定的 diffusers 格式版本(base/post-trained 分支),不是默认的 main 分支模型格式。务必指定 revision:
huggingface-cli download nvidia/Cosmos-Predict2.5-2B \
--revision diffusers/base/post-trained \
--local-dir playground/Pretrained_models/Cosmos-Predict2.5-2B
下载完成后验证目录结构:
ls playground/Pretrained_models/Cosmos-Predict2.5-2B/
# 应包含: transformer/, vae/, text_encoder/, scheduler/, model_index.json
2.9 LD_LIBRARY_PATH
这也是一个高频踩坑点。conda 环境下 torch_npu 的 .so 文件不会自动加入动态库搜索路径,导致启动训练时找不到 libc10.so 或 libtorch_npu.so。
方案一:启动脚本中显式设置(推荐,简单直接)
export LD_LIBRARY_PATH=/opt/conda/envs/torch2.7.1/lib/python3.10/site-packages/torch/lib:/opt/conda/envs/torch2.7.1/lib/python3.10/site-packages/torch_npu/lib:$LD_LIBRARY_PATH
路径中 torch2.7.1 需替换为你的实际 conda 环境名。
方案二:写入 conda activate 钩子(一劳永逸)
mkdir -p /opt/conda/envs/torch2.7.1/etc/conda/activate.d
cat > /opt/conda/envs/torch2.7.1/etc/conda/activate.d/ld_path.sh << 'EOF'
export LD_LIBRARY_PATH=/opt/conda/envs/torch2.7.1/lib/python3.10/site-packages/torch/lib:/opt/conda/envs/torch2.7.1/lib/python3.10/site-packages/torch_npu/lib:$LD_LIBRARY_PATH
EOF
两种方案本质一样,方案二的好处是 conda activate 后自动生效。但如果环境名变了(比如你又克隆了一个新环境),需要重新创建钩子文件。
三、代码适配详解
适配只需修改三个文件。以下按文件逐一说明改动点和原因。
3.1 train.py(训练入口)— 6 处修改
① 导入 torch_npu + patcher
在 import torch.distributed 之后插入:
import torch_npu
from mx_driving.tools.patch.patcher import default_patcher_builder
import torch_npu 的作用是注册 "npu" device,之后 torch.device("npu") 和 torch.autocast("npu") 才能正常使用。即使你不直接调用 torch_npu 的任何 API,也必须 import,否则 accelerate 在初始化时会因为找不到 "npu" device 而 fallback 到 CPU。
② autocast 替换
全文搜索 torch.autocast("cuda") → torch.autocast("npu")。这一步就是字面替换,但需要注意:
- 不要用全局 sed:有些代码里的 "cuda" 是注释或字符串内容,比如
print("Running on cuda"),替换了反而会误导日志 - autocast 参数保持一致:torch_npu 的 autocast 支持与 CUDA 相同的
dtype、enabled参数,不需要改参数
③ AdamW 融合优化器
在创建 AdamW 优化器时,加入 fused=True:
optimizer = torch.optim.AdamW(params, lr=lr, fused=True)
torch_npu 从 2.1 开始支持 fused AdamW,将参数更新和梯度缩放融合为单个 NPU kernel,能减少 ~15% 的优化器耗时。不过 fused 模式对参数 dtype 有一定要求——如果某组参数是 fp32,fused 可能回退到普通实现,不会报错只会慢一些。
④ patcher 初始化
在 if __name__ == "__main__": 入口处,所有业务逻辑之前调用:
default_patcher_builder.build().__enter__()
注意 API 变更:旧版 StarVLA 的 patch 使用 Patcher().add(TransformersNPU).apply(),但新版 mx_driving(1.0.20260421+)已废弃 TransformersNPU 类,改为 default_patcher_builder。这个 builder 内部预设了 mmcv、torch、numpy、mmdet 的全套 patch,不需要手动指定 patch 项。
⑤ DeepSpeedPlugin 去硬编码
原代码中 DeepSpeedPlugin 硬编码了 hf_ds_config 参数指向 HuggingFace 的默认 ZeRO-2 配置。NPU 场景下需要用自己的配置(不同卡数显存策略不同),删除该参数:
# 旧:
deepseed_plugin = DeepSpeedPlugin(hf_ds_config="path/to/config.json", ...)
# 新:
deepseed_plugin = DeepSpeedPlugin() # 通过 accelerate config 文件指定
⑥ 模型保存
ZeRO-3 场景下,accelerator.save_state() 会尝试 gather 所有分片参数到 rank 0,2.3B 参数在 8 卡 HCCL 网络上 gather 耗时数分钟且容易超时。改用 DeepSpeed 原生的 model.save_checkpoint():
# 替换 accelerator.save_state()
model.save_checkpoint(save_dir, tag=f"step_{global_step}")
save_checkpoint 每张卡各自写自己的参数分片,不跨卡通信。加载时需要同样的 ZeRO-3 配置,用 model.load_checkpoint() 恢复。
3.2 DiT4DiT.py(模型框架)— 4 处修改
① autocast 替换(4 处)
DiT4DiT 框架中有 4 处 torch.autocast("cuda"),分别在 forward()、training_step()、predict_action() 和 compute_loss() 中,全部替换为 "npu"。
② action_mask 截断对齐
这是一个隐蔽的 bug:模型输出的 action_mask 维度为 8(action_horizon),而 LeRobot 数据集中的 mask 维度为 16(obs_horizon)。原 CUDA 代码在训练时通过 broadcast 隐式对齐了,但 NPU 的 HCCL broadcast 语义有所不同,导致形状不匹配抛出 RuntimeError。
# 修复:在 loss 计算前显式截断
action_mask = action_mask[:, :action_horizon] # 截断到 8
③ bf16 → numpy 转换
predict_action() 中,模型输出是 bf16 tensor,代码尝试 .cpu().numpy()。CUDA 的 bf16 tensor move 到 CPU 后会自动转 float32,但 NPU 不会——NPU 的 bf16 tensor 迁移到 CPU 后仍是 bf16,而 numpy 不支持 bf16 dtype,抛出 TypeError: Got unsupported ScalarType BFloat16。
修复:先显式 .float() 再 .numpy():
actions = pred_actions.float().cpu().numpy() # bf16 → fp32 → numpy
④ @torch.inference_mode() → @torch.no_grad()
评估时原代码使用 @torch.inference_mode() 装饰器。这个与 @torch.no_grad() 的区别是 inference_mode 会禁用 autograd engine 的 version counter 跟踪,性能略优。但 ZeRO-3 的 LinearFunctionForZeroStage3 内部需要保存 tensor 用于反向传播中的参数 gather——inference_mode 创建的 tensor 无法被 DeepSpeed 保存,触发 AssertionError: tensor has no storage。
改为 @torch.no_grad() 即可,性能差异可忽略。
3.3 Cosmos25.py(backbone)— 1 处关键修改
Cosmos-Predict2.5-2B 包含三部分:
- VAE(视频编解码器,~80M 参数,fp32 权重)
- Text Encoder(T5-XXL,~4.7B 参数,fp32 权重)
- Transformer(DiT backbone,~2B 参数,bf16 权重)
问题:VAE 和 Text Encoder 是 fp32,Transformer 是 bf16。ZeRO-3 初始化时会遍历所有 trainable 参数,要求 dtype 一致。混合 dtype 触发 defragment 阶段报错:RuntimeError: expected scalar type BFloat16 but found Float。
修复:冻结 VAE + Text Encoder:
# Cosmos25.py 中
self.vae.requires_grad_(False)
self.text_encoder.requires_grad_(False)
这样做有两个好处:
- 解决 ZeRO-3 的 defragment 报错(冻结参数不参与 ZeRO 分片)
- 显著降低显存占用——T5-XXL 的 4.7B 参数不再需要梯度存储,训练显存从 ~58GB 降到 ~42GB(per_device_batch_size=2 时)
对于 Robotwin 这类双臂操作数据,视频 VAE 和语言编码器在预训练阶段已经学得足够好,冻结影响可忽略。
四、DeepSpeed 分布式配置
我们同时配置了 ZeRO-2 和 ZeRO-3 两套方案,放在 DiT4DiT/config/deepseeds/ 下。
4.1 ZeRO-2 配置(deepspeed_zero2.yaml)
ZeRO-2 对代码侵入最小——只分片 optimizer states + gradients,参数完整保留在各卡上。适合快速验证和调试:
{
"zero_optimization": {
"stage": 2,
"overlap_comm": true,
"contiguous_gradients": true,
"reduce_bucket_size": 5e8,
"allgather_bucket_size": 5e8
},
"bf16": {"enabled": true},
"train_batch_size": "auto",
"train_micro_batch_size_per_gpu": "auto",
"gradient_accumulation_steps": "auto"
}
关键参数说明:
- overlap_comm: true:通信与计算重叠,HCCL 网络够快时能减 ~10% step 时间
- reduce_bucket_size / allgather_bucket_size:设为 500MB,适配 910B 的高带宽(HCCS 互联 ~56GB/s),桶太小反而增加通信次数
- gradient_accumulation_steps: "auto":由 accelerate 根据 per_device_batch_size 和总 batch size 自动计算
4.2 ZeRO-3 配置(deepspeed_zero3.yaml)
ZeRO-3 额外分片参数,适合需要更大 batch size 的正式训练:
{
"zero_optimization": {
"stage": 3,
"overlap_comm": true,
"contiguous_gradients": true,
"reduce_bucket_size": 5e8,
"allgather_bucket_size": 5e8,
"stage3_prefetch_bucket_size": 5e8,
"stage3_param_persistence_threshold": 1e6,
"stage3_max_live_parameters": 1e9,
"stage3_max_reuse_distance": 1e9,
"stage3_gather_16bit_weights_on_model_save": true
},
"bf16": {"enabled": true},
"train_batch_size": "auto",
"train_micro_batch_size_per_gpu": "auto",
"gradient_accumulation_steps": "auto"
}
ZeRO-3 特有参数说明:
- stage3_param_persistence_threshold: 1e6:小于 1M 参数的小 tensor(如 LayerNorm 的 γ/β)保留在本地不参与分片,减少 allgather 调用次数
- stage3_max_live_parameters / stage3_max_reuse_distance:控制参数预取窗口,设大一些减少参数 gather 频率
- stage3_gather_16bit_weights_on_model_save: true:保存时 gather 为 bf16(而不是 fp32),节省一半存储
五、8 卡启动脚本与训练配置
5.1 启动脚本 train_8p.sh
完整的 train_8p.sh:
#!/bin/bash
# ========== 环境变量 ==========
export WANDB_MODE=offline
export WANDB_OFFLINE=true
export TASK_QUEUE_ENABLE=2
export PYTORCH_NPU_ALLOC_CONF=expandable_segments:True
export CPU_AFFINITY_CONF=1
export LD_LIBRARY_PATH=/opt/conda/envs/torch2.7.1/lib/python3.10/site-packages/torch/lib:/opt/conda/envs/torch2.7.1/lib/python3.10/site-packages/torch_npu/lib:$LD_LIBRARY_PATH
# ========== 训练参数 ==========
num_processes=8
max_train_steps=150000
per_device_batch_size=2
gradient_accumulation_steps=8
data_mix=robotwin
base_model=./playground/Pretrained_models/Cosmos-Predict2.5-2B
config_yaml=./DiT4DiT/config/robotwin/dit4dit_robotwin.yaml
Framework_name=DiT4DiT
run_root_dir=./results/Checkpoints
run_id=run_$(date +%Y%m%d_%H%M%S)
LOG_FILE=./train_${run_id}.log
# ========== 启动训练 ==========
accelerate launch \
--config_file DiT4DiT/config/deepseeds/deepspeed_zero2.yaml \
--num_processes ${num_processes} \
DiT4DiT/training/train.py \
--config_yaml ${config_yaml} \
--framework.name ${Framework_name} \
--framework.cosmos25.base_model ${base_model} \
--datasets.vla_data.per_device_batch_size ${per_device_batch_size} \
--datasets.vla_data.data_mix ${data_mix} \
--trainer.gradient_accumulation_steps ${gradient_accumulation_steps} \
--trainer.max_train_steps ${max_train_steps} \
--run_root_dir ${run_root_dir} \
--run_id ${run_id} \
2>&1 | tee ${LOG_FILE}
环境变量说明:
| 变量 | 作用 |
|---|---|
| WANDB_MODE=offline | NPU 训练机通常不连外网,离线模式避免 wandb 初始化超时 |
| TASK_QUEUE_ENABLE=2 | 启用 HCCL task queue,减少通信同步开销 |
| PYTORCH_NPU_ALLOC_CONF=expandable_segments:True | 允许 NPU 显存段动态扩展,减少 OOM(类似 CUDA 的 expandable_segments) |
| CPU_AFFINITY_CONF=1 | 绑定 CPU 核心亲和性,避免 8 个训练进程争抢同一 NUMA 节点 |
5.2 训练配置(Robotwin 数据集)
新建配置文件 DiT4DiT/config/robotwin/dit4dit_robotwin.yaml,关键参数:
| 参数 | 值 | 说明 |
|---|---|---|
| base_model | Cosmos-Predict2.5-2B | 视频 backbone |
| data_mix | robotwin | RoboTwin 双臂数据集 |
| action_dim | 14 | 双臂 7 关节 × 2 |
| state_dim | 14 | 与 action 维度一致 |
| training | action | 仅动作训练(不含视频辅助 loss) |
| future_loss_type | flow_matching | Flow Matching 动作预测头 |
| future_action_window_size | 7 | 预测未来 7 步动作 |
| action_horizon | 8 | 动作时间窗口 |
| obs_horizon | 16 | 观测时间窗口 |
Robotwin 是 StarVLA 团队发布的双臂操作数据集,包含 50 个桌面操作任务(拿取、放置、推拉等)。每个 episode 包含 2~3 个视角的视频 + 双臂关节角序列,与 DiT4DiT 的多视角视频输入天然匹配。
六、ZeRO-2 vs ZeRO-3 性能对比
在 Robotwin 数据集上,同等 batch size(等效 128)下的对比:
| ZeRO-3 | ZeRO-2 | |
|---|---|---|
| 训练速度 | ~3.14s/it | ~0.74s/it |
| per_device_batch_size | 16 | 2 |
| gradient_accumulation | 1 | 8 |
| 等效 batch | 128 (16×8×1) | 128 (2×8×8) |
| 单卡显存占用 | ~32GB | ~42GB |
| 数据吞吐量 | 更高(大 batch,少 accum) | 较低(小 batch × 多 accum) |
| 模型保存 | save_checkpoint(分片) | 正常 |
| 配置复杂度 | 较高(多个兼容性问题) | 较低 |
| 推荐场景 | 正式训练(吞吐优先) | 调试 / 快速验证 |
为什么 ZeRO-3 单步慢但总体吞吐高? ZeRO-3 每步 3.14s 看着比 ZeRO-2 的 0.74s 慢很多,但 ZeRO-3 每步处理 128 个样本(batch=16×8 卡),ZeRO-2 每步只处理 16 个样本(batch=2×8 卡)。换算成单样本耗时:ZeRO-3 = 3.14/128 = 0.025s/sample,ZeRO-2 = 0.74/16 = 0.046s/sample。ZeRO-3 的实际吞吐量是 ZeRO-2 的 1.8 倍。
如果 910B 的 64GB 显存对你来说足够大(batch=2 就能占满),那 ZeRO-2 的绝对训练时长更短(150000 步 × 0.74s = 30.8h vs 150000 步 × 3.14s = 130.8h)。但如果你想跑更大的 batch(对 DiT 训练收敛有帮助),ZeRO-3 是唯一选择。
七、踩坑记录
以下是适配过程中遇到的 12 个典型问题,按出现顺序记录。
7.1 num_processes 必须匹配物理卡数
现象:train_8p.sh 中设置 --num_processes=16,启动后 HCCL 报错 hcom init failed: rank 8 assigned to same device as rank 0。
原因:accelerate 的 --num_processes 指的是进程数,不是卡数。当进程数 > 物理卡数时,accelerate 会循环分配(rank 0→卡0, rank 8→卡0),导致多进程抢占同一张卡。HCCL 不允许同一张卡上有多个通信端点。
修复:--num_processes 必须等于 npu-smi info 显示的物理卡数(这里是 8)。如果你在踩坑记录里看到别的 StarVLA README 里写 --num_processes=16,那是因为他们用了 16 卡环境。
7.2 mx_driving.patcher API 变更
现象:从 StarVLA patch 复制了 patcher 调用代码,运行时报 ImportError: cannot import name 'TransformersNPU' from 'mx_driving.tools.patch'。
原因:StarVLA 的 starvla.patch 基于旧版 mx_driving(≤0.9.x),使用 Patcher().add(TransformersNPU).apply() 模式。新版(≥1.0.20260421)重构了架构,TransformersNPU 类已被移除,功能整合进 default_patcher_builder.build().__enter__()。
修复:替换为 default_patcher_builder.build().__enter__()。
7.3 libc10.so / libtorch_npu.so 未自动发现
现象:import torch_npu 成功,但训练启动时报 OSError: libc10.so: cannot open shared object file。
原因:pip 安装的 torch_npu .whl 将 .so 文件放在 site-packages 目录下,但 conda 环境不会自动将该路径加入 LD_LIBRARY_PATH。这和 CUDA 的 pip 安装行为一致,不过在 NPU 集群上更容易忽略(因为多数用户习惯了 conda 安装的 torch 会自动配置好路径)。
修复:见 2.9 节。
7.4 Cosmos-Predict2.5-2B 显存占用
现象:加载模型后显存直接占满 64GB,无法启动训练。
原因:Cosmos2.5 包含 VAE(~80M)+ Transformer(~2B)+ Text Encoder(T5-XXL,~4.7B)。Text Encoder 的 T5-XXL 虽然不参与训练,但 PyTorch 默认加载全部参数到显存。三部分加起来约 7B 参数,bf16 下 ~14GB 权重 + 优化器状态和中间激活会把 64GB 吃满。
修复:冻结 VAE + Text Encoder(见 3.3 节),显存从 ~58GB 降至 ~42GB。
7.5 VAE / Text Encoder 混合 dtype → ZeRO-3 defragment 报错
现象:ZeRO-3 训练启动时在 defragment 阶段报 RuntimeError: expected scalar type BFloat16 but found Float。
原因:VAE 和 Text Encoder 权重为 fp32(HuggingFace 默认格式),Transformer 为 bf16。ZeRO-3 初始化时遍历所有 trainable 参数做内存预分配,要求 dtype 一致。混合 dtype 触发了 defragment 的类型检查。
修复:冻结 VAE 和 Text Encoder,使其不进入 ZeRO-3 的参数管理范围。注意不能简单地将它们 cast 到 bf16——VAE 的视频编解码对 fp32 精度有硬性要求,cast 后解码画面会出现明显的颜色失真。
7.6 eval_action_model 维度不匹配
现象:训练中途评估时报 RuntimeError: The size of tensor a (16) must match the size of tensor b (8) at non-singleton dimension 1。
原因:训练阶段 dataloader 返回的 actions 和 action_mask 已按 action_horizon(8)截断,但 eval 阶段走的是一条不同的代码路径,直接用了原始的 obs_horizon(16)维度的数据。
修复:在 eval_action_model 函数中添加与训练阶段一致的截断逻辑。
7.7 bf16 → numpy 不支持
现象:eval 阶段 predict_action() 输出报 TypeError: Got unsupported ScalarType BFloat16。
原因:见 3.2 节 ③。
修复:.float().cpu().numpy()。
7.8 @torch.inference_mode() 与 ZeRO-3 冲突
现象:ZeRO-3 模式下 eval 报 AssertionError: tensor has no storage。
原因:见 3.2 节 ④。
修复:@torch.no_grad()。
7.9 ZeRO-3 模型保存卡死
现象:训练到 checkpoint 保存步骤时,rank 0 内存飙升然后 OOM,其余 7 张卡处于等待状态。
原因:accelerator.get_state_dict() 在 ZeRO-3 下会把 8 卡各自持有的参数分片 gather 到 rank 0,相当于 rank 0 需要加载完整的 2.3B 参数。gather 过程通过 HCCL allgather 实现,在 64GB 的卡上 gather 后 rank 0 显存占用会翻倍,超过 64GB 限制。
修复:改用 model.save_checkpoint(),每卡各自写自己的分片到磁盘,不跨卡通信。但保存后的 checkpoint 是分片格式,加载时需要相同的 ZeRO-3 配置。
7.10 accelerate config 中的 mixed_precision 与 DeepSpeed bf16 冲突
现象:训练启动时报 ValueError: mixed_precision is set to 'bf16' in both accelerate config and DeepSpeed config。
原因:accelerate 的默认配置文件和 DeepSpeed 的 ds_config 都声明了 bf16,冲突。
修复:accelerate config 中设置 mixed_precision: "no",bf16 完全交由 DeepSpeed 管理。因为 ZeRO 的混合精度训练有自己的 cast 逻辑,accelerate 的 autocast 层会与之冲突。
7.11 视频解码 decord 的 CUDA 默认设备
现象:dataloader 第一次取数据时报 RuntimeError: CUDA error: no CUDA-capable device is detected。
原因:decord 的 VideoReader 默认尝试使用 GPU 解码(ctx=decord.gpu(0)),在 NPU 环境下 GPU 不存在所以报错。decord 目前没有 NPU 解码支持。
修复:在 dataloader 创建 VideoReader 时指定 CPU 解码:
vr = decord.VideoReader(video_path, ctx=decord.cpu(0))
CPU 解码在 8 卡训练场景下会成为瓶颈(占用 DataLoader worker 的 CPU 时间),但鉴于 910B 的单步训练耗时在秒级,视频解码的毫秒级开销可以忽略。
7.12 HCCL 通信超时(偶发)
现象:训练执行到某个 allreduce 时卡住,最终 hccl timeout after 1800s。
原因:偶发性问题,通常与集群网络抖动或节点负载有关。HCCL 默认超时 1800s,对于正常的 gradient allreduce(几百 MB 的数据)来说远远足够,超时基本就是网络层的问题而非代码。
排查步骤:
npu-smi info -t topo检查卡间互联是否正常hccl_test跑一次 allreduce 基准测试- 检查 /var/log/npu/slog 中的 HCCL 警告
- 如果持续复现,降低 bucket_size(比如从 5e8 降到 1e8),减小单次通信量
八、适配 Patch
基于 commit 1ae6efd 生成的两个 patch 文件,已覆盖全部修改:
- dit4dit_zero3.patch:ZeRO-3 适配版(含 NPU 移植 + 全部 Bug 修复)
- dit4dit_zero2.patch:ZeRO-2 适配版(同上 + train.py 去硬编码 + ZeRO-2 配置)
使用方法:
git clone https://github.com/Mondo-Robotics/DiT4DiT.git
cd DiT4DiT
git checkout 1ae6efd
git apply dit4dit_zero3.patch # 或 dit4dit_zero2.patch
pip install -r requirements.txt
pip install -e .
两个 patch 下载链接:待上传至网盘或服务器公开目录。
九、总结
DiT4DiT 的 NPU 适配整体比预期顺利。得益于三个因素:
- DiT4DiT 没有深度绑定 CUDA 的自定义算子——所有核心计算都是 PyTorch 高层 API,torch_npu 已经做了适配
- StarVLA 的 NPU 适配可以作为模板——训练框架一模一样,只需处理 backbone 的差异
- 华为 mx_driving patcher 成熟度高——隐式 CUDA 调用基本被拦截修复,不需要手写算子映射
真正耗时的是 ZeRO-3 的兼容性踩坑(7 个问题里有 5 个跟 ZeRO-3 相关)。建议先在 ZeRO-2 上快速验证,确认训练流程无误后再切 ZeRO-3。如果显存够用,ZeRO-2 的绝对训练时长更短,省下的卡时可以多跑几组实验。
下一步计划:将 patch 贡献回 DrivingSDK,补充 Robotwin 数据集上的完整评测结果(当前仅验证了 loss 收敛,未跑测评流程)。
简记。








COMMENTS | NOTHING