大模型面经——超细节大模型训练与微调实操经验总结(上)
大模型的训练和微调,确实和以前做 NLP 微调的路数不太一样。光是那些新冒出来的“坑”,就够让人头疼一阵子。更不用说,做个简单的消融实验,试错成本也比原来高出一大截。很多搞算法的小伙伴,日常工作几乎都扑在处理数据上,真没太多余力去做探索性的实验。
所以,在动手实践之前,多看一些通用的实践经验,带着“先验知识”去摸索,就变得格外重要。这能帮你少踩很多无意义的坑。
接下来这个系列,会尽量把大模型训练和微调中的实操经验掰开揉碎了讲。作为开篇,这次主要聚焦在
训练数据预处理、模型结构、训练参数设置与错误处理
- 拿到一批新的业务对话数据要做 SFT,怎么优化这批数据?
- 模型训练时,历史对话长度是不是永远越长越好?一般设多少?
- 训练样本量一大,任务直接 OOM 了,怎么办?
- 微调大模型时,在模型结构方面有什么经验?
- 微调时训练配置一般怎么设?
- 微调过程中要是崩溃报错了,该怎么处理?
拿到业务产生的一批对话数据,需要进行 SFT,怎样对这批数据进行优化?
这是最实操的问题,可以从几点入手:
1. 上下文内容处理
得考虑具体模型能处理的历史对话长度。输入时,对历史对话数据进行左阶段,尽量保留最新的对话记录。历史的太远,模型可能就记不住了。
2. 语句顺滑处理
把一些口语化的语气词、语法错误顺一顺。比如那些频繁出现的“嗯嗯”、“呃”、“啊啊”,这些对模型理解语义没什么帮助,反而可能引入噪声。
3. 去掉一些敏感或不合适的内容
可以从整句和词两个层面来做。
-
整句层面
-
词层面
4. 扩充用户特征标签
如果业务允许,可以基于用户的年龄、性别、地域、人群等维度,给对话数据打上标签。这些标签对于后续分析、做其他实验,都是很有用的资产。
模型训练时,历史对话长度是不是设置得越长越好,一般设置多少?
这个问题可以做个简单的消融实验来验证。选同一个模型,分别用两种方案训练,变量就是 max_source_length 和 max_target_length。然后从 Loss、Bleu 指标、离线人工评估几个角度去对比。
直接说结论:基于现有显存条件,从人工评估少量样本和 Loss 下降趋势来看,历史对话长度设置得越长越好。
1024 长度明显优于 512 长度
建议直接扩大到 1024 长度
模型训练样本量规模增大,导致训练任务直接报 OOM 了,该怎么办?
遇到这种情况,常规手段是数据并行处理。核心思想是
让数据向量化的耗时,随着处理进程数的增加而线性下降
具体操作也很清晰:
1. 将完整数据集平均分到所有进程(也就是总的 GPU 卡数)上。
2. 每个 epoch 训练时,整体数据分片做一次 shuffle。在每个进程的同一时间,只加载单个分段大小的数据集。
3. 如果重新训练,可以直接加载之前向量化好的数据,省去重复处理的时间。
微调大模型的时候在模型结构方面有哪些经验?
这里有几条行业里已经验证过、可以放心用的经验:
- :目前主流都是用 Causal Decoder + LM 的架构,它的 zero-shot 和 few-shot 能力很突出,也有比较明显的涌现效应。
模型结构
- :推荐使用 Pre RMS Norm,实践证明它更稳定。
Layer normalization
- :GeGLU 或 SwiGLU 是当前的主流选择,效果比传统的 ReLU 好。
激活函数
- :Embedding 层之后不要加 Layer Normalization,否则会损害 LLM 的性能。
Embedding 层
- :ROPE 或 ALiBi 都可以,目前 ROPE 的应用范围更广一些。
位置编码
- :建议去掉 Dense 层和 Layer Norm 中的偏置项,这有助于提升训练的稳定性。
去除偏置项
微调大模型的时候在训练配置方面有哪些经验?
这些配置参数有比较通用的最佳实践:
- :在硬件显存允许的情况下,batch size 越大越好。甚至可以后期动态增加 batch size,GPT-3 当时的做法就是从 32K token 逐渐增加到 3.2M token。
Batch size
- :采用 warmup 再衰减的策略。先让学习率线性增长,到达预设最大值后,再通过余弦衰减降到最大值的 10%。这个最大值通常在 5e-5 到 1e-4 之间。
学习率设置
- :一个常用的安全策略,通常将梯度裁剪为 1.0,防止梯度爆炸。
梯度裁剪
- :采用 AdamW 优化器,权重衰减系数设置为 0.1。AdamW 相当于给 Adam 加了一个 L2 正则项。
权重衰减
- :推荐使用 bfloat16 而不是 float16 来训练,前者在数值范围上更友好,不容易出现溢出。
混合精度训练
微调大模型时出现错误崩溃该怎么办?
训练训得好好的,结果过某个 shard(数据分片)时突然崩了,大概率是数据出了问题。这时候的救急办法是:
选择一个好的断点,跳过导致崩溃的数据段,进行断点重训
怎么才算一个好的断点?有两个判断标准:
1. 损失标度(loss scale)大于 0。
2. 梯度的 L2 范数小于某个固定值,并且波动很小。如果这两个条件都满足,就可以放心地从这里接续训练。
-
- 关于宇宙的好的网名有哪些
- 角色扮演 | 1
- 网名