返回

LLM后训练知识:如何判断 SFT 已到位

一、核心内容总结

1.1 核心问题(Motivation)

在大模型后训练(post-training)中,SFT(Supervised Fine-Tuning,监督微调)通常是模型对齐和能力注入的第一阶段。一个被普遍低估的关键问题是:

SFT 到底应该在什么时候停止?

如果只盯着 train loss、validation loss 或单一 benchmark 分数,很容易把”更好地拟合标注答案“误判为”模型能力持续提升“。

尤其在 reasoning(推理)、code(代码)、tool-use(工具调用)、agentic task(智能体任务)中,过度 SFT(over-SFT)会带来三个典型副作用:

  • 输出模板化:回答趋于固定套路,缺乏灵活性;
  • 采样空间收缩:生成分布变窄,多样性下降;
  • 后续 RL 收益下降:策略分布被过度压缩,RL 阶段缺乏探索空间。

1.2 判断 SFT 是否到位的关键指标(原文)

不能只看监督指标(loss)是否下降,应综合观察以下信号:

指标含义 / 信号
Hard/OOD 泛化继续训练是否还能提升 hard set、OOD set 和真实任务指标
pass@k对 math/code/tool-use/agentic 任务尤为重要
多样性(Diversity)采样多样性是否被压缩
协议稳定性格式、协议类能力是否稳定
RL probe gainRL 探针增益是否明确

关键判断逻辑:

  • 进入低收益区间的信号:train loss 继续下降,但 OOD/hard set 不涨,且输出开始模板化 → SFT 已到头。
  • pass@k 信号:如果 pass@1 不高但 pass@k 高 → 应优先考虑 RL 或 preference optimization,而不是继续堆 SFT。

1.3 其他要点(Takeaways)

  • 协议类能力最适合用 SFT 解决:chat template、JSON schema、工具调用格式、拒答格式、role consistency。
  • 过度 SFT 会压缩策略分布,削弱后续 RL 的探索空间和训练收益。
  • SFT 数据应按能力分桶构建与评估:通用指令、领域任务、工具调用、多轮对话、安全拒答、结构化输出、长上下文等。
  • 最终 checkpoint 选择:不应只选 SFT 分数最高的版本,而应选择——协议稳定、OOD 不退化、pass@k 保持较高、RL probe gain 明确的版本。

1.4 SFT 与 RL 的边界

SFT 解决”会不会按正确方式做”,RL 解决”多个可行行为中哪个更优”。

对后续要做 RL 的模型,SFT 的目标不是最大化监督拟合,而是构造一个稳定、可控、保留优化空间的 policy 初始化。


二、关键指标详解与实现

下面对每一个指标给出为什么重要 → 怎么度量 → 参考实现的完整链路。所有代码假设使用 HuggingFace transformers + vllm(用于高效采样)。

2.1 Hard / OOD 泛化(Generalization Gap)

为什么重要:train loss 下降只说明模型在”记忆”训练分布。真正衡量能力的是它在未见过的难样本(hard set)分布外样本(OOD set)上的表现。当二者停止提升甚至回退,而 train loss 仍在降,说明模型开始”过拟合标注风格”而非”学到能力”。

怎么度量

  • 维护三个独立评测集:in-domain val、hard set、ood set。
  • 每隔 N 个 step 记录三者的任务指标(accuracy / exact-match / reward)。
  • 观察 泛化 gap = train_metric - ood_metric 是否随训练持续拉大。

三个集合怎么选(关键,选错会导致误判):

核心原则——三者必须与训练集严格去重、互不重叠,并且分别代表”同分布 / 更难 / 分布外”三个梯度。

