跳转至

一面

项目提问

你整套后训练训练架构是怎么搭建的?分布式训练遇到过什么问题?

后训练架构整体设计:

我们的后训练架构分为三个核心阶段——SFT(监督微调)、Reward Modeling(可选)、以及偏好对齐(DPO/RLHF),并辅以数据引擎、评估体系和分布式训练基础设施。

  • 数据引擎层:负责从线上日志、人工标注、合成数据等多渠道收集原始数据,经过清洗、去重、脱敏、格式标准化后,按任务类型(对话、知识问答、写作、代码等)分类存储。使用动态采样策略,保证各能力域的均衡和长尾覆盖。

  • 训练调度层:基于Ray和Kubernetes构建弹性训练集群,支持多节点多卡分布式训练。使用DeepSpeed ZeRO-3和FSDP作为分布式策略,能够支持70B以上模型的微调。训练任务通过CI/CD流水线触发,自动拉取Docker镜像、挂载分布式存储、启动训练作业。

  • SFT阶段:基座模型选型(如LLaMA-3、Qwen-2)后,使用高质量指令数据进行监督微调。训练格式统一为ChatML或ShareGPT,支持system prompt、多轮对话。利用FlashAttention-2和序列打包(packing)提高GPU利用率。训练时监控loss、perplexity、token accuracy以及下游评测集上的指标。

  • 偏好对齐阶段:采用DPO(Direct Preference Optimization)作为主要对齐方法。构建偏好对数据时,使用当前SFT模型对同一prompt采样多个回复,由GPT-4或人工根据帮助性、安全性、真实性等多维度打分,形成chosen和rejected对。DPO训练中引入reference model的KL散度约束,防止模型偏离太远。

  • 评估与反馈层:离线评测覆盖MMLU、GSM8K、HumanEval、AlpacaEval、自建评测集等。在线评测通过A/B测试对比核心业务指标(满意度、完成率、留存)。评测结果反馈至数据引擎,形成闭环。

分布式训练遇到过的问题及解决:

  • 通信瓶颈:在使用ZeRO-3时,参数的all-gather通信成为瓶颈,尤其在跨机多卡场景。通过增加InfiniBand带宽、使用梯度累积、调整ZeRO的stage(ZeRO-2在部分层上更优)来缓解。

  • 显存碎片与OOM:变长输入序列打包后,最大序列长度设置不当容易导致OOM。我们根据数据长度分布动态调整packing策略,设置合理的max_seq_length和micro batch size,并开启activation checkpointing。

  • loss spike与发散:在DPO训练中,偶发loss尖峰,原因多为偏好对中正负样本差异过大或KL惩罚系数过小。通过clip gradient、调大KL系数、增加warmup step予以解决。

  • 多节点训练不稳定:偶发NCCL超时或节点掉线。通过启用NCCL重试、健康检查自动剔除故障节点、设置更长超时时间,使训练能自动恢复。


LoRA训练时超参怎么选择,学习率、批次大小依据是什么?

LoRA超参数选择:

  • 秩 r:通常选择 8、16 或 32。r 越大,可训练的参数量越多,拟合能力越强,但过拟合风险也增大。经验上,简单任务 r=4~8 足够,复杂多任务可到 16~32。我们根据参数量预算和下游数据量,通过小规模网格搜索确定。

  • alpha (缩放系数):LoRA 的实际更新幅度为 ΔW = (alpha / r) * B A。alpha 通常设置为 r 的两倍(如 r=8, alpha=16),这样缩放系数为2。增大 alpha 等效于增大学习率,需要配合调整。保持 alpha/r 比值不变时,改变 r 对初始更新幅度影响不大。

  • 学习率 (lr):LoRA 的学习率通常比全参微调高 1-2 个数量级,因为只训练少量参数,需要更大步长才能有效更新。典型值在 1e-4 到 5e-4 之间。最终通过逐步尝试(如 1e-4, 2e-4, 5e-4)并观察训练 loss 和验证集指标确定。

  • 批次大小 (batch size):受限于显存,通常 micro batch size 设为 1~4,配合梯度累积达到等效较大的全局 batch size(如 32~128)。较小的 batch 可能引入更多噪声,有助于泛化;但过小会导致训练不稳定。我们根据数据量和模型规模,选择能使训练稳定的最小全局 batch size,再逐步增大观察性能。

  • Target Modules:通常选择所有 attention 层的 q_proj, v_proj(有时也包含 k_proj, o_proj, 甚至 FFN 的 gate/up/down)。增加模块能提升容量但也增加参数。根据下游任务复杂度选择,我们会在必要时对 FFN 也加 LoRA。

  • Dropout:一般不加或设置极小值(如 0.05)。LoRA 参数少,dropout 可能阻碍有效学习。

