昇腾 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(没升级)
CANN9.0.0
PyTorch / torch-npu2.7.1+cpu / 2.7.1.post4
Transformers5.6.0
PEFT / Accelerate0.18.1 / 1.11.0
fla-core / flash-linear-attention0.5.2 / 0.5.2
triton-ascend3.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 后面还可能改回去。

最终就两刀:

  1. 对 GradientCheckpointingLayer 一律开检查点,包括冻结但传激活梯度的层。
  2. 在第一次真正的 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
output0.003616
dq0.003687
dk0.004031
dv0.003795
dg0.003824
dbeta0.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
1002.2275
5002.1512
8002.2403
11002.3991
16002.6346
20002.7945
24002.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。每层都用日志和数值钉死,才知道卡在哪,也知道结论能说到哪一步。

延伸阅读与复现依据

具体数据来自本次的 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/) 转载或引用必须申明原指尖魔法屋来源及源地址!