FLUX GRPO 文生图强化学习训练 NPU 适配实战:基于 DanceGRPO 与 HPSv2 的 Atlas A3 训练指南
【免费下载链接】cann-recipes-train本项目针对LLM与多模态模型训练业务中的典型模型、加速算法,提供基于CANN平台的优化样例项目地址: https://gitcode.com/cann/cann-recipes-train
本篇指南以 cann-recipes-train 仓库multimodal_rl/flux_grpo样例为主线,系统讲解如何在昇腾 Atlas A3 系列 NPU 上完成 FLUX 文生图模型的 GRPO 强化学习训练。样例参考 DanceGRPO 的训练流程,以 HPSv2 作为奖励模型优化生成图片质量,通过两枚适配补丁将 CUDA/NCCL 训练链路迁移到 NPU/HCCL,并补齐 diffusers 中 FLUX 相关算子的 NPU 实现。读完本文,你将掌握 NPU 环境构建、上游源码补丁应用、模型与数据准备、8 die / 16 die 训练脚本运行以及常见问题排查的完整实战方法。
1. 样例概览:将 GRPO 强化学习引入文生图模型
FLUX 是开源生态中代表性的一类文生图扩散模型(以 FLUX.1-dev 权重为基础)。传统 SFT 微调只能让模型模仿数据分布,而 GRPO(Group Relative Policy Optimization)强化学习可以直接针对「生成图片与提示词的匹配质量」这类难以用交叉熵表达的目标进行优化,其奖励信号由奖励模型提供。
本样例参考 DanceGRPO 的 FLUX GRPO 训练流程,在昇腾 NPU 上完成以下四类适配:
- 训练链路设备迁移:将训练与预处理流程从 CUDA/NCCL 适配为 NPU/HCCL,包括
torch.device("cuda")→"npu"、torch.cuda.set_device→torch.npu.set_device、dist.init_process_group("nccl")→"hccl"、torch.autocast("cuda")→torch.autocast("npu")等。 - diffusers 算子级适配:为 diffusers 中 FLUX 相关算子补充 NPU 支持,包括新增
NpuFusedRMSNorm融合归一化算子与 NPU 上的 rotary position embedding(RoPE)处理。 - NPU 训练脚本与环境变量:新增 Atlas A3 16 die 训练启动脚本,并配置
TASK_QUEUE_ENABLE、COMBINED_ENABLE、PYTORCH_NPU_ALLOC_CONF等 NPU 训练环境变量。 - 训练 batch 组织优化:新增
rollout_batch_size、train_micro_batch_size训练参数,把 rollout 采样阶段与训练阶段(含梯度累积)的 batch 切分逻辑独立出来,便于在显存受限时精细调控。
需要特别说明的限制:当前默认使用 HPSv2 作为奖励模型;use_pickscore参数暂未适配 NPU,开启后不会生效(训练代码中会打印 Warning),因此训练脚本建议保持--use_hpsv2。
在仓库根目录 README.md 中,本样例被归入「🎨 多模态强化学习」类别,定位为「基于 FLUX GRPO,覆盖文生图模型 GRPO 训练、HPSv2 奖励优化及 NPU 适配」,是 cann-recipes-train 在 LLM 强化学习(llm_rl/)之外的第一个多模态强化学习样例。
2. 补丁结构与源码级适配原理
2.1 补丁文件说明
当前目录(multimodal_rl/flux_grpo)仅保存适配补丁,不直接包含上游框架源码。使用时需要先准备 DanceGRPO 和 diffusers 源码,再将补丁应用到对应仓库。
| 文件路径 | 说明 |
|---|---|
| multimodal_rl/flux_grpo/patches/DanceGRPO.patch | 适配 DanceGRPO/FastVideo 工程,使能 NPU 预处理、FLUX GRPO 训练、HPSv2 奖励计算和 A3 16 die 启动脚本 |
| multimodal_rl/flux_grpo/patches/diffusers.patch | 适配 diffusers FLUX 相关模块,增加NpuFusedRMSNorm、torch_npu导入和 NPU 上的 RoPE 处理 |
两枚补丁均采用标准git format-patch格式(含From ... Mon Sep 17 00:00:00头部),因此推荐使用git am应用,可以保留提交信息并便于后续git am --continue处理冲突。仓库的 ci/check_patch_names.sh 中维护了一套补丁命名规范(NNNN-(bugfix|feature)-description.patch格式、序号从 0001 连续递增、描述用下划线),本文样例的补丁也遵循这一工程约定。
2.2 设备与通信栈切换:CUDA/NCCL → NPU/HCCL
DanceGRPO.patch 的第 4 个提交「更新代码以支持 NPU」集中完成了设备栈替换,涉及三个核心文件:
- 数据预处理:preprocess_flux_embedding.py 对应补丁段 将
torch.device("cuda" if torch.cuda.is_available() else "cpu")改为torch.device("npu" if torch.npu.is_available() else "cpu"),并把torch.cuda.set_device(local_rank)、backend="nccl"同步替换为torch.npu.set_device(local_rank)、backend="hccl"。 - GRPO 训练主程序:train_grpo_flux.py 对应补丁段 在文件头
import torch_npu,main()中dist.init_process_group("hccl")、torch.npu.set_device(local_rank)、device = torch.npu.current_device();所有torch.autocast("cuda", torch.bfloat16)统一改为torch.autocast("npu", torch.bfloat16);并将torch.backends.cuda.matmul.allow_tf32 = True替换为torch_npu.npu.matmul.allow_hf32 = True,启用 NPU 的 HF32 混合精度矩阵乘。 - FSDP 工具:fsdp_util.py 对应补丁段 将两处
torch.cuda.current_device()改为torch.npu.current_device(),保证 FullyShardedDataParallel 在 NPU 上正确绑定设备。
此外,load.py 对应补丁段 注释掉了 HunyuanVideo、Mochi 等无关模型模块的导入路径,get_no_split_modules()仅保留 FLUX 分支(FluxTransformer2DModel→FluxTransformerBlock, FluxSingleTransformerBlock),既减少不必要的依赖加载,也确保 FSDP 自动切分策略只针对 FLUX Transformer 生效。
2.3 diffusers 算子级适配:NpuFusedRMSNorm 与 RoPE
diffusers.patch 由两个提交组成,改动集中在src/diffusers/models/下的三个文件:
新增NpuFusedRMSNorm(normalization.py):补丁在 normalization.py 中新增了一个 NPU 融合 RMSNorm 模块:
class NpuFusedRMSNorm(torch.nn.Module): def __init__(self, dim, eps: float = 1e-6): super().__init__() self.weight = nn.Parameter(torch.ones(dim)) self.eps = eps def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: return torch_npu.npu_rms_norm(hidden_states.to(self.weight.dtype), self.weight, epsilon=self.eps)[0]核心是调用 torch-npu 提供的torch_npu.npu_rms_norm融合算子(FusedRMSNorm),将归一化运算下沉到 NPU 融合算子实现,避免逐元素 Python 循环,并显式import torch_npu(第二个提交)保证算子注册表可用。
FLUX Attention 的 q/k 归一化替换(attention_processor.py):在 attention_processor.py 中,Attention模块所有qk_norm == "rms_norm"与"rms_norm_across_heads"分支的RMSNorm均替换为NpuFusedRMSNorm(包括norm_q、norm_k、norm_added_q、norm_added_k)。FLUX 的联合注意力(joint attention)正是依赖这一类 q/k RMSNorm,这是适配的核心路径。
RoPE 的 NPU 处理(embeddings.py):embeddings.py 包含两处修改:
get_1d_rotary_pos_embed中,FLUX 分支的 cos/sin 频率计算由freqs.cos().repeat_interleave(2, dim=1)调整为freqs.cos().T.repeat_interleave(2, dim=0).T.contiguous(),即在转置维度上做 repeat 后再转置并显式contiguous(),保证 NPU 上的内存排布与算子要求一致;FluxPosEmbed前向中,freqs_dtype的判断由is_mps扩展为is_mps or is_npu,使 NPU 上同样使用 float32 而非 float64 计算位置编码频率,规避精度与性能开销。
从这些改动可以看出,NPU 适配不止是设备名替换,还包括算子融合(RMSNorm)、内存排布(contiguous)与数据类型(float32 频率)等与硬件算子行为强相关的细节。
2.4 rollout 与训练阶段 batch 组织:新增两个训练参数
原版脚本以「每设备 batch_size=1」组织 rollout,显存利用率低。补丁在 train_grpo_flux.py 参数区 新增两个参数:
parser.add_argument("--rollout_batch_size", type=int, default=1, help="Batch size (per device) for the rollout sampling.") parser.add_argument("--train_micro_batch_size", type=int, default=1, help="Micro batch size (per device) in training.")- rollout_batch_size:控制
sample_reference_model中每个采样微步处理的 batch 大小。补丁将batch_size = 1改为batch_size = args.rollout_batch_size,初始潜变量通过.repeat(batch_size, 1, 1, 1)广播,图像 ID 按 batch 展开,VAE 解码后的图片按flux_{rank}_{idx}.png逐张保存(对应补丁段)。 - train_micro_batch_size:在
train_one_step中替代原先的「dict of lists → list of dicts」转换逻辑,改为按range(0, samples['latents'].shape[0], train_mbs)直接对 tensor 切片做微批训练(对应补丁段),并将梯度累积判定从(i+1) % gradient_accumulation_steps == 0改为(i // train_mbs + 1) % gradient_accumulation_steps == 0,使梯度累积与微批边界严格对齐;同时 rank 0 打印的 reward / advantage 改为sample["rewards"].mean().item()、sample["advantages"].mean().item()的批内均值。
两者的配合关系:rollout_batch_size × 设备数决定每轮 GRPO 采样的样本组规模(组内生成num_generations张图用于组内相对奖励计算),train_micro_batch_size × gradient_accumulation_steps × 设备数决定训练阶段实际生效的全局 batch(GBS)。
2.5 奖励模型适配:HPSv2 生效、PickScore 暂不支持
训练默认使用 HPSv2(Human Preference Score v2)作为奖励模型,它基于 CLIP 图像/文本编码器输出对齐分数,衡量生成图片与提示词的人类偏好一致程度。补丁的适配点:
- 权重加载路径写死为
./hps_ckpt/HPS_v2.1_compressed.pt,torch.load(cp, map_location=f'npu:{device}')加载到 NPU(对应补丁段); - 奖励打分循环迁移到
with torch.amp.autocast('npu'):下执行,image_features @ text_features.T得到 logits,取torch.diagonal作为逐样本 HPS 分数,batch 化处理多个 rollout 样本(对应补丁段); use_pickscore分支被整体替换为main_print("[Warning] 'use_pickscore' is not supported yet, ...")的警告输出,参数帮助文本同步标注了该限制。
2.6 NPU 训练环境变量
补丁为 8 die / 16 die 训练脚本与预处理脚本统一注入了一组 NPU 环境变量(以 16 die 脚本 为例),运行前应了解其作用并可按需调整:
| 环境变量 | 脚本中的值 | 说明 |
|---|---|---|
ASCEND_RT_VISIBLE_DEVICES | 0..15(16 die) | 指定训练使用的 NPU 编号列表 |
TASK_QUEUE_ENABLE | 2 | 开启任务队列(Task Queue)特性,提升算子下发效率 |
COMBINED_ENABLE | 1 | 使能通信与计算融合/组合优化 |
CPU_AFFINITY_CONF | 2 | 配置进程 CPU 亲和策略 |
HCCL_CONNECT_TIMEOUT | 1200 | HCCL 建链超时(秒),大集群下建议调大 |
NPU_ASD_ENABLE | 0 | 关闭芯片自检(ASD)相关逻辑,避免干扰训练 |
ASCEND_LAUNCH_BLOCKING | 0 | 关闭算子同步阻塞,保持异步下发 |
ACLNN_CACHE_LIMIT | 100000 | ACLNN 算子编译缓存上限 |
MULTI_STREAM_MEMORY_REUSE | 2 | 多流内存复用策略 |
PYTORCH_NPU_ALLOC_CONF | expandable_segments:True | 使能 NPU 显存的可扩展段分配,缓解碎片化 |
HCCL_BUFFSIZE | 800 | HCCL 通信缓冲区大小(MB) |
3. 支持的产品型号与脚本说明
本样例支持Atlas A3 系列产品。补丁将原仓scripts/finetune/finetune_flux_grpo_8gpus.sh适配为 8 NPU die 训练脚本,并额外提供 16 NPU die 训练脚本scripts/finetune/finetune_flux_grpo_a3_16die.sh。
由于不同机器的 NPU 编号、CANN 安装路径、通信网卡和启动方式可能不同,请务必根据实际环境修改脚本中的环境变量(ASCEND_RT_VISIBLE_DEVICES、source .../set_env.sh路径等)和torchrun参数(--nnodes、--nproc_per_node、--master_addr、--master_port)。
4. 环境准备
4.1 方式一:Docker 构建环境
请使用已安装 CANN、驱动和torch-npu==2.7.1的镜像,或基于昇腾官方镜像自行安装对应版本依赖。以下命令以 Atlas A3 16 die 为例,${IMAGE_NAME}、${HOST_WORKSPACE}和${YOUR_CONTAINER_NAME}请替换为实际值:
docker run \ --device=/dev/davinci0 --device=/dev/davinci1 --device=/dev/davinci2 --device=/dev/davinci3 \ --device=/dev/davinci4 --device=/dev/davinci5 --device=/dev/davinci6 --device=/dev/davinci7 \ --device=/dev/davinci8 --device=/dev/davinci9 --device=/dev/davinci10 --device=/dev/davinci11 \ --device=/dev/davinci12 --device=/dev/davinci13 --device=/dev/davinci14 --device=/dev/davinci15 \ --device=/dev/davinci_manager --device=/dev/devmm_svm --device=/dev/hisi_hdc \ -v /etc/localtime:/etc/localtime \ -v /usr/local/dcmi:/usr/local/dcmi \ -v /usr/local/Ascend/driver:/usr/local/Ascend/driver \ -v /usr/local/bin/npu-smi:/usr/local/bin/npu-smi \ -v ${HOST_WORKSPACE}:${HOST_WORKSPACE} \ -w ${HOST_WORKSPACE} \ --shm-size=100g \ --privileged=true \ -itd \ --net=host \ --name ${YOUR_CONTAINER_NAME} \ ${IMAGE_NAME} \ /bin/bash docker exec -it ${YOUR_CONTAINER_NAME} bash要点说明:16 个--device=/dev/davinciN逐一映射 NPU 设备;--shm-size=100g为分布式通信的共享内存预留空间;--net=host保证 torchrun 多机通信端口可达;挂载/usr/local/Ascend/driver、/usr/local/dcmi与npu-smi是 NPU 容器运行的标准配置。
4.2 方式二:Conda 构建环境
conda create -n flux_grpo python=3.11 -y conda activate flux_grpo安装torch-npu==2.7.1(该版本与补丁中pyproject.toml锁定的torch==2.7.1+torch-npu==2.7.1一致):
pip install torch==2.7.1 torch-npu==2.7.15. 源码准备与补丁应用
5.1 下载本仓库源码
cd path-to-cann-recipes-train/ git clone https://gitcode.com/cann/cann-recipes-train.git本案例内容位于path-to-cann-recipes-train/cann-recipes-train/multimodal_rl/flux_grpo。下文以${FLUXGRPO_PATH}指代该目录。
5.2 准备 DanceGRPO 源码并应用补丁
${FLUXGRPO_PATH}请替换为所使用的本案例地址;如果源码已经下载,也可以直接进入对应目录执行git am。应用补丁后安装相关环境:
cd "${FLUXGRPO_PATH}" git clone https://github.com/XueZeyue/DanceGRPO.git cd DanceGRPO git checkout 15cc71d git am "${FLUXGRPO_PATH}"/patches/DanceGRPO.patch bash env_setup.sh cd ..为什么必须git checkout 15cc71d:补丁是基于 DanceGRPO 特定提交生成的上下文,git am对上下文匹配敏感。若上游代码已有较大变化会应用失败,此时需要手动解决冲突后执行git am --continue(详见第 9 节常见问题)。
补丁对 DanceGRPO 工程的依赖做了系统性替换(env_setup.sh / pyproject.toml 补丁段),核心版本锁定如下:
- 训练栈:
torch==2.7.1、torch-npu==2.7.1、diffusers==0.32.0、transformers==4.46.1、accelerate==1.9.0、bitsandbytes==0.46.1、peft==0.13.2、liger_kernel==0.4.1; - 通用工具:
packaging==25.0、ninja==1.11.1.4、ml-collections==1.1.0、absl-py==2.3.1、inflect==6.0.4、huggingface_hub==0.34.0、protobuf==3.20.0、pybind11==3.0.4(为 aarch64 编译 decord 等服务)。
同时env_setup.sh不再安装 CUDA 版 torch 与 flash-attn,.gitignore追加data、hps_ckpt、images、kernel_meta、wandb等目录,避免训练产物污染 Git 工作区。
5.3 准备并安装 diffusers
DanceGRPO.patch依赖diffusers==0.32.0,需先将 diffusers 切换到对应版本再应用补丁:
git clone https://github.com/huggingface/diffusers.git cd diffusers git checkout v0.32.0 git am "${FLUXGRPO_PATH}"/patches/diffusers.patch pip uninstall diffusers pip install -e . cd ..注意pip uninstall diffusers是为了移除可能已安装的发布版,随后以 editable 方式安装补丁后的源码版本。
5.4 准备 HPSv2 奖励模型依赖
训练脚本默认开启--use_hpsv2,需要安装 HPSv2 并将权重放到hps_ckpt/HPS_v2.1_compressed.pt:
git clone https://github.com/tgxs002/HPSv2.git cd HPSv2 git checkout 866735ecaae999fa714bd9edfa05aa2672669ee3 pip install -e . cd ..5.5 安装 decord(aarch64 环境)
aarch64 版本无法直接pip install decord==0.6.0,推荐从源码编译安装:
mkdir ffmpeg-decord cd ffmpeg-decord # 安装依赖包 ffmpeg wget https://ffmpeg.org/releases/ffmpeg-4.0.1.tar.bz2 --no-check-certificate tar -xvf ffmpeg-4.0.1.tar.bz2 mv ffmpeg-4.0.1 ffmpeg cd ffmpeg ./configure --enable-shared make -j 64 make install cd .. # 安装 decord git clone --recursive https://github.com/dmlc/decord cd decord if [ -d build ];then rm -rf build;fi && mkdir build && cd build cmake .. -DUSE_CUDA=0 -DCMAKE_BUILD_TYPE=Release make -j 64 make install cd ../python pip install -e . cd ..编译要点:FFmpeg 必须以--enable-shared编译(否则运行时找不到libswscale.so.5等动态库);decord 使用-DUSE_CUDA=0关闭 CUDA 支持,适配纯 CPU/NPU 环境。
6. 数据集与模型权重准备
6.1 准备 FLUX 模型权重
训练脚本默认从data/flux加载模型和 VAE,请将所需 FLUX 权重下载到该目录:
cd "${FLUXGRPO_PATH}"/DanceGRPO mkdir -p data/flux # 根据模型许可和实际网络环境下载 FLUX 权重到 data/flux权重来源可参考 huggingface 的black-forest-labs/FLUX.1-dev或 ModelScope 同名模型仓库(modelscope.cn/models/black-forest-labs/FLUX.1-dev),请按模型许可与网络环境自行选择。
6.2 准备 HPS 和 CLIP 权重
训练脚本默认从hps_ckpt加载,请将 HPS 评分权重与 CLIP 编码器权重下载到该目录:
mkdir -p hps_ckpt cd hps_ckpt # HPS_v2.1_compressed.pt 下载方式 1 wget https://huggingface.co/xswu/HPSv2/resolve/main/HPS_v2.1_compressed.pt?download=true # HPS_v2.1_compressed.pt 下载方式 2 wget https://www.modelscope.cn/models/AI-ModelScope/HPSv2/resolve/master/HPS_v2.1_compressed.pt # open_clip_pytorch_model.bin 下载方式 1 wget https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K/resolve/main/open_clip_pytorch_model.bin?download=true # open_clip_pytorch_model.bin 下载方式 2 wget https://www.modelscope.cn/models/laion/CLIP-ViT-H-14-laion2B-s32B-b79K/resolve/master/open_clip_pytorch_model.bin cd ..文件用途:HPS_v2.1_compressed.pt是 HPSv2 奖励模型的压缩权重(训练时由train_grpo_flux.py直接torch.load加载);open_clip_pytorch_model.bin是 CLIP-ViT-H-14 的 open_clip 格式权重,供 HPSv2 内部的 CLIP 编码器使用。
6.3 准备训练 prompt 数据,生成 FLUX RL embeddings
预处理脚本默认读取assets/prompts.txt中的内容作为 prompts(逐行一条),通过 FLUX 的文本编码器(T5 与 CLIP)生成文本 embeddings,默认生成位置和训练脚本默认读取位置为data/rl_embeddings/videos2caption.json:
# 如需修改 NPU 数量、CANN 路径或模型路径,请先编辑脚本 bash scripts/preprocess/preprocess_flux_rl_embeddings.sh预处理脚本被适配为 NPU 版本(对应补丁段):GPU_NUM改为NPU_NUM,torchrun --nproc_per_node=$NPU_NUM驱动fastvideo/data_preprocess/preprocess_flux_embedding.py,脚本内置了与训练脚本一致的环境变量与data/flux、data/rl_embeddings默认路径。该步骤是训练前必须完成的离线步骤——GRPO 训练阶段直接消费预计算好的 embeddings,不再重复文本编码。
7. 运行 GRPO 训练
在 DanceGRPO/FastVideo 源码根目录下执行训练脚本。运行前请重点检查脚本中的如下配置:
source /usr/local/Ascend/cann/set_env.sh:根据实际 CANN 安装路径修改;ASCEND_RT_VISIBLE_DEVICES:根据实际使用的 NPU 编号修改;--pretrained_model_name_or_path、--vae_model_path:默认指向data/flux;--data_json_path:默认指向data/rl_embeddings/videos2caption.json;--use_hpsv2:默认使用 HPSv2 奖励模型,需要提前准备hps_ckpt/HPS_v2.1_compressed.pt。
7.1 16 NPU die 训练示例
bash scripts/finetune/finetune_flux_grpo_a3_16die.sh该脚本(由补丁全新创建,完整内容见补丁)以torchrun --nnodes=1 --nproc_per_node=16启动,核心训练参数如下,可直接作为 Atlas A3 16 die 的基线配置:
| 参数 | 值 | 说明 |
|---|---|---|
--seed | 42 | 全局随机种子 |
--rollout_batch_size | 4 | 每设备 rollout 采样 batch |
--train_micro_batch_size | 1 | 训练微批大小 |
--train_batch_size | 1 | 每设备(训练)batch,与梯度累积配合 |
--gradient_accumulation_steps | 4 | 梯度累积步数,与 micro batch 对齐 |
--max_train_steps | 300 | 最大训练步数 |
--learning_rate/--weight_decay | 1e-5 / 0.0001 | 优化器超参 |
--mixed_precision | bf16 | 混合精度训练 |
--num_generations | 12 | 每组 prompt 生成的图片数(组内相对奖励) |
--h/--w/--t | 720 / 720 / 1 | 生成图像分辨率与帧数 |
--sampling_steps/--eta | 16 / 0.3 | rollout 采样步数与随机性参数 |
--timestep_fraction | 0.6 | 训练时采样的时间步比例 |
--cfg | 0.0 | 关闭 classifier-free guidance |
--clip_range/--adv_clip_max | 1e-4 / 5.0 | 优势裁剪范围 |
--shift/--use_group/--ignore_last/--init_same_noise | 3 / 开 / 开 / 开 | 噪声调度偏移、组奖励、GRPO 细节控制 |
--checkpointing_steps | 40 | 每 40 步保存 checkpoint |
--gradient_checkpointing | 开 | 激活重计算以省显存 |
7.2 8 NPU die 训练示例
bash scripts/finetune/finetune_flux_grpo_8gpus.sh该脚本由原仓finetune_flux_grpo_8gpus.sh适配而来,torchrun --nproc_per_node=8 --master_port 19002,环境变量与 16 die 版本一致,适合单机 8 die 环境;其训练参数(如--train_batch_size 2)与 16 die 版本略有差异,可对比阅读选择。
7.3 多机训练
多机训练可参考scripts/finetune/finetune_flux_grpo.sh(原多机脚本),并根据实际集群修改--nnodes、--node_rank、--master_addr和--master_port,同时确认各节点HCCL_CONNECT_TIMEOUT充足(脚本中默认 1200 秒)。
7.4 训练输出
训练输出默认保存在如下目录:
data/outputs/grpo:训练 checkpoint 和输出结果;images:rollout 过程中生成的图片样例(命名格式flux_{rank}_{idx}.png,便于按 rank 与样本索引回溯生成质量);wandb:如开启 wandb,则保存日志相关文件(脚本默认WANDB_DISABLED=true,也可按需配置WANDB_BASE_URL/WANDB_MODE)。
8. 训练效果与性能参考
文章开头的图片展示了训练过程中 reward 指标的变化趋势(横轴为 Training Iterations,纵轴为 Reward,奖励值从约 0.30 波动上升至接近 0.40),可用于判断 GRPO 优化是否收敛。
在 Atlas A3 16 die 上,设置train_micro_batch_size=1、rollout_batch_size=4、gradient_accumulation_steps=4,对应全局训练批次大小(train GBS)为1 × 4 × 16 = 64。该配置下的实测训练耗时如下:
| 训练迭代数 | 训练耗时 |
|---|---|
| 200 iterations | 17.22 小时 |
| 300 iterations | 25.83 小时 |
由于硬件平台、软件栈及具体运行环境存在差异,以上数据用于展示本样例在 Atlas A3 上的实际训练性能,不作为严格的同条件性能对比。
9. 常见问题与排查
git am失败:请确认上游源码版本与补丁生成版本匹配(DanceGRPO 需git checkout 15cc71d,diffusers 需git checkout v0.32.0);若上游代码已有较大变化,需要手动解决冲突后继续执行git am --continue。找不到 HPSv2 权重:请确认文件路径为
hps_ckpt/HPS_v2.1_compressed.pt(训练代码中写死该相对路径),或在训练代码中修改cp路径。更换 CANN 或 torch-npu 版本后出现编译缓存问题:建议清理如下目录后重新运行:
rm -rf kernel_meta rm -rf .cache rm -rf /root/.cache--use_pickscore不生效:当前该参数暂未适配 NPU,训练脚本建议保持使用--use_hpsv2(开启 pickscore 只会打印 Warning,不会报错,但也不会产生 PickScore 奖励)。GCC 版本问题:可在 conda 环境中安装指定版本 GCC(以 aarch64 和 GCC11 为例):
conda install -c conda-forge gcc_linux-aarch64=11 gxx_linux-aarch64=11 export CC="$CONDA_PREFIX/bin/aarch64-conda-linux-gnu-gcc" export CXX="$CONDA_PREFIX/bin/aarch64-conda-linux-gnu-g++"OSError: libswscale.so.5: cannot open shared object file:表示运行时找不到 FFmpeg 的动态库libswscale.so.5。先执行find /usr -name libswscale.so.5(通常会出现在/usr/local/lib的子目录下),再执行export LD_LIBRARY_PATH=/usr/local/lib:$LD_LIBRARY_PATH,然后运行程序即可。ImportError: .../libGLdispatch.so.0: cannot allocate memory in static TLS block:推荐预加载libGLdispatch。先定位系统libGLdispatch.so.0:find /usr -name libGLdispatch.so.0,若出现/usr/lib/aarch64-linux-gnu/libGLdispatch.so.0,执行export LD_PRELOAD=/usr/lib/aarch64-linux-gnu/libGLdispatch.so.0,随后运行程序即可。
10. 总结
本文基于 cann-recipes-train 的multimodal_rl/flux_grpo样例,完整梳理了 FLUX 文生图模型 GRPO 强化学习在昇腾 Atlas A3 NPU 上的落地路径:
- 适配思路清晰可复用:从设备栈(CUDA→NPU)、通信栈(NCCL→HCCL)、混合精度(TF32→HF32)到算子层(
NpuFusedRMSNorm融合归一化、RoPE 内存排布与 float32 频率),这套「训练链路 + 算子 + 脚本参数」三层适配方法论同样适用于其他基于 diffusers/FastVideo 的文生图、文生视频模型迁移; - 工程化实践完整:Docker / Conda 双路径环境构建、补丁应用(
git am+ 固定上游 commit)、HPSv2 奖励权重准备、RL embeddings 离线预处理、8 die / 16 die 训练脚本运行、性能基线(200 iterations 约 17.22 小时)与 7 类常见问题排障,均可直接照做。
若希望进一步了解 CANN 在预训练、LLM RL、SFT 等场景的更多原子化能力与优化特性,可继续阅读仓库 docs/cann_capability_overview.md 与其他样例目录(如 llm_rl/qwen3/verl-mindspeed/README.md),将本样例的 NPU 适配经验迁移到更大规模的训练任务中。
【免费下载链接】cann-recipes-train本项目针对LLM与多模态模型训练业务中的典型模型、加速算法,提供基于CANN平台的优化样例项目地址: https://gitcode.com/cann/cann-recipes-train
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考