选择依据:最终超参由小规模实验确定。固定基座模型和数据,在几个关键评测指标上(如微调任务准确率、通用能力保持度)进行网格搜索,选择帕累托最优的配置。


多轮对话任务,如何稳定保持模型的输出格式一致性?

保持多轮对话中输出格式一致性,需从数据构造、训练策略和解码约束三个层面入手。

数据构造层面:

  • 所有训练数据严格遵守统一的对话模板,例如:
<|im_start|>system
You are a helpful assistant.<|im_end|>
<|im_start|>user
...<|im_end|>
<|im_start|>assistant
...<|im_end|>
  • 每轮结束均有明确的终止符。训练时 loss 只计算 assistant 部分,且强制模型学会在恰当位置生成终止符。

  • 对于结构化输出(如 JSON、Markdown表格),在训练数据中大量增加此类样本,并使格式严格一致。通过数据增强,对同一内容生成不同表述但格式相同的数据,提升模型对格式的鲁棒性。

训练策略层面:

  • 在 SFT 阶段加入“格式指令遵循”的专项数据,如“请用 JSON 格式回答”等指令,并检查模型输出是否能被解析。在 DPO 中,若格式错误直接视为 rejected。

  • 使用 token-level 的 loss mask,对格式控制 token(如 {, }, \n, <tag> 等)施加更高的 loss 权重,迫使模型精确生成这些 token。

  • 引入语法约束的监督信号,在训练时若模型输出的 JSON 无法解析,给出惩罚。

解码层面:

  • 在推理时,使用 guided generation / constrained decoding 技术(如 outlines、guidance 库或 vLLM 的 logits processor)。对于要求 JSON 输出的场景,实时构建有限状态机,强制模型只能生成符合 JSON 语法的 token,从根本上杜绝格式错误。

  • 将格式要求明确写入 system prompt,并开启 json_mode(如 OpenAI API 的 response_format),让模型和推理引擎协同保证格式。

后处理兜底:

  • 对输出进行正则匹配和解析,若失败则自动重试,并增加更明确的格式纠正指令。多次重试仍失败则降级为纯文本回复并告知用户。

多次迭代之后效果分化,你用了哪些消融实验定位问题?

效果分化是指在不同能力维度上,某些指标显著提升而另一些下降,或不同评测集表现不一致。我们采用分层消融来定位。

(1)数据层面消融:

  • 数据源隔离:将训练数据按来源或能力域切分(如百科、代码、对话、安全),逐一移除某一类数据重新训练,观察各指标变化,识别是哪类数据导致了能力退化或提升。

  • 数据量梯度:对关键数据子集,按比例(如 10%、30%、50%、100%)加入训练,绘制能力曲线,找到饱和点或毒化阈值。

  • 数据质量分析:使用奖励模型或 GPT-4 对训练数据评分,剔除低分样本,对比训练效果。若某项能力突然下降,检查是否因为引入了低质量甚至错误标注的数据(如幻觉样本)。

(2)训练策略消融:

  • LoRA 模块消融:对比仅训练 Q/V 投影 vs 训练所有 attention 参数 vs 加入 FFN。若通用能力下降,可能因 LoRA 扩展至 FFN 导致过拟合。

  • 损失权重消融:检查 SFT loss 和 DPO loss 的配比,以及 DPO 中 KL 惩罚强度。调整 KL 系数,观察对齐税的变化。

  • 混合精度与随机种子:固定其他条件,改变随机种子和混合精度设置,排除偶然因素。

(3)模型内部表征分析:

  • 使用探针(probing)检查模型隐藏层在旧任务和新任务上的表征差异。如果表征发生明显偏移,可能是灾难性遗忘的信号,需要通过回放(replay)或 EWC 类方法约束。

  • 分析 logits 分布,看模型是否变得极端(过自信)或过于平坦(能力衰退)。

(4)迭代级联消融:

  • 构建“版本链”:模型 v1 → v2 → v3,分别进行评测。若 v2 到 v3 分化,回退到 v2 并只添加单项改动,确认是哪个改动引入的问题。这样可快速定位罪魁祸首的更新。