集合定义怎么构建常见坑
in-domain val(同分布验证集)与训练数据同来源、同分布,只是没参与训练从训练数据里按样本随机切出 1%~5%(切分要在去重之后做,避免近重复泄漏);保持任务类型、难度、领域比例与训练集一致与 train 存在近重复(n-gram/embedding 相似)→ 指标虚高;它涨只能说明”拟合得更好”,不能单独作为停止依据
hard set(难样本集)同领域、但难度显著更高的样本①按 base 模型/教师模型的低置信度或高 loss 挑选;②取真实业务中出错率高的样本;③人工筛选长链推理、多步依赖、易混淆的 case。规模几百~几千即可若”难”只是标注噪声大,会把噪声当难度;要保证答案本身是可验证/高质量的
ood set(分布外集)来源/分布与训练集不同,检验真实泛化①换不同数据源(不同标注团队、不同网站、不同风格);②换相邻但未训练的子任务/领域;③直接用公开 benchmark 当 OOD(如 chat 用 IFEval、code 用未训练的题库);④时间上用训练截止后新产生的数据与 train 领域”偷偷重叠”(例如同一数据源换了个名字)→ 失去 OOD 意义

判定要点:

  • 三者去重是前提,推荐用 n-gram 重叠 + embedding 余弦相似度 双重过滤,剔除与训练集近重复的样本。
  • 看趋势而非绝对值:in-domain 持续涨、但 hard/ood 走平甚至回退,才是”SFT 到头/开始过拟合”的可靠信号。
  • 三个集合的指标口径要一致(都用同一 accuracy / EM / reward 定义),否则趋势对比无意义。
  • 规模建议:in-domain val 500~2000 条即可稳定;hard/ood 各几百条起步,越贴近真实业务越好。

举例:Agent / 工具调用任务如何划分三个集合

假设训练目标是「让模型学会调用一套内部 API 完成用户任务」,训练数据主要是电商域(查订单、退款、改地址)的多轮工具调用轨迹。

集合具体怎么选(Agent 场景)
in-domain val从电商域轨迹里随机切出一批没参与训练的样本:同样的工具集、同样的任务类型(查订单/退款)、同样的对话风格。用于确认”同分布上格式与成功率是否正常”。
hard set同为电商域,但挑更难的轨迹:①需要 5 步以上工具链才能完成;②含歧义指令(”帮我处理一下上次那个问题”需先澄清);③需要错误恢复(前一步工具报错后要重试/换路径);④参数嵌套复杂的调用。
ood set换未训练过的域,但用同一套/相似工具协议:①航旅域(订机票、改签、退票,即 τ-bench 的 airline 子集);②全新工具(训练时没出现过的 API,考验对 schema 的泛化);③直接用公开基准 BFCL / τ-bench / WebArena 当 OOD。
判定示例:若继续训练后,电商域(in-domain)成功率还在涨,但航旅域(ood)成功率走平、且 5 步以上难轨迹(hard)不再提升,同时输出的工具调用越来越模板化 → 说明协议已学到位、决策能力靠 SFT 已到头,应转 RL。

举例:视频生成任务如何划分三个集合

假设训练目标是「文生视频(T2V)指令微调」,训练数据主要是日常真实场景(人物、街景、动物)的短视频 + caption。

集合具体怎么选(视频生成场景)
in-domain val从训练分布里留出一批未训练的 prompt:同样的题材(日常场景)、同样的时长/分辨率、相近的 caption 风格。用于看 VBench 各维度在同分布上是否正常。
hard set同为真实场景,但挑更难生成的 prompt:①大幅度/快速运动(奔跑、跳跃、镜头快速摇移,考验时序一致性);②多主体交互(两人握手、球类碰撞);③复杂物理(水流、火焰、烟雾);④长 prompt / 多属性绑定(”红衣女孩牵着黑狗在雪地里走”,考验属性不串色)。
ood set换未训练过的分布:①新风格/新题材(动漫、赛博朋克、水墨,而训练只用真实视频);②未见过的相机运动或构图;③直接用公开基准 VBench / EvalCrafter 的 prompt 套件当 OOD;④训练截止后新出现的热点题材 prompt。
判定示例:若继续训练后,日常场景(in-domain)的美学分、成像质量还在涨,但动漫/水墨等新风格(ood)不涨、大运动难样本(hard)时序一致性不升反降,且 dynamic degree(动态程度)下降、镜头趋同(多样性坍塌)→ 说明进入过拟合区间,应停止 SFT,转视频 DPO / reward 微调。

