AI 实践 / NOTES

微调代码补全模型三个月

从数据构造、评测体系到超参实验,还有几个不报错、却让分数暴跌的坑。

今年 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 资源,就能在特定任务上取得更好的效果。

常见的三种方式:

  1. 全参更新:训练所有参数。最简单,但计算成本最高,还有灾难性遗忘的问题,模型会“忘记”预训练阶段学到的有用信息。
  2. 迁移学习:保留大部分参数,替换网络的“头部”。能降低计算成本,但不一定能解决灾难性遗忘。
  3. 参数高效微调(PEFT):只用少量可训练参数增强基础模型,以很小的计算和存储成本,达到和全参更新相当的效果。LoRA 是其中最流行的方法,此外还有 QLoRA、AdaLoRA、P-Tuning 等。

全参更新和 PEFT 所需的显存差异巨大,再加上灾难性遗忘,这次我选了 LoRA。

用到的框架:PyTorch、transformers、datasets、peft、accelerate,训练过程用 wandb 可视化,后期加了 DeepSpeed。

微调实践

步骤

  1. 明确目标
  2. 选择预训练模型
  3. 选择微调方式和框架
  4. 构建特定任务数据集
  5. 设置合适的超参
  6. 建立评价体系,评估微调后的模型
  7. 重复 4、5、6 步,直到得到符合预期的模型
  8. 部署上线

明确目标

我先分析了现有补全场景的数据,挑出表现最差的 26 种补全场景,作为这次针对性提升的对象。目标有两条:

  • 泛化能力保持不变或提升;
  • 这些特定场景的补全能力明显提升。

基础模型选 StarCoder-15B,微调方式是 LoRA。

构建数据集

处理流程:

  1. 收集代码样本:收集大量 Java 代码文件。
  2. 预处理:
    • 过滤掉开源依赖的代码,只保留团队自己开发的文件;
    • 挑选长度在 1500 到 2000 token 之间的文件;
    • 用 tree-sitter 分析代码,按那 26 种语法节点类型切分;
    • 给切出的代码段加上 <fim_prefix>、<fim_middle>、<fim_suffix> 这些 special token;
    • 保留仓库、文件名、节点类型等额外信息,用于训练和评估。
  3. 构造输入输出对:从代码样本中挖掉一段作为待补全部分。
  4. 切分数据集:训练集用于更新权重,验证集在训练过程中用来纠偏,测试集在训练完成后评估。
  5. 格式化:转成 transformers 能直接读取的 JSON。

为什么推理单卡够,微调却要好几张卡

训练时显存要装下四样东西:

  1. 模型参数;
  2. 激活值,前向传播时每一层计算出的 tensor;
  3. 梯度,反向传播时计算出的梯度;
  4. 优化器状态,比如 Adam,微调时这一部分占大头。

推理只需要前向传播;微调还要反向传播,保存每层的激活值和梯度,内存需求因此大幅增加。

建立评价体系

评估分两部分:

  1. 泛化能力:用 MultiPL-E 跑 HumanEval 的多语言版本。
  2. 特定下游任务:用自建测试集评估,对测试集的补全结果打分。

自建评测是我自己写的脚本:按语法节点类型分别统计,补全结果和原文完全一致才记 1 分。这样能看清模型在哪一类场景最差,而不只是一个总分。

实验结论

几轮实验下来,得到这些结论:

  1. 试过不同的 batch size 和 learning rate 之后,batch size 为 64、learning rate 为 5e-5 时效果最好:HumanEval-Java 从 0.298 提升到 0.312,自建测试集也有明显提升。
  2. learning rate 对微调效果影响很大。 依次验证了 6e-5、5e-5、5e-4、1e-4、5e-3、1e-3,无论是 loss 曲线还是实际效果,都有显著变化。取到 1e-3、5e-3 时,模型基本崩溃。
  3. batch size 影响一般。32、64、128 三种效果差距不明显,64 略好。
  4. 只针对一种语言微调,也会提升其他语言的补全效果。
  5. 效果较好的那组,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。

难点和后续

难点有三个:

  1. 充足的 GPU 资源;
  2. 数据构造,高质量的数据对微调影响非常大;
  3. 超参的种类多且杂。

后续计划:

  1. 尝试更好的模型,比如 StarCoder2、CodeLlama;
  2. 尝试更多超参组合;
  3. 跨文件代码片段的微调;
  4. 推理优化,量化和蒸馏;
  5. 验证微调模型能否复现效果,再上线做 AB 测试,看它对采纳率的实际影响。