通过上述消融,我们发现过某次迭代中安全能力下降是因为过度裁剪了安全对齐数据,而加入了大量未经过滤的网络数据;通过增加安全数据回放得以修复。


依旧八股:

对比普通多头注意力,GQA是如何兼顾速度与效果的

多头注意力 每个头都有独立的 Q、K、V 投影和各自的 KV 缓存,推理时 KV 缓存占用 = 2 × 层数 × 头数 × 序列长度 × 头维度。

多查询注意力 所有头共享同一组 K、V 投影(因此共享 KV 缓存),只有 Q 保持多头。缓存大小降为原来的 1/头数,极大降低显存,但对注意力质量有一定损失。

分组查询注意力 是两者的折中:将 h 个头分为 g 个组(g < h),组内共享 K、V 投影。每组内的所有头使用同一组 K、V,但每个头有自己独立的 Q。KV 缓存大小降为标准的 g/h。

兼顾速度与效果的原理:

  • 速度:减少 KV 缓存量和每次推理时 K、V 的计算量。GQA 的 KV 缓存仅需存储 g 份,是 MHA 的 g/h。在长序列和大batch推理时显存带宽瓶颈大大缓解,推理吞吐成倍提升。

  • 效果:共享 KV 会使模型损失一部分表达多样性。但实验表明,当 g 选择合适(如 8 头中分 2 组),性能下降极小。因为不同头的注意力模式常存在冗余,共享 K、V 不仅没有显著伤害效果,反而起到一定正则化作用。同时仍保留了多查询的多样性(通过多 Q 实现)。

实际应用:LLaMA 2 70B 使用 h=64, g=8,即每 8 个头共享 KV。它在速度上接近多查询,在质量上接近标准多头,实现了平衡。


RMSNorm相比于LayerNorm优化点在哪里?

LayerNorm 计算公式:

image.png

通常不再使用 β 偏移。

优化点:

  • 计算效率更高:省去了均值计算和减均值操作,减少了约 30% 的计算量。在大模型训练和推理中,每层的微小节省累积后影响显著。

  • 数值稳定性相似甚至更好:RMSNorm 实验证明其训练稳定性和收敛速度与 LayerNorm 相当,有时更优。去掉均值中心化不会损害训练,因为后续的线性层可以补偿平移。

  • 更少的参数和内存:通常没有 β 参数,参数量减半,内存访问更少。

  • 与Pre-Norm结合:主流 LLM 采用 Pre-Norm 结构,RMSNorm 作为其中的归一化层,使整体架构更轻量。

总结:RMSNorm 是 LayerNorm 的精简高效版,在保持归一化效果的同时显著降低计算和内存开销,成为现代大模型(LLaMA、Qwen等)的标准选择。


DPO为什么不需要单独训练奖励模型?

DPO(Direct Preference Optimization)将偏好对齐问题转化为直接对策略(语言模型)的优化,无需显式训练奖励模型。

核心推导:

在 RLHF 的 PPO 阶段,优化目标是最大化奖励同时约束策略与参考策略的 KL 散度。其最优策略满足:

image.png

优势:避免了奖励模型的训练开销和可能出现的奖励hacking问题,训练更稳定,计算成本更低。


长文本场景下RoPE会出现衰减,有哪些改进手段?

RoPE(Rotary Position Embedding)通过旋转变换编码相对位置,但在超出预训练长度时,高频分量(远距离)的点积会衰减,导致模型无法有效利用长距离信息。

改进手段:

  • 线性位置插值(PI):将超出训练长度的位置索引按比例压缩到训练窗口内。简单有效,但会降低对局部依赖的分辨率。

  • NTK-Aware 插值:将高频分量少压缩、低频分量多压缩,利用神经正切核(NTK)理论对RoPE的基频进行缩放。相比线性插值能更好地保持短距离性能。

  • YaRN:结合NTK插值和温度缩放,对注意力logits进行动态调整,进一步改善外推效果,是目前效果最稳健的方法之一。

  • ReRoPE / Self-Extend:在推理时,对于长输入,将序列分块,块内使用原始精确位置编码,块间使用插值或零向量,实现近似无限长度扩展。

  • 动态缩放:训练时使用随机长度或逐步增大长度,并在不同长度上使用对应的插值因子,使模型适应多尺度位置编码。

  • 改用其他位置编码:如ALiBi、xPos等,它们天然具有更好的外推能力。

  • 训练时混合长文本:在SFT和DPO阶段混入长文本数据,使模型在原生窗口内充分学习长距离依赖,减少对外推的依赖。