实现(评测循环 + 早停信号):

``` ``` import numpy as np class GeneralizationMonitor: """ 监控 in-domain / hard / ood 三个集合的指标变化, 当 ood/hard 连续 patience 次不再提升时,判定 SFT 进入低收益区间。 """ def __init__(self, patience=3, min_delta=0.005): self.patience = patience self.min_delta = min_delta self.history = {"indomain": [], "hard": [], "ood": []} self.best_ood = -np.inf self.stall = 0 def update(self, indomain_score, hard_score, ood_score): self.history["indomain"].append(indomain_score) self.history["hard"].append(hard_score) self.history["ood"].append(ood_score) # 只要 ood 有明显提升就重置计数 if ood_score > self.best_ood + self.min_delta: self.best_ood = ood_score self.stall = 0 else: self.stall += 1 gap = indomain_score - ood_score saturated = self.stall >= self.patience return { "gen_gap": gap, "ood_stall_steps": self.stall, "sft_saturated": saturated, # True 表示建议停止 SFT } # 使用示例 monitor = GeneralizationMonitor(patience=3) for step, (idm, hard, ood) in enumerate(eval_stream): sig = monitor.update(idm, hard, ood) if sig["sft_saturated"]: print(f"[step {step}] OOD 已停滞 {sig['ood_stall_steps']} 次,gap={sig['gen_gap']:.3f},建议停止 SFT") break ``` ```

2.2 pass@k(采样命中率)

为什么重要pass@1 衡量”一次答对”,pass@k 衡量”k 次采样内至少答对一次”。二者的差距揭示了模型的潜在能力(capability)与稳定性(reliability)

  • pass@1 低但 pass@k 高 → 模型”能做对,只是不稳定” → 这是 RL/偏好优化的甜点区,继续 SFT 收益低。
  • pass@k 本身随 SFT 下降 → 说明采样空间被压缩,属于过度 SFT 的危险信号。

无偏估计公式(Chen et al., 2021, Codex 论文):

对每题采样 n 个样本,其中 c 个正确,则:

\text{pass@}k = 1 - \frac{\binom{n-c}{k}}{\binom{n}{k}}

实现:

``` ``` import numpy as np from itertools import accumulate def pass_at_k(n: int, c: int, k: int) -> float: """ n: 每题总采样数 c: 其中正确的样本数 k: pass@k 的 k 使用无偏估计,数值稳定版本。 """ if n - c < k: return 1.0 return 1.0 - np.prod(1.0 - k / np.arange(n - c + 1, n + 1)) def eval_pass_at_k(model, dataset, k_list=(1, 5, 10), n_samples=20, temperature=0.8): """ 对每道题采样 n_samples 次,计算各 k 的 pass@k。 checker(prompt, completion) -> bool,任务相关的正确性判定。 """ results = {k: [] for k in k_list} for item in dataset: completions = model.generate( item["prompt"], n=n_samples, temperature=temperature ) c = sum(item["checker"](item, out) for out in completions) for k in k_list: results[k].append(pass_at_k(n_samples, c, k)) return {f"pass@{k}": float(np.mean(v)) for k, v in results.items()} # 判断逻辑:pass@1 低而 pass@k 高 → 转 RL def should_switch_to_rl(metrics, gap_threshold=0.15): gap = metrics.get("pass@10", 0) - metrics.get("pass@1", 0) return gap > gap_threshold # True: 建议转 RL 而非继续 SFT ``` ```

2.3 多样性(Diversity / Distribution Collapse)

为什么重要:过度 SFT 会让模型在相同 prompt 下反复输出几乎一样的答案(mode collapse)。多样性是 RL 阶段”探索”的燃料,一旦坍塌,RL 几乎无法产生有效梯度。

常用度量:

  • Distinct-n:生成文本中 unique n-gram 占比。
  • Self-BLEU:同一 prompt 多次采样间的 BLEU,越高越同质(多样性越低)。
  • 采样熵 / logits 熵:token 分布的平均熵。

