昇腾 910B 微调 Qwen3.8-27B:从 OOM 到 4096 上下文 LoRA
cutoff_len=16时启动、加载、分布式初始化全绿;一改成 4096,MLP 前向直接 OOM。当时脑子里第一反应是:LoRA 可训练参数才几百万,8 卡 FSDP2 怎么还会爆?
目标其实很朴素:在 8 张昇腾 910B(约 64 GB HBM/卡)上,给 Qwen3.8-27B 做 LoRA,学 SoulChat-R1 那套心理咨询对话。约束也清楚——尽量只动 Docker,不升级共享机宿主驱动;正式训练上下文别低于 4096;只存 adapter,不合并整模。
最后跑通的组合是:8 卡 FSDP2 + BF16 LoRA + NPU 上的 FLA DeltaNet + 全层非重入梯度检查点。rank 8,打在全注意力的 Q/K/V/O 上,cutoff_len=4096,3 个 epoch、2475 步。训练跑完了,但验证 loss 最低点在前面,后面主要是在继续拟合训练集。
下面的数字和修复都绑在这套具体版本上,别当成任意昇腾环境的通用菜谱。心理健康相关数据和模型只适合授权研究开发,没经过专业评估就别往临床场景推。
实验边界先钉死
基座目录叫 Qwen3.8-27B,来源记的是 Qwen/Qwen3.8-27B(国内镜像)。训练前核对过权重索引里的 18 个 safetensors 分片都在、且非空——能抓缺片,代替不了哈希校验;真要复现还得锁 revision。
容易绕的一点:这份权重走 Transformers 的 qwen3_5 实现,加载类是 Qwen3_5ForConditionalGeneration。别光看类名就说目录标错了,来源、配置、实际结构一起对。
语言部分 64 层:48 层 Gated DeltaNet,16 层全注意力(第 3、7、11……63 层,零起算)。只修普通 Attention,盖不住整棵计算图。
| 组件 | 本次实测 |
|---|---|
| 硬件 | 8 × Ascend 910B,约 64 GB HBM/卡 |
| 宿主驱动 | 25.3.rc1(没升级) |
| CANN | 9.0.0 |
| PyTorch / torch-npu | 2.7.1+cpu / 2.7.1.post4 |
| Transformers | 5.6.0 |
| PEFT / Accelerate | 0.18.1 / 1.11.0 |
| fla-core / flash-linear-attention | 0.5.2 / 0.5.2 |
| triton-ascend | 3.2.1.dev20260530 |
| 训练框架 | LLaMA-Factory(容器内版本)+ 进程内适配入口 |
torch 带 +cpu 不代表没用 NPU——后端是 torch-npu,看设备和日志。Triton 包元数据和 triton.__version__(报 3.2.0)也对不太齐,只能当「这次测过的组合」,别包装成官方兼容矩阵。
16 token 短测只能证明「能启动」
早年为了验证链路,把序列缩到 16。它能查启动、加载、分布式初始化,验不了心理咨询长对话。
cutoff_len=16 是处理后样本长度上限约 16 token,不是 16K,也不是 batch size。带 system、history、问答的样本,这点长度连有效监督内容都未必留得住。
后来验收底线钉在 4096,并分清三件事:
- 配置上限 4096 只管张量长度;超了照样截断,单条对话不一定完整。
- 动态 batch 里出现 1552、2280,只说明样本本身更短,不是被偷偷下调了。
- 短样本 pad 到 4096,测的是这个长度下的算力/显存压力;有效监督 token 有多少是另一本账。
数据在远端现成:训练 6598、验证 300。转换时保留 system、history,模板用兼容的 qwen3_vl_nothink,确认无思考版输出里没有 <think>。正式复现最好再统计 template 后的长度分位数和超长比例——光听「长上下文数据集」不够。
LoRA 省的是可训练状态,不是激活
LoRA 砍的是可训练参数及其梯度、优化器状态,基座前向还在。峰值内存大致是:
基座权重及分片临时展开
+ LoRA 参数 / 梯度 / 优化器状态
+ 反向要留的激活
+ 算子临时区与通信缓冲
长序列上,后两项经常才是主角。权重冻住了,模块仍可能处在梯度回传路径上——后面的 LoRA 要更新,梯度往往得穿过中间冻结层。
所以这几句听起来很顺、其实站不住:
- 「可训练参数才几百万,总显存一定小。」
- 「FSDP2 切到 8 卡,就不会 OOM。」
- 「换成 QLoRA,长序列内存就全解决。」
FSDP2 主要分参数等状态,不会自动把每层序列激活均摊到 8 张卡。QLoRA 能压量化基座的存放,还得看 NPU 上量化算子和训练是否真能用;它也解不了「激活保存策略错了」。这次最后上的是 BF16 LoRA,不是 QLoRA,也没做真正的序列并行。
DeltaNet:fast path 警告别一口咬死
DeltaNet 会报 fast path 不可用。直觉容易理解成:缺 causal-conv1d,整套高效路径全废。
翻过当前版实现后发现:因果卷积、DeltaNet 核心、归一化可以分开选。我们最后留下的是:
| 模块 | 实际路径 |
|---|---|
| 因果深度卷积 | Transformers 原生卷积 + SiLU + 切片 |
| Gated RMSNorm | 原生实现 |
| DeltaNet 分块计算 | NPU 上的 FLA chunk_gated_delta_rule |
另一个坑在检测逻辑:Transformers 对 FLA 的可用性判断带 CUDA 条件。容器里明明装了能在 NPU 上跑的 FLA,也可能不会自动选上。
做法是在隔离训练进程里显式挂上验证过的函数,不动宿主驱动:
from fla.ops.gated_delta_rule import chunk_gated_delta_rule
from transformers.models.qwen3_5 import modeling_qwen3_5 as modeling
# 避开 CUDA 导向的融合实现
modeling.FusedRMSNormGated = None
modeling.causal_conv1d_fn = None
modeling.causal_conv1d_update = None
def npu_delta(*args, **kwargs):
if args[0].device.type != "npu":
raise RuntimeError("Expected NPU tensors")
return chunk_gated_delta_rule(*args, **kwargs)
modeling.chunk_gated_delta_rule = npu_delta
真实入口还会按 LOCAL_RANK 设设备,并打首次调用的设备和 shape。所有 rank 都出现 FLA_NPU_ACTUAL_CALL,比一句 warning 靠谱。训练里的 chunk 路径验过了,推理适配还得单独做。
梯度检查点:冻层也得算进去
接上 FLA 并没立刻通。第一轮还是在 MLP 前向 OOM。
翻 LLaMA-Factory 自定义 checkpoint 逻辑:会按「模块有没有可训练参数」决定要不要检查点。LoRA 只打在少数注意力投影上时,大量 DeltaNet / MLP 是冻的,但梯度路径还在;这些层一跳过,激活照样堆。
第二轮覆盖这个行为还是挂。回溯发现反向重计算仍在用 reentrant checkpoint,MLP 反向重算时 OOM。只在初始化改一次函数不够——Trainer 后面还可能改回去。
最终就两刀:
- 对
GradientCheckpointingLayer一律开检查点,包括冻结但传激活梯度的层。 - 在第一次真正的
Trainer.training_step里设use_reentrant=False,防止被框架盖掉。
from functools import partial
from torch.utils.checkpoint import checkpoint
from transformers.modeling_layers import GradientCheckpointingLayer
# 放在首个 training_step 里执行一次,别只在建模前改
for module in model.modules():
if isinstance(module, GradientCheckpointingLayer):
module.gradient_checkpointing = True
module._gradient_checkpointing_func = partial(
checkpoint, use_reentrant=False
)
这是绑版本的内部属性修补,升级 Transformers / LLaMA-Factory 后要重查实现并冒烟。运行时每个 rank 看到 91 个 checkpoint 层、92 个 FSDP 模块——别和「91 个语言层」混;语言层还是 64。
非重入在所有环境是不是都更省,我不下结论。这次日志和对照实验只说明:全层覆盖 + 运行时强制非重入,让前面挂掉的配置过了关。
验收别停在 import 成功
我把验收拆成几层:算子能不能跑、数值是否离谱、整模能不能训、adapter 有没有真更新。
DeltaNet 前向与梯度
同组 BF16 舍入输入,CPU FP32 参考 vs NPU BF16,输出和五组输入梯度的相对 L2:
| 比较对象 | 相对 L2 |
|---|---|
| output | 0.003616 |
| dq | 0.003687 |
| dk | 0.004031 |
| dv | 0.003795 |
| dg | 0.003824 |
| dbeta | 0.004132 |
大约 0.36%–0.41%,过了这次的 5% 门槛。有限规模数值测替代不了数学证明。测试自身也踩过坑:FLA 布局是 [B, T, H, D],时间维和 head 维搞反,算完也不等于验到了想要的序列长度。
4096 长度
单算子用 [1, 4096, 48, 128],输出/梯度有限,峰值 allocated 约 1.39 GiB——这是算子测,不是整模。
整模先跑 rank8、只打 Q 的 10 步冒烟,再跑 pad 到 4096 的压力测。pad 区 attention_mask=0、labels=-100,不进监督损失。
扩到 Q/K/V/O,并核对 adapter
只打 Q 是排查时的窄范围,不是 LoRA「默认正确解」。扩到 Q/K/V/O 仍 rank8,避免同时改覆盖面和 rank,说不清开销从哪来。
可训练参数 5,242,880(约 0.0192%)。Q/K/V/O 各命中 16 层,共 64 个投影,存出 128 个 A/B 张量。名字只打到全注意力层,另外 48 层 DeltaNet 投影模块名不同,写 q_proj,k_proj,v_proj,o_proj 不等于「全语言层都挂了 LoRA」。
Q/K/V/O 的 10 步测跑完,训练段约 356 秒;固定 4096 压力下峰值 allocated 约 32.3 GiB/卡。保存后逐项查有限值,并确认四类投影每层的 B 权重非零——真有更新,不是一份初始化壳。300 条验证跑完,验证 loss 4.8649,退出码 0;这是流程检查,效果另说。
正式训练配置
同一套适配入口,去掉冒烟步数上限和强制 padding,按真实长度动态补齐。主配置摘录:
stage: sft
finetuning_type: lora
lora_rank: 8
lora_alpha: 16
lora_dropout: 0.05
lora_target: q_proj,k_proj,v_proj,o_proj
template: qwen3_vl_nothink
cutoff_len: 4096
bf16: true
gradient_checkpointing: true
per_device_train_batch_size: 1
gradient_accumulation_steps: 1
learning_rate: 5.0e-5
lr_scheduler_type: cosine
warmup_ratio: 0.03
num_train_epochs: 3.0
max_grad_norm: 1.0
save_only_model: true
save_strategy: steps
save_steps: 100
eval_strategy: steps
eval_steps: 100
每卡 batch 1、累积 1,有效全局 batch 8。6598 条数据,每 epoch 825 步,3 epoch 共 2475 步。
FSDP2 按 Qwen3_5DecoderLayer / Qwen3_5VisionBlock 包,开 reshard_after_forward 和 CPU RAM 高效加载,没开参数 CPU offload。视觉塔和多模态 projector 冻着——这次是文本对话适配。
检查点是反向重算换内存,没有用 detach 砍跨层梯度,也不是截断 BP。输入仍封顶 4096;「超长截断」和「截断反向」别混。
Docker 后台跑,日志写挂载目录,launcher 加锁防重复启动,结束后记退出码。SSH 断了不一定杀任务;主机重启、容器挂、进程崩照样会停。
save_only_model: true 只存 adapter,优化器/调度器状态没了——能接着训,严格无损续训谈不上。
训练跑完了,验证集却先降后升
正式任务北京时间 9 月 6 日 07:16 起,7 日 09:05 左右 2475 步跑完。Trainer 累计约 25 小时 48 分,含周期性评估,不是纯算时间。
最大 allocated 约 34,958,087,680 字节(≈32.6 GiB/卡)。别和 npu-smi 直接对表——后者还有缓存池、运行时等。看危险程度要综合 allocated、reserved、申请失败和其他进程。
验证曲线比「最后训练 loss 很低」更有用:
| step | 验证 loss |
|---|---|
| 100 | 2.2275 |
| 500 | 2.1512 |
| 800 | 2.2403 |
| 1100 | 2.3991 |
| 1600 | 2.6346 |
| 2000 | 2.7945 |
| 2400 | 2.8168 |
| 2475 | 约 2.819 |
末段训练日志 loss 约 1.362,只是最近窗口,不是全训练集均值。验证 loss 从第 500 步后总体往上走——继续拟合训练集,验证侧没跟着受益,过拟合味道很重。
这次没配早停、也没自动加载最佳 checkpoint,就是硬跑满 3 epoch。更合理的下一步是先比 checkpoint-500 和后面几份,别直接拿最后一份当宝。最低验证 loss 也不等于咨询能力最好,还得留出对话看连贯性、追问、套话和危机回应,必要时找人评;本文不下临床结论。
进度条 100% 不等于任务结束
2475/2475 只说明优化步走完。后面还有完整验证、分片汇总、存 adapter、写指标、清进程。
末步后的评估 300 条、38 个分布式 batch,实测约 514 秒。配置同时开了训练中评估和 do_eval,日志确认:存模后又跑了一遍最终评估。所以收尾预留十几到二十多分钟,比看到 100% 就杀容器稳。想省重复评估可以以后改工作流,别在收尾中途动刀。
9 月 7 日 09:18 查日志:最终 adapter 09:15 已写入顶层输出目录(20,990,496 字节),最终评估还在跑,退出码还没定——还不算完整交付。
我的完成清单:步数到位、进程退出码 0、最终 adapter 可读、关键张量有限、指标文件齐。进度条 100% 或目录里有旧 checkpoint,都不能单独过关。
若重来一次
先统计 template 后长度和有效监督 token,再定上下文和验收用例。16-token 只配当启动检查。
兼容性分层验:能 import、算子能跑、数值梯度像样、整模前后向通、保存评估完成——别从其中一环直接跳到「全通」。
OOM 先定位阶段:前向、反向重算、参数展开、验证,内存来源可能完全不同;别只会把 rank、batch、cutoff 往小拧。
共享机升级尽量困在容器里,但宿主驱动仍是共同依赖。这次跑通,跟新 CANN 能不能配旧驱动,得另测。
验收目标也该从「训完」挪到「选出值得用的那份」:提前定早停、缩短周期或调学习率,再用留出集比;别在同一验证集上反复调参,还当它是独立测试。
这回真正记牢的几条:LoRA 消不掉激活内存;FSDP2 不会自动给你序列并行;CUDA 导向的检测不等于 NPU 算子不可用;冻层也可能必须 checkpoint。每层都用日志和数值钉死,才知道卡在哪,也知道结论能说到哪一步。
延伸阅读与复现依据
- PyTorch activation checkpoint:重入 / 非重入差别。
- PyTorch FSDP2:参数分片与 forward 后 reshard。
- PEFT LoRA:目标模块、rank、量化相关配置。
- Flash Linear Attention 与 Triton Ascend:本次 DeltaNet NPU 路径。
具体数据来自本次的 FLA_NPU_VALIDATION.md、qwen38_fla_entry.py、qwen38_qkvo_formal_20260906.yaml、数值校验日志和正式训练 trainer_log.jsonl。外链解释概念,不替代本机源码和实测。为少暴露共享环境,服务器地址、用户目录和原始咨询样本都略了。
版权声明: 本文首发于 指尖魔法屋-昇腾 910B 微调 Qwen3.8-27B:从 OOM 到 4096 上下文 LoRA(https://blog.thinkmoon.cn/post/1038-ascend-910b-qwen38-27b-lora/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。