实践中,我们常采用YaRN或自扩展方法,以最小性能损失将上下文窗口扩展到目标长度。

二面

微调后模型旧能力下降,除了继续用LoRA,还有哪些方案?

灾难性遗忘的缓解方案:

  • 数据回放:在微调数据中混合一定比例的通用指令数据或旧任务的代表性样本,让模型在学新知识时不断回顾旧知识。

  • 弹性权重巩固:计算旧任务上各参数的重要性矩阵,在微调时对重要参数的更新施加二次惩罚,限制其偏离旧解。

  • 渐进式网络:为新任务添加小型的Adapter或Prompt模块,冻结原网络主体,实现知识与能力的完全隔离。

  • 多任务学习:将新旧任务数据混合,通过联合训练和多任务损失平衡,使模型在多个能力轴上同步优化。

  • 知识蒸馏:用旧模型作为教师,微调模型作为学生,在旧任务数据上计算蒸馏损失,约束学生不要遗忘教师的关键输出分布。

  • 梯度手术:在反向传播后,将梯度投影到旧任务损失不增的方向,或选择性屏蔽对旧能力有害的梯度分量。

  • 模块冻结与稀疏微调:仅微调模型的一部分(如顶层、特定层、新添加的模块),冻结其他部分。如 LLaMA-Adapter 的做法。

实际中常用的是数据回放+LoRA组合,成本低且效果显著。


批量生成文本容易高度同质化,从模型解码层面如何优化多样性?

解码层面优化多样性:

  • 温度采样:提高温度参数(如 0.8~1.2),使 softmax 分布更平坦,增大低概率 token 被采样的机会,直接提升多样性。但过高会引入乱码。

  • Top-k / Top-p 采样:限制候选 token 集合。k 和 p 设置过小(如 k=10)会导致同质化;适当增大 k(如 50)或 p(如 0.95)能增加多样性。

  • 典型采样:选择与平均信息量接近的 token,抑制过于罕见或过于常见的 token,在流畅性和多样性间平衡。

  • 重复惩罚:对已生成的 token 降低其 logits,避免陷入循环。可用于减少模型反复输出相同句式。

  • 对比解码:同时运行一个较弱模型(如较小版本),放大主模型与弱模型输出概率的差异,强化那些主模型特有、弱模型预测不到的 token,提升内容的独特性和信息量。

  • 多次采样+重排序:生成 N 个候选回复,通过多样性指标(如 distinct-n)或语义聚类选择最具差异性的结果。

  • 引导性解码:在 prompt 中加入“请用与之前完全不同的角度/措辞”等指令,从源头影响生成。

工程实践中,我们常采用较高温度(0.9)+ Top-p(0.95)+ 重复惩罚(1.1) 的组合作为默认,再按需配合对比解码或多次采样。


上线后推理时延过高,分别从模型结构和解码策略给出优化方案。

模型结构优化:

  • 量化:将模型权重量化到 INT8 或 INT4,减少内存占用和计算量,利用硬件对整型运算的加速。结合 GPTQ、AWQ 等算法保持精度。

  • 蒸馏:用大模型教师蒸馏出更小的学生模型,在保持性能的同时大幅降低推理成本。

  • 剪枝:移除冗余的注意力头和 FFN 神经元,减少计算量和参数量。

  • 高效变体:如使用 GQA/MQA 替代多头注意力,减少 KV 缓存;采用 RMSNorm 等轻量组件。

  • FlashAttention / PagedAttention:利用 IO-aware 的注意力算法,极大加速长序列自注意力计算,并优化显存使用。

  • 投机采样:使用小模型快速生成候选 token,再由大模型验证,在无损精度下加速生成。

解码策略优化:

  • KV 缓存:启用并优化 KV 缓存管理,避免重复计算。

  • 批处理:将多个请求组成 batch,利用 GPU 并行性,提高吞吐,降低单请求平均时延。

  • 流式输出:让首个 token 尽快返回,提升感知速度。

  • 提前停止:根据生成长度分布,设置合理的 max_new_tokens;对简单问题早期停止。

  • 推测解码:已经提及的投机采样,本质也是解码优化。

  • 约束解码:在特定场景下(如 JSON 输出),有限状态机缩小候选 token 空间,加速生成。