详解 1:Distinct-n(词汇/表达多样性)

  • n-gram 是什么:把文本按词切分后,连续 n 个词组成的片段。例如句子 我 喜欢 吃 苹果:

      - 1-gram(unigram):我、喜欢、吃、苹果
    • 2-gram(bigram):我 喜欢、喜欢 吃、吃 苹果
    - Distinct-n 的定义:Distinct-n = 不重复的 n-gram 数量 / n-gram 总数量(在一批生成文本上统计)。
  • 直观含义:衡量”用词/搭配”重复的程度。值越接近 1,说明几乎没有重复的 n-gram → 表达越丰富;值越低,说明大量 n-gram 反复出现 → 表达越单调。
  • 举例(假设有两条生成结果,统计 2-gram):

      - 多样:今天 天气 很 好 + 我 想 出去 玩 → 7 个 bigram 全不重复 → Distinct-2 = 7⁄7 = 1.0
    • 单调:好 的 我 明白 + 好 的 我 知道 → 6 个 bigram 中 好 的、的 我 各重复一次 → unique=4,total=6 → Distinct-2 = 4⁄6 ≈ 0.67
    - 在 SFT 中的作用:随训练推进若 Distinct-n 持续下降,说明模型开始用固定套话/模板 → 采样空间收缩,是过度 SFT 的信号。

详解 2:Self-BLEU(样本间相似度 / 同质化)

  • BLEU 是什么:机器翻译常用指标,衡量一句”候选文本”与”参考文本”的 n-gram 重合度,取值 0~1,越高越相似。通常综合 1~4 gram。
  • Self-BLEU 的做法:对同一个 prompt 采样出多条回答后,把其中每一条当作”候选”,把其余所有条当作”参考”,算 BLEU,再对所有条取平均。

      - 直白说:就是”这批回答彼此之间有多像”。
    - 含义:
      - Self-BLEU 高 → 多条回答互相很像(措辞高度重合)→ 模型对同一问题只会给几乎一样的答案 → 多样性低。
    • Self-BLEU 低 → 多条回答各不相同 → 多样性高。
    • 注意它和 Distinct-n 方向相反:Distinct-n 越高越好,Self-BLEU 越低越好。

    - 举例(对同一问题采样 3 条):

      - 都答 你可以先重启一下试试 / 你可以先重启看看 / 建议你先重启一下 → 措辞高度重合 → Self-BLEU 高(同质化严重)。
    • 分别答 重启设备 / 检查网络连接 / 更新到最新版本 → 几乎无重合 → Self-BLEU 低(多样性好)。
    - 在 SFT 中的作用:训练过程中若 Self-BLEU 持续上升,说明模型对同一 prompt 的输出越来越趋同(mode collapse)→ 后续 RL 缺乏可探索的候选,收益下降。
一句话对比:Distinct-n 看”一批文本内部用词有多不重复”,Self-BLEU 看”同一问题的多个回答彼此有多像”。 前者越高越好,后者越低越好,二者互为佐证。

实现:

``` ``` from collections import Counter import numpy as np def distinct_n(texts, n=2): """越接近 1 越多样;SFT 过程中若显著下降 = 采样空间收缩。""" total, uniq = 0, set() for t in texts: toks = t.split() grams = list(zip(*[toks[i:] for i in range(n)])) total += len(grams) uniq.update(grams) return len(uniq) / max(total, 1) def self_bleu(samples): """同一 prompt 的多个采样两两 BLEU 均值,越高 = 越同质。""" from nltk.translate.bleu_score import sentence_bleu, SmoothingFunction sm = SmoothingFunction().method1 scores = [] for i, hyp in enumerate(samples): refs = [s.split() for j, s in enumerate(samples) if j != i] scores.append(sentence_bleu(refs, hyp.split(), smoothing_function=sm)) return float(np.mean(scores)) if scores else 0.0 def sampling_entropy(prompt, model, n=20, temperature=1.0): """对同一 prompt 多次采样,统计答案分布熵。熵下降 = 多样性坍塌。""" outs = model.generate(prompt, n=n, temperature=temperature) counts = Counter(outs) p = np.array(list(counts.values()), dtype=float) p /= p.sum() return float(-(p * np.log(p + 1e-12)).sum()) ``` ```

