news 2026/9/18 19:30:05

FLUX GRPO 文生图强化学习训练 NPU 适配实战:基于 DanceGRPO 与 HPSv2 的 Atlas A3 训练指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
FLUX GRPO 文生图强化学习训练 NPU 适配实战:基于 DanceGRPO 与 HPSv2 的 Atlas A3 训练指南

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_devicetorch.npu.set_devicedist.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_ENABLECOMBINED_ENABLEPYTORCH_NPU_ALLOC_CONF等 NPU 训练环境变量。
  • 训练 batch 组织优化:新增rollout_batch_sizetrain_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 相关模块,增加NpuFusedRMSNormtorch_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_npumain()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 分支(FluxTransformer2DModelFluxTransformerBlock, 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_qnorm_knorm_added_qnorm_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.pttorch.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_DEVICES0..15(16 die)指定训练使用的 NPU 编号列表
TASK_QUEUE_ENABLE2开启任务队列(Task Queue)特性,提升算子下发效率
COMBINED_ENABLE1使能通信与计算融合/组合优化
CPU_AFFINITY_CONF2配置进程 CPU 亲和策略
HCCL_CONNECT_TIMEOUT1200HCCL 建链超时(秒),大集群下建议调大
NPU_ASD_ENABLE0关闭芯片自检(ASD)相关逻辑,避免干扰训练
ASCEND_LAUNCH_BLOCKING0关闭算子同步阻塞,保持异步下发
ACLNN_CACHE_LIMIT100000ACLNN 算子编译缓存上限
MULTI_STREAM_MEMORY_REUSE2多流内存复用策略
PYTORCH_NPU_ALLOC_CONFexpandable_segments:True使能 NPU 显存的可扩展段分配,缓解碎片化
HCCL_BUFFSIZE800HCCL 通信缓冲区大小(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_DEVICESsource .../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/dcminpu-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.1

5. 源码准备与补丁应用

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.1torch-npu==2.7.1diffusers==0.32.0transformers==4.46.1accelerate==1.9.0bitsandbytes==0.46.1peft==0.13.2liger_kernel==0.4.1
  • 通用工具packaging==25.0ninja==1.11.1.4ml-collections==1.1.0absl-py==2.3.1inflect==6.0.4huggingface_hub==0.34.0protobuf==3.20.0pybind11==3.0.4(为 aarch64 编译 decord 等服务)。

同时env_setup.sh不再安装 CUDA 版 torch 与 flash-attn,.gitignore追加datahps_ckptimageskernel_metawandb等目录,避免训练产物污染 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_NUMtorchrun --nproc_per_node=$NPU_NUM驱动fastvideo/data_preprocess/preprocess_flux_embedding.py,脚本内置了与训练脚本一致的环境变量与data/fluxdata/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 的基线配置:

参数说明
--seed42全局随机种子
--rollout_batch_size4每设备 rollout 采样 batch
--train_micro_batch_size1训练微批大小
--train_batch_size1每设备(训练)batch,与梯度累积配合
--gradient_accumulation_steps4梯度累积步数,与 micro batch 对齐
--max_train_steps300最大训练步数
--learning_rate/--weight_decay1e-5 / 0.0001优化器超参
--mixed_precisionbf16混合精度训练
--num_generations12每组 prompt 生成的图片数(组内相对奖励)
--h/--w/--t720 / 720 / 1生成图像分辨率与帧数
--sampling_steps/--eta16 / 0.3rollout 采样步数与随机性参数
--timestep_fraction0.6训练时采样的时间步比例
--cfg0.0关闭 classifier-free guidance
--clip_range/--adv_clip_max1e-4 / 5.0优势裁剪范围
--shift/--use_group/--ignore_last/--init_same_noise3 / 开 / 开 / 开噪声调度偏移、组奖励、GRPO 细节控制
--checkpointing_steps40每 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=1rollout_batch_size=4gradient_accumulation_steps=4,对应全局训练批次大小(train GBS)为1 × 4 × 16 = 64。该配置下的实测训练耗时如下:

训练迭代数训练耗时
200 iterations17.22 小时
300 iterations25.83 小时

由于硬件平台、软件栈及具体运行环境存在差异,以上数据用于展示本样例在 Atlas A3 上的实际训练性能,不作为严格的同条件性能对比。

9. 常见问题与排查

  1. git am失败:请确认上游源码版本与补丁生成版本匹配(DanceGRPO 需git checkout 15cc71d,diffusers 需git checkout v0.32.0);若上游代码已有较大变化,需要手动解决冲突后继续执行git am --continue

  2. 找不到 HPSv2 权重:请确认文件路径为hps_ckpt/HPS_v2.1_compressed.pt(训练代码中写死该相对路径),或在训练代码中修改cp路径。

  3. 更换 CANN 或 torch-npu 版本后出现编译缓存问题:建议清理如下目录后重新运行:

    rm -rf kernel_meta rm -rf .cache rm -rf /root/.cache
  4. --use_pickscore不生效:当前该参数暂未适配 NPU,训练脚本建议保持使用--use_hpsv2(开启 pickscore 只会打印 Warning,不会报错,但也不会产生 PickScore 奖励)。

  5. 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++"
  6. 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,然后运行程序即可。

  7. ImportError: .../libGLdispatch.so.0: cannot allocate memory in static TLS block:推荐预加载libGLdispatch。先定位系统libGLdispatch.so.0find /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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/18 19:27:04

27英寸显示器选购:4K与高刷取舍、Mini LED与OLED对比解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/18 19:26:33

VLM、VLA、VLN三者区别详解:从视觉理解到具身智能导航

最近后台和群里收到不少提问,都是围着三个缩写转:VLM、VLA、VLN。有人把VLA当成VLM的升级版,有人以为VLN是VLA的一个数据集,还有人干脆把三个词混着用。其实这三个概念在具身智能、多模态大模型和机器人导航领域各占一个位置&…

作者头像 李华
网站建设 2026/9/18 19:26:27

AI陪伴长期记忆架构:事实-模式-意图三层设计

1. 为什么“AI陪伴”必须解决长期记忆,而不是只靠上下文窗口?我第一次在真实产品中部署AI陪伴对话模块时,团队里所有人都觉得“用好大模型的上下文长度就够了”——毕竟主流模型现在都能塞进32K甚至128K token,聊个几十轮对话、记…

作者头像 李华
网站建设 2026/9/18 19:26:18

图像内容自适应滤波:原理、实现与参数调优指南

简介:PDF文档《一种基于内容的图像自适应滤波算法》是一篇面向图像处理与人工智能方向研究者的算法论文。该论文针对高斯白噪声和椒盐噪声干扰下的图像去噪问题,提出基于图像分块内容自适应调整滤波系数的思路,融合均值滤波与中值滤波优势&am…

作者头像 李华
网站建设 2026/9/18 19:22:55

Go 服务内存泄漏定位实战:pprof 与 inuse_space 分析

Go 服务内存泄漏定位实战:pprof 与 inuse_space 分析在很多人印象中,Go 拥有现代化的垃圾回收器(GC),基本不会发生内存泄漏。但在线上长期运行的高并发服务中,“Goroutine 泄漏”和“未释放的切片底层数组引…

作者头像 李华
网站建设 2026/9/18 19:20:53

自组织视角下智能制造系统演进与仿真实践

简介:一份围绕智能制造系统技术演进的学术文献,以自组织方法论为分析框架,面向智能制造研究者、产业规划人员以及系统开发从业者。内容从传统制造业痛点出发,剖析现有研究方法的局限,并基于系统开放性、非线性等特征&a…

作者头像 李华