今年 1 月到 3 月,我做了一轮代码补全大模型的微调。这篇把过程和结论整理出来:准备什么、怎么做、效果如何,以及踩过的坑。
资源准备
模型准备
目标是选一个性能良好的预训练模型。可以先看 BigCode 的代码大模型排行榜,榜单把模型分成三类:外部标准评估模型、在预训练基础上做过指令微调的模型、纯预训练模型。
目前主流的代码大模型主要是 StarCoder 和 CodeLlama 及它们的变体。StarCoder 系列是开源社区 BigCode 发布的代码补全大模型;CodeLlama 基于 Llama 2,又做了大量代码训练。具体选多大的模型,要结合补全场景、GPU 种类和推理性能综合考虑。
数据准备
主流大模型基本已经把网上能拿到的代码都训练过了,所以用开源数据再做微调,效果通常不好。要解决特定下游任务的补全,就得用自己的代码来构造数据。构造时既要参考模型原始的数据结构,也要贴合下游的补全场景,两者都要兼容。
微调需要高质量的数据集,通常是由 input 和 label 组成的数据对。代码补全模型属于自回归的因果语言模型:每一个输出都基于前面所有的输出。
- 分类、判断类模型:input 和 label 不一致,label 需要额外标注。
- 自回归补全模型:
input_ids和labels通常是同一段序列。模型从[token1, token2]预测token3,从[token1, token2, token3]预测token4,依此类推。
回到代码补全,就是截取代码片段,加入合适的 special token。构造好的片段既是 input,也是 label。
理论准备
微调是拿一个预训练模型,至少训练其中一部分内部参数。 相比从头预训练,它只需要少得多的标注数据和 GPU 资源,就能在特定任务上取得更好的效果。
常见的三种方式:
- 全参更新:训练所有参数。最简单,但计算成本最高,还有灾难性遗忘的问题,模型会“忘记”预训练阶段学到的有用信息。
- 迁移学习:保留大部分参数,替换网络的“头部”。能降低计算成本,但不一定能解决灾难性遗忘。
- 参数高效微调(PEFT):只用少量可训练参数增强基础模型,以很小的计算和存储成本,达到和全参更新相当的效果。LoRA 是其中最流行的方法,此外还有 QLoRA、AdaLoRA、P-Tuning 等。
全参更新和 PEFT 所需的显存差异巨大,再加上灾难性遗忘,这次我选了 LoRA。
用到的框架:PyTorch、transformers、datasets、peft、accelerate,训练过程用 wandb 可视化,后期加了 DeepSpeed。
微调实践
步骤
- 明确目标
- 选择预训练模型
- 选择微调方式和框架
- 构建特定任务数据集
- 设置合适的超参
- 建立评价体系,评估微调后的模型
- 重复 4、5、6 步,直到得到符合预期的模型
- 部署上线
明确目标
我先分析了现有补全场景的数据,挑出表现最差的 26 种补全场景,作为这次针对性提升的对象。目标有两条:
- 泛化能力保持不变或提升;
- 这些特定场景的补全能力明显提升。
基础模型选 StarCoder-15B,微调方式是 LoRA。
构建数据集
处理流程:
- 收集代码样本:收集大量 Java 代码文件。
- 预处理:
- 过滤掉开源依赖的代码,只保留团队自己开发的文件;
- 挑选长度在 1500 到 2000 token 之间的文件;
- 用 tree-sitter 分析代码,按那 26 种语法节点类型切分;
- 给切出的代码段加上
<fim_prefix>、<fim_middle>、<fim_suffix>这些 special token; - 保留仓库、文件名、节点类型等额外信息,用于训练和评估。
- 构造输入输出对:从代码样本中挖掉一段作为待补全部分。
- 切分数据集:训练集用于更新权重,验证集在训练过程中用来纠偏,测试集在训练完成后评估。
- 格式化:转成 transformers 能直接读取的 JSON。
为什么推理单卡够,微调却要好几张卡
训练时显存要装下四样东西:
- 模型参数;
- 激活值,前向传播时每一层计算出的 tensor;
- 梯度,反向传播时计算出的梯度;
- 优化器状态,比如 Adam,微调时这一部分占大头。
推理只需要前向传播;微调还要反向传播,保存每层的激活值和梯度,内存需求因此大幅增加。
建立评价体系
评估分两部分:
- 泛化能力:用 MultiPL-E 跑 HumanEval 的多语言版本。
- 特定下游任务:用自建测试集评估,对测试集的补全结果打分。
自建评测是我自己写的脚本:按语法节点类型分别统计,补全结果和原文完全一致才记 1 分。这样能看清模型在哪一类场景最差,而不只是一个总分。
实验结论
几轮实验下来,得到这些结论:
- 试过不同的 batch size 和 learning rate 之后,batch size 为 64、learning rate 为 5e-5 时效果最好:HumanEval-Java 从 0.298 提升到 0.312,自建测试集也有明显提升。
- learning rate 对微调效果影响很大。 依次验证了 6e-5、5e-5、5e-4、1e-4、5e-3、1e-3,无论是 loss 曲线还是实际效果,都有显著变化。取到 1e-3、5e-3 时,模型基本崩溃。
- batch size 影响一般。32、64、128 三种效果差距不明显,64 略好。
- 只针对一种语言微调,也会提升其他语言的补全效果。
- 效果较好的那组,loss 曲线反而不收敛。 loss 的收敛程度和最终效果无法对应,这一点我目前还解释不了。
换成已经做过指令微调的 WizardCoder 再练,专项分数上涨,泛化分数有所下降。全参更新时,HumanEval 直接掉到接近 0。
后期用 DeepSpeed 加速训练,一轮从大约 50 小时缩短到 10 小时左右,效果基本不变,可以更快地验证想法。训练 4800 步和 700 步的效果也差不多,说明这批数据很快就学饱和了。
踩过的坑
很多次分数暴跌,最后查下来都是工程问题,而且都不报错:
- LoRA 合并选错 base 模型。 前两个版本合并时,误选了 StarCoder 的 base 去合并 WizardCoder 的 LoRA 权重,HumanEval 评分暴跌。第三版改对之后,分数才回到正常范围。
- EOT token 写错。 从别的脚本抄过来的结束 token 不对,生成结果末尾总多一段。
- peft 版本不一致。 训练和合并用了不同版本,合并直接失败。
- 不同模型的 special token id 不同。 SantaCoder 和 StarCoder 的特殊 token id 不一样,换模型时要重新核对。
- 评测脚本的 prompt 格式。 同一个模型,评测时加不加 FIM token,分数能差出将近一半。换用 StarCoder2 时,自建集几乎是 0 分,很可能也是格式没对上。
最后一条提醒我:评测流程本身的写法,就可能决定分数高低。 拿分数下结论之前,先确认评测没问题。
顺带纠正一个我之前理解错的概念:采样 20 次、只要有一次通过就算通过,算的是 pass@k,不是 pass@1。
难点和后续
难点有三个:
- 充足的 GPU 资源;
- 数据构造,高质量的数据对微调影响非常大;
- 超参的种类多且杂。
后续计划:
- 尝试更好的模型,比如 StarCoder2、CodeLlama;
- 尝试更多超参组合;
- 跨文件代码片段的微调;
- 推理优化,量化和蒸馏;
- 验证微调模型能否复现效果,再上线做 AB 测试,看它对采纳率的实际影响。