判断信号:随训练推进,distinct-n / 采样熵单调下降、self-BLEU 单调上升 → 分布正在坍塌,应考虑停止 SFT。

2.4 协议稳定性(Protocol / Format Compliance)

为什么重要:协议类能力(JSON schema、工具调用格式、chat template、拒答格式)是 SFT 最擅长且应该收敛到 100% 的能力。它是能力的”下限保证”——协议不稳定则一切下游解析都会崩。

怎么度量:格式合规率(format pass rate)——生成结果能否被严格解析 / 通过 schema 校验。

实现(以 JSON schema 与工具调用为例):

``` ``` import json from jsonschema import validate, ValidationError def json_schema_pass_rate(model, dataset, schema): ok = 0 for item in dataset: out = model.generate(item["prompt"], temperature=0.0)[0] try: obj = json.loads(out) validate(instance=obj, schema=schema) ok += 1 except (json.JSONDecodeError, ValidationError): pass return ok / len(dataset) # 期望收敛到 ~1.0 且稳定不抖动 def tool_call_pass_rate(model, dataset): """校验工具调用是否满足:函数名合法、参数齐全、类型正确。""" ok = 0 for item in dataset: out = model.generate(item["prompt"], temperature=0.0)[0] try: call = json.loads(out) name_ok = call["name"] in item["allowed_tools"] args_ok = set(item["required_args"]).issubset(call.get("arguments", {})) ok += int(name_ok and args_ok) except Exception: pass return ok / len(dataset) ``` ```

判断信号:协议合规率应快速升到接近 1.0 并稳定。若已饱和(例如连续多个 ckpt 都 ~0.99),说明协议能力这一维度上 SFT 已到头。

2.5 RL Probe Gain(RL 探针增益)

为什么重要:SFT 的最终目的(若后续要 RL)是提供一个保留优化空间的初始化。直接的检验方法是:拿当前 SFT ckpt 做一次短程 RL 探针(几百 step 的 PPO/GRPO/DPO),看 reward 能否明显上升。

  • 探针增益大 → 说明还有优化空间,当前 SFT ckpt 是好的 RL 起点。
  • 探针增益接近 0 → 策略已被压死,说明 SFT 过头(或该 ckpt 不适合做 RL 起点)。

实现(轻量 GRPO 探针的度量骨架):