我们上线时主要应用INT8量化+GQA+FlashAttention+vLLM的连续批处理,使得延迟降低70%以上,满足实时交互要求。


业务问题:用户输入语句混乱,如何借助大模型统一用户意图?

处理流程:

  • 预处理与清洗:首先去除输入中的无意义符号、乱码、重复内容,但保留关键的口语化表达和拼写错误,作为上下文。

  • 意图澄清提示词:设计专门的 system prompt,要求模型对混乱输入进行改写,抽取核心意图。例如: “你是一个意图理解助手。用户输入可能混乱,请完成:1) 纠正常见错别字和语序;2) 识别用户的核心意图和关键实体;3) 如意图不明确,输出需要澄清的具体问题。”

  • 少样本示例:在 prompt 中给出几个典型混乱输入与清晰意图输出的示例,帮助模型快速进入改写模式。

  • 多轮确认:若模型置信度低(输出包含“可能”、“不确定”等),自动触发追问,向用户确认意图。例如:“您是想查询订单,还是想了解退货流程?”

  • 槽位填充:定义业务所需的关键槽位(如产品名、日期、金额),让模型从改写后的语句中提取并返回结构化 JSON,供下游执行。

  • 反馈闭环:记录意图理解结果与用户后续行为的匹配度,将误判 case 作为训练数据持续优化意图模型。

这样通过“清洗-改写-抽取-确认”的流水线,即使面对混乱输入也能较准确地统一用户意图。


多模态图文信息出现矛盾,设计简单的信息取舍策略。

图文矛盾是常见挑战,需根据应用场景决定取信策略。

策略设计:

  • 权威优先级机制:定义来源的可靠度等级。例如,官方文档中的文字说明 > 图像中的文字 > 模型的常识推断。在信息冲突时,高优先级来源覆盖低优先级。

  • 图像为主,文本为辅(视觉任务):对于目标检测、OCR、场景理解等以视觉为核心的业务,图像是 ground truth。当文本描述与视觉内容矛盾时,以视觉检测结果为准,并生成提示“根据图像,发现与描述不一致:...”。

  • 文本为主,图像为辅(语义任务):对于意图理解、逻辑推理,用户的文字指令代表真实意图。即使图像中某个属性与文本不符,也应遵循文本指令,并说明图像可能不匹配。

  • 置信度加权融合:对图像和文本各自的提取信息赋予置信度分数(如OCR的置信度、文本描述的确定性),通过加权投票决定最终答案。若双方置信度都低,启动人工确认。

  • 透明化沟通:无论采用哪种取舍,都在最终输出中告知用户矛盾点及我们的依据,让用户知情并可介入调整。

示例:用户上传商品图并问“这个红色的包多少钱?”,但检测到包是黑色的。策略:以视觉为准,回复“图中该款包为黑色,未找到红色包。黑色款的价格是...”。


带因果掩码的缩放点积注意力前向传播代码。

以下是简洁的 PyTorch 实现:

import torch
import torch.nn.functional as F
import math

def scaled_dot_product_attention_with_causal_mask(Q, K, V):
    """
    Q, K, V: (batch, num_heads, seq_len, d_k)
    返回: (batch, num_heads, seq_len, d_k)
    """
    d_k = Q.size(-1)
    # 计算注意力分数
    scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)  # (B, H, L, L)

    # 构建因果掩码(下三角为0,上三角为 -inf)
    seq_len = scores.size(-1)
    causal_mask = torch.triu(torch.ones(seq_len, seq_len, device=scores.device), diagonal=1).bool()
    scores.masked_fill_(causal_mask, float('-inf'))

    # Softmax 并 dropout(可选)
    attn_weights = F.softmax(scores, dim=-1)
    # attn_weights = dropout(attn_weights)  # 可选

    # 加权求和
    output = torch.matmul(attn_weights, V)  # (B, H, L, d_k)
    return output

要点:

  • scores 计算后进行缩放。

  • causal_mask 使用 torch.triu 生成上三角为 True 的布尔掩码,将上三角位置填充为 -inf,使得 softmax 后这些位置的注意力权重为零,实现“当前位置只能注意自身及以前”的因果条件。

  • 支持 batch 和 multi-head,张量维度为 (B, H, L, d_k)

  • 在训练时可一次计算整条序列的注意力,推理时结合 KV 缓存逐 token 生成。