``` ``` def rl_probe_gain(sft_ckpt, reward_fn, prompts, steps=200): """ 对候选 SFT ckpt 跑一段短 RL,返回 reward 提升幅度。 用于对比多个 ckpt,选 probe_gain 明确为正的那个。 """ policy = load_policy(sft_ckpt) base_reward = evaluate_reward(policy, prompts, reward_fn) trainer = GRPOTrainer(policy=policy, reward_fn=reward_fn, kl_coef=0.05) for _ in range(steps): trainer.step(sample_prompts(prompts)) final_reward = evaluate_reward(policy, prompts, reward_fn) return { "base_reward": base_reward, "final_reward": final_reward, "probe_gain": final_reward - base_reward, # 明确为正 = 优秀的 RL 起点 } def select_best_ckpt(ckpt_metrics): """ 综合选 ckpt:协议稳定 + OOD 不退化 + pass@k 高 + probe_gain 明确为正。 而非简单取 SFT 分数最高者。 """ def score(m): return ( 1.0 * m["protocol_pass_rate"] # 协议稳定 + 1.0 * m["ood_score"] # OOD 不退化 + 1.0 * m["pass@k"] # 保留能力 + 2.0 * max(m["probe_gain"], 0) # RL 优化空间(加权更高) ) return max(ckpt_metrics, key=score) ``` ```

2.6 综合监控面板(汇总判断)

``` ``` def sft_stop_decision(m): """ 综合五项信号给出停止建议。 m 为单个 ckpt 上聚合的指标字典。 """ signals = { "ood_saturated": m["ood_stall_steps"] >= 3, # OOD 停滞 "diversity_drop": m["distinct_2"] < m["distinct_2_init"] * 0.85, # 多样性坍塌 "protocol_ready": m["protocol_pass_rate"] > 0.98, # 协议已稳定 "passk_gap": (m["pass@10"] - m["pass@1"]) > 0.15, # 该转 RL "probe_flat": m["probe_gain"] < 0.01, # 无 RL 空间(过头) } stop = signals["ood_saturated"] and signals["protocol_ready"] switch_to_rl = signals["passk_gap"] and not signals["probe_flat"] over_sft = signals["diversity_drop"] or signals["probe_flat"] return { "stop_sft": stop, "switch_to_rl": switch_to_rl, "over_sft_warning": over_sft, "signals": signals, } ``` ```

三、常见任务的判断方法、指标与数据集

不同任务对”SFT 到位”的判定重心不同。下表先给全局对照,再逐一展开。

任务首要指标辅助指标常用数据集
Chat / 通用对话人类/模型偏好胜率、指令遵循率多样性、拒答合规、多轮一致性MT-Bench、AlpacaEval 2.0、Arena-Hard、IFEval
Agent / 工具调用工具调用格式合规率、任务成功率、pass@k轨迹长度、无效调用率、多轮状态一致性ToolBench、BFCL、τ-bench、AgentBench、WebArena
视频生成人类偏好 / VBench 维度分时序一致性、文本对齐、运动质量VBench、EvalCrafter、内部人评集

3.1 Chat / 通用对话

判断方法:Chat 任务没有唯一正确答案,核心看”是否符合人类偏好 + 是否遵循指令“。SFT 到位的标志是:偏好胜率与指令遵循率提升趋缓,同时多样性未坍塌、拒答/安全协议稳定。

关键指标:

  • 偏好胜率(Win Rate):用 GPT-4/裁判模型或人评,对比当前 ckpt vs 参考模型的胜率(AlpacaEval 2.0、Arena-Hard 的 LC win-rate)。
  • 指令遵循率(IFEval):可自动验证的指令约束(如”用 3 个 bullet”、”包含关键词 X”)的通过率。
  • 多轮一致性 / role consistency:多轮中是否保持人设与上下文。
  • 拒答合规率:不安全请求的正确拒答率(协议类,应稳定接近 1.0)。

实现要点(IFEval 风格自动校验):

``` ``` def ifeval_score(model, dataset): """dataset 每项含可编程校验的约束列表 constraints。""" total, passed = 0, 0 for item in dataset: out = model.generate(item["prompt"], temperature=0.0)[0] for c in item["constraints"]: total += 1 passed += int(c["verify_fn"](out)) # 例如 len(bullets)==3 return passed / total ``` ```

常用数据集:MT-Bench(多轮打分)、AlpacaEval 2.0(LC win-rate)、Arena-Hard-Auto、IFEval(指令遵循)、以及内部人评集。

到位信号:win-rate / IFEval 曲线走平,distinct-n 不再下降,拒答合规稳定 → SFT 到位;若 win-rate 停但输出明显模板化 → 过头,转 DPO/RLHF。

3.2 Agent / 工具调用

判断方法:Agent 任务是协议 + 多步决策的结合。SFT 主要解决”格式/协议对不对”(该收敛到近 100%),而”多步路径优不优”往往要靠 RL。因此:协议合规率饱和 + pass@1 与 pass@k 出现明显差距,就是”SFT 该收、转 RL”的典型信号。

关键指标:

  • 工具调用格式合规率:函数名、参数、JSON 结构是否合法(见 2.4)。
  • 任务成功率(Task Success Rate):端到端完成任务的比例(τ-bench、WebArena)。
  • pass@k:多次采样能否命中正确工具链(见 2.2)。
  • 无效调用率 / 幻觉工具率:调用不存在的工具或参数缺失的比例。
  • 多轮状态一致性:跨轮次是否正确维护环境状态。

实现要点(端到端轨迹评测):

``` ``` def agent_task_success(model, env_list, max_steps=15): success = 0 for env in env_list: obs, done, steps = env.reset(), False, 0 while not done and steps < max_steps: action = model.act(obs) # 输出工具调用 if not is_valid_tool_call(action, env.tools): break # 无效调用 -> 失败 obs, done, info = env.step(action) steps += 1 success += int(info.get("task_completed", False)) return success / len(env_list) ``` ```

常用数据集 / 基准:ToolBench、BFCL(Berkeley Function Calling Leaderboard)、τ-bench(tau-bench,含 retail/airline 多轮环境)、AgentBench、WebArena / VisualWebArena、GAIA。

到位信号:格式合规率饱和(~1.0)、任务成功率提升趋缓,但 pass@k − pass@1 差距大 → SFT 已到头,转 GRPO/PPO 提升决策质量。

3.3 视频生成

判断方法:视频生成的”标注答案”不唯一,且质量维度多(画质、时序一致性、文本对齐、运动)。因此不能只看重建/扩散 loss。SFT(或指令微调阶段)到位的标志是:VBench 各维度分与人类偏好提升趋缓,且不以牺牲多样性/运动幅度为代价(常见过拟合表现为”画面变干净但运动僵硬、镜头趋同”)。

关键指标:

  • VBench 维度分:主体一致性(subject consistency)、背景一致性、时序闪烁(temporal flickering)、运动平滑度(motion smoothness)、美学质量(aesthetic quality)、成像质量(imaging quality)、动态程度(dynamic degree)、文本对齐(对 T2V)。
  • 文本-视频对齐:CLIP-Score / ViCLIP / 内部对齐评分。
  • 时序一致性:相邻帧特征相似度、光流一致性。
  • 人类偏好:成对人评胜率(最终仲裁指标)。

实现要点(时序一致性 + 文本对齐骨架):

``` ``` import torch, torch.nn.functional as F def temporal_consistency(frames, feature_extractor): """相邻帧特征余弦相似度均值,越高越连贯(但过高可能=运动僵硬)。""" feats = torch.stack([feature_extractor(f) for f in frames]) # [T, D] sims = F.cosine_similarity(feats[:-1], feats[1:], dim=-1) return sims.mean().item() def clip_text_video_alignment(frames, text, clip_model): """逐帧 CLIP 相似度均值,衡量文本对齐。""" txt = clip_model.encode_text(text) scores = [F.cosine_similarity(clip_model.encode_image(f), txt, dim=-1) for f in frames] return torch.stack(scores).mean().item() ``` ```

常用数据集 / 基准:VBench / VBench++(细粒度多维评测)、EvalCrafter、UCF-101 与 Kinetics(FVD 计算的参考分布)、WebVid / Panda-70M(训练侧)、以及内部 prompt 人评集。

到位信号:VBench 综合分与人评胜率走平;若继续训练导致 dynamic degree(动态程度)下降、镜头/构图趋同(多样性坍塌)→ 过拟合,应停止并转偏好优化(如视频 DPO / reward 微调)。


四、一句话总结

判断 SFT 是否到位,本质是从”拟合视角”切换到”能力 + 可优化性视角”: 当 train loss 还在降,但 OOD/hard 不涨、多样性开始坍塌、协议已稳定、pass@k 明显高于 pass@1、RL 探针仍有增益时——就该停止 SFT,把剩下的”哪个更优”交给 RL。 SFT 管”会不会按正确方式做”,RL 管”多个可行方案里哪个最好”。
分享到
QQ 微信 微博 复制链接
微信分享二维码

微信扫一扫,分享给好友

本文来自投稿,不代表本站立场,如若转载,请注明出处:

发表评论

发表评论

作者信息

va
作者有点忙,还没写简介
TA的最新作品
    请配置好页面缩略名选项

推荐话题

关于本站

来跟我聊聊吧~