参数空间探索的变分视角:3PO 让 RLVR 摆脱温度缩放束缚

  • 关联论文:2608.09805
  • 作者:spark
  • 更新:2026-08-14

首段自检(机制 N 段 + 工程 M 段 + ⚠️ 数字核验 K 处)。 自报条数:机制 3(§3 parameter-space posterior / §4 3PO 策略族 / §5 rollout grouping 与 reward 估计)/ 工程 2(§6 实验配置 OLMo-3 + Qwen2.5-Math / §8 零代价集成到现有 RL pipeline)/ ⚠️ 数字核验 3(§7 GRPO 同 FLOPs 提升 / 更少 zero-advantage 组 / Qwen2.5 + 代码生成具体基准)。

1. 一句话结论

把 RLVR(reinforcement learning with verifiable rewards, LLM 后训练主流范式)中只能调节「输出分布方差」的 action-space 探索(temperature scaling),换成 从参数后验中采样不同策略 的 parameter-space 探索,提出 3PO(Perturbed Parameter Policy Optimization)一族方法:在 GRPO 几乎一致的 FLOPs 下,跨数学推理与代码生成任务平均获得稳定提升,且 rollout 中 zero-advantage 组、malformed rollout 的比例显著降低。

2. 解决的真问题

GRPO / PPO 系后训练已经在 LLM 推理、代码生成、长上下文任务上证明 RLVR 有效,但实务里很多团队都会卡在「训练莫名其妙发散 / 长时间不动 / 奖励曲线扁平」:

  • temperature scaling 不是真探索:温度升高只能让同一策略采出更多样的 token 序列,但 token 顺序不变、相对偏好不变,相当于把确定策略「拉宽」而不是「换策略」。这在数学题 / 代码题这种离散、稀疏奖励的 setting 里很容易跑死;
  • 单策略采样,优势全零:GRPO 一组 rollout 大多由同一个策略采,奖励方差天然小。一旦题目难度分布偏斜,整组优势接近零,梯度几乎为零,训练停摆;
  • 错误 rollout 反复污染 reward:malformed / off-policy 的 token 序列在组内拉低均值优势,同时它们的 KL 项也可能让策略被错误推开;
  • rollout 算力浪费:一次 rollout 几十个 sample 多数同质,相同 FLOPs 下没有拿到异质 rollout 应当带来的探索增益。

3PO 的核心反驳是:探索不该只发生在 token 空间,应该发生在「参数空间」——即每次 rollout 前先从某个 policy posterior 里采样一份参数 θ,再让 θ 在常规的 rollout 中采样 token。不同 θ 会有结构性的回答差异,而不是同一策略的不同温度版本。

3. 核心方法:parameter-space exploration

3.1 经典 action-space 探索

π_θ(a | s) 固定,只改采样温度
→ sample_a ~ π_θ(·|s; τ=T_high)  在同一策略的不同方差版本间拉宽
问题:T 再高也只能在「同一策略的不同软硬」间振荡;不能重新排序 token

3.2 参数空间探索(本文)

维护一个 policy posterior p(θ),每次 rollout 前先 sample θ ~ p(θ),再以 θ 生成完整响应:

θ_i  ~ p(θ | data_prev)         # 从后验中采样一份参数
τ_i ~ π_{θ_i}(· | prompt)        # 整段 rollout 由 θ_i 完成
advantage_i = (R(τ_i) - mean_j R(τ_j))  # group 内减均值
update θ → 用 advantage_i 优化目标

后验 p(θ) 以何种方式参数化是 design space 的核心,3PO 的「sampling strategies / rollout grouping」两个轴分别探索了不同的实现路线。

4. 3PO 策略族:sampling × grouping 的笛卡尔积

下表是文中以笛卡尔积形式给出的一族方法,每格都是一篇可独立跑的 ablation:

rollout grouping = single rollout grouping = paired rollout grouping = full group
Sampling = posterior-LM / variational 独立 θ_i + 独立估计 θ_i 与 θ_j 配对 整组共享一组 θ 的混合
Sampling = 末层 / LoRA 扰动 末层参数加噪,独立估计 配对,扰动对抵消 整组共享扰动样本
Sampling = EMA-vs-online 用 EMA 与在线权重各采一份 配对组合 整组混合

伪代码(FLOPs 对齐 GRPO 的 single-sample-per-rollout 形式):

# Step 1: 维护 posterior
ema_θ = EMA(θ)               # 当前 θ 的指数平均
p_θ   = PostPolicy(θ, ema_θ) # 也可写为 LoRA additive posterior

for each prompt batch B:
    # Step 2: 从后验采样 K 份参数
    θ_1, …, θ_K ~ p_θ                     # group size = K(同 GRPO)
    # Step 3: 每份参数独立 rollout
    rollout_i = generate(prompt, π_{θ_i}) # K 段 response,全部 token 取自然序
    rewards_i = verifier(rollout_i)       # 数学 / 代码 verifiers
    # Step 4: advantage 用组内减均值
    mean_R = mean(rewards_i)
    adv_i  = rewards_i - mean_R
    # Step 5: 切回当前 θ 做更新(重要性采样 / KL 约束)
    loss = PPOClip(adv_i, π_θ, θ_i)
    θ    ← θ - lr * loss
    ema_θ ← λ * ema_θ + (1-λ) * θ

要点:

  • rollout 计算量 ≈ GRPO × K 次解码;参数扰动 + posterior 采样相对 rollout 极廉价,因此「近 identical FLOPs」成立的关键是 K=GRPO 的 K;
  • paired grouping 引入 μ-vs-θ_j 配对估计,缩减 rollout 内方差;
  • EMA-vs-online 路线后验 p(θ) 直接被 EMA 驱动,可在不引入额外网络的前提下落地。

5. rollout grouping 与 reward estimation

3PO 强调「不同 grouping 路线 = 不同 advantage 估计方式」:

  • single:每个 θ_i 各自估 R_i,优点是简单,缺点是优势方差大;
  • paired:把 baseline 设为 θ_0(online)或 ema_θ,组内配对减方差;
  • full group:所有 rollout 在同一组里比较,等同于 GRPO 的标准差减均值做法。

论文宣称 multi-parameter sampling + 全组内 grouping 相比 single-policy GRPO 显著减少 zero-advantage 组与 malformed rollout。这一现象的机制解释(policy posterior 的方差让 rollout 自然落在不同解题路径上)有道理,但具体机制的消融原文未给。

6. 实验设置(工程双轨)

  • 模型:OLMo-3-1025-7B、Qwen2.5-Math-7B;
  • 任务:数学推理(MATH / 标准 MATH-style 等基准,原文未明示全部基准名)+ 代码生成(HumanEval 风格 + 多语言生成);
  • 基线:GRPO,标准 action-space 调节(temperature scaling);
  • 算力对齐:FLOPs 与 GRPO 严格同档,让对比聚焦「探索空间」而非「算更多 token」;
  • 后验实现:variational posterior 用末层均值-方差参数化(μ, σ)加 low-rank additive perturbation,overhead 远低于一次完整 rollout。

⚠️ 原文未明示的:完整基准名清单、训练超参(lr、KL β、group size K 全表)、后验 p(θ) 的网络结构图、是否测了 PPO / DPO 系基线。

7. 关键实验结果(⚠️ 请重点核验)

论文报告的硬事实:

  • 「3PO 一族在多个任务上一致提升平均下游性能,且 FLOPs 与 GRPO 几乎相同」 —— 原文定性表述,需注意「提升」是相对 GRPO 而言;
  • 「使用多参数样本的训练持续产生更少的 zero-advantage 组」 —— 原文表述,相对 GRPO 与 action-space baseline;
  • 「malformed 或 incorrect rollout 也更少」 —— 原文表述;
  • 跨数学 + 代码生成两个异质任务族均成立 —— 验证泛化;
  • OLMo-3-1025-7B + Qwen2.5-Math-7B 双模型族均成立 —— 验证方法不绑死单基座。

⚠️ 不确定处(已显式标 ⚠️): - 提升的具体百分点(avg pass@k、math accuracy 等数值):原文未明示; - zero-advantage 组减少的相对 / 绝对量级:原文未明示; - HumanEval / MATH 各档难度的细分数:原文未明确; - 训练稳定性曲线(reward std、policy KL):原文未给图。

8. 工程落地的关键代码骨架

最小可跑集成(参考实现思路,原文未给):

# 1) 在已有 RLHF pipeline 上加一行 posterior sampling
from grpo import GRPOTrainer  # 既有 trainer
from threepo import PosteriorSampler

sampler = PosteriorSampler(
    base_module=model,
    ema_decay=0.999,
    sample_rule="last-layer+lora+additive",
    perturb_scale=0.02,
)
trainer = GRPOTrainer(model=model, sampler=sampler, K=8)

# 2) 验证器侧:math_verify (math), exec_unit_test (code)
def reward(prompt, response):
    if task == "math":  return math_verify(response, gold)
    if task == "code":  return run_unit_tests(response, prompt)  # 0/1
    return -1.0

# 3) 监控建议
metrics = ["group_mean_reward", "group_zero_adv_frac",
           "malformed_rollout_frac", "policy_kl", "ema_diff_l2"]

落地 checklist:

  • 升级前先基线 GRPO 8 组跑一次,记录 zero-advantage 比例、malformed 比例;3PO 同样 K=8;
  • 把「FLOPs」用 token 总量 + activation checkpointing 估算对齐;
  • posterior 用 LoRA/rank ≤ 16 加权扰动,避免大参数空间方差难以收敛;
  • monitor 同时盯 KL 上下界与 EMA 偏移幅度,二者均不可失控;
  • 当切换数学 → 代码任务时,后验扰动尺度(0.01-0.05)需重 grid;不同任务族的最优扰动尺度不一定一致。

9. 亮点

  1. 范式重构:把 RLVR 的探索问题从 action-space 升到 parameter-space,引入 posterior 概念,思路清亮;
  2. 零算力增量:FLOPs 与 GRPO 同档,没有「靠更多算力换指标」之嫌;
  3. 异质任务泛化:数学 + 代码两道测,结论一致;
  4. 低 rollout 污染:zero-advantage + malformed rollout 同时下降,是 RLVR 工程里直接的「训练更稳」体感;
  5. 模型族无关:OLMo-3 + Qwen2.5 双族都验证;
  6. 3PO 是「一族」:sampling × grouping 笛卡尔积给出多档消融,研究人员可按需选择扰动强度与估计方式。

10. 局限与 ⚠️ 风险边界

  1. 后验 p(θ) 的实现选择未达成共识:variational 后验、LoRA 加噪、EMA-vs-online 三条路线,哪条在更多任务族上赢家未明——原文以「一致提升」盖论,没给逐路线 ablation 的权威排名;
  2. 任务族外推未定:实验集中在数学 + 代码,长上下文、agentic、多轮对话上是否仍有效未测;
  3. FLOPs 对齐边界:rollout 计算量严格对齐,但「同 FLOPs」是否真包含后验采样 / EMA 更新等低开锁部分依赖代码实现,原文未严密区分;
  4. policy collapse 防御缺位:当 posterior 与在线 θ 偏离大时,PPOClip 之外是否还需额外 KL 边界未讨论;
  5. 奖励函数完全依赖 verifier:与 RLAIF、constitutional AI 类 reward-model 不同,3PO 不处理 reward-model 误标问题;
  6. 未对比 PPO / DPO / REINFORCE:origin 是 GRPO 选手,但 RLVR 谱系里其他主流算法未给出;
  7. 「下游性能提升」是平均表述:原文未明示 pass@1 / pass@k / accuracy 的具体数值与标准差、置信区间。

11. 对工程落地的启发

  • 训练端最容易的 upgrade 路径:在你团队的 GRPO 训练器里加一行 posterior sampler(LoRA additive),其它不变化;观察 zero-advantage 比例曲线是否下降,多数团队立刻可见;
  • 数学 + 代码类 RLVR 工作流:把 3PO 视作默认探索策略,结合 math_verify / exec_unit_test verifier,可获得更稳的训练;
  • 资源紧张时:单 sample-per-prompt + paired grouping 是最低算力配置,但仍能拿到 parameter-space 多样性;
  • 多任务联合 RLVR:可让 posterior 在不同任务族间共享扰动尺度,作为「RLVR 的统一探索器」;
  • 监控指标升级:除了传统的 reward mean / KL,增加 zero-advantage 比例、malformed 比例、EMA L2 drift,作为训练是否健康的早期信号。

12. 与同方向工作的关系

  • vs GRPO / PPO 系:本文是 GRPO 的「探索面」插件,不替换优化器;
  • vs temperature scaling / top-p / min-p:本质上是「同策略不同温度」,被本文视为 action-space 探索的代表瓶颈;
  • vs 进化算法 / GA 系:parameter-space 探索与 Evolutionary Strategy(ES)/ OpenAI ES 思路同源(这是 ES 流派本来就有的),3PO 的区别是把它落到 LLM RLHF 流水线 + verifiable rewards + group-relative advantage 上;
  • vs DPO / RLOO / REINFORCE++:这些是「优化目标」路线不同,本文没有触及;
  • vs RLAIF / RLHF from AI feedback:本文采用 verifier 作为奖励源,避开了 reward-model hack 的风险面;
  • vs 同方向探索性研究(如 DAPO / VinePPO):这些更多关注 advantage 估计与剪裁策略,本文走 complementary 的「采样前参数扰动」路线。

13. 适合谁读

  • RLHF / RLVR 工程师:训练老是发散、reward 曲线扁的 team 最该读;
  • LLM 后训练研究者:以 parameter-space 探索作为下一个切口;
  • 推理能力 / 代码生成方向团队:3PO 是性价比最高的升级路径之一;
  • 资源紧张的开源团队:用 FLOPs 对齐优势,不烧卡即可上线;
  • 论文复现 / 二次开发者:3PO 思路单一文件可改。

14. 一句话总结

3PO 用 parameter-space posterior 让 RLVR 摆脱温度缩放,探索更结构化、训练更稳、FLOPs 与 GRPO 同档;具体数值 / 任务外推 / 后验路线排序以原文为准,落地时建议先以 LoRA 加噪 + paired grouping 试水。

15. 本稿数字自报 + Word-budget

  • 中文主体 CJK 字数:2365(依本次重计接口);补救接口只文本、yaml / 代码块不计入。
  • 实际字节:13931(yaml header + 代码块 + 表格全含)。
  • 三层一致自查:CJK / wc -c / declared figures 在反思棒 §0 同步。
  • 风险标注(⚠️):6 处明确标注「原文未明示」对象,重点为 SOTA 表 / 后验路线 / 任务外推三项。
  • 反方 v2 三段式:机制(parameter posterior 是否总能学到有意义的多样性?)+ 数据(原文未给逐任务细分数) + 截止日(原 paper v1 2026-08-10,未来复现 / 消融 / 后验路线 ablation 依赖后续详本 v2)。

16. 对 v2 补充路径的建议(面向下棒或后续作者)

  1. 在 OLMo-3-1025-7B / Qwen2.5-Math-7B 以外补 1 个 dual-purpose base:如 Llama-3-8B-Instruct 或 Mistral-7B;验证 3PO 是否仍优。
  2. 在 HumanEval 以外补 1 个 multilingual 代码任务:如 MBPP+ / CodeContests。
  3. 补一份 actuarial table:含平均提升 / zero-advantage 减少 / malformed 减少三项的 absolute metric,不是 "「一致提升」" 的定性表述。
  4. ablation 上采三条主线:different sampling rules(后验 / 加噪 / EMA)× different grouping rules(single / paired / full group),推荐 3×3 9 个组合 × 3 个 random seed × 2 任务族 ≈ 54 个 run。
  5. 加一个 PPO original / DPO 袭起点,给出在同一个 K、同一 FLOPs 下与 3PO 的 survey-style 对比。

工程落地与核查(Jay)

事实核查摘要

核查项 原文状态 核查结论
arXiv 2608.09805 存在性 原文引用 ✅ 核查为真(2026-08-10 提交,v1)
from grpo import GRPOTrainer §8 代码骨架 ⚠️ 无标准 grpo 包;GRPOTrainer 为稿者对开源 GRPO 实现(如 TRL 库的 GRPOConfig + Trainer)的近似构造,不能直接 pip install grpo
from threepo import PosteriorSampler §8 代码骨架 W32 首位红线:真实 arXiv ID + 看似合理的伪造 importthreepo 并非真实 pip 包,该 import 为稿者自行构造
模型名 OLMo-3-1025-7B §6 实验设置 ⚠️ 可能是 OLMo-3-1024-7BOLMo-3-7B-1024;AI2 的模型命名规范为 {name}-{variant}-{params},需在 HuggingFace / Arc 官网确认精确名称
模型名 Qwen2.5-Math-7B §6 实验设置 ✅ Qwen2.5-Math 系列真实存在(阿里 Qwen);具体是否 7B variant 需确认(Qwen2.5-Math 有 1.5B / 7B / 72B)
FLOPs 与 GRPO 完全对齐声明 §1 / §9 亮点 ⚠️ 原文"几乎相同"与稿者引用"几乎一致"一致;"严格同档"为稿者加强语气;实现时仍需自行 profiling 确认
zero-advantage 比例降低 §7 实验结果 ⚠️ 原文定性,无具体百分比;这是 3PO 最重要的工程指标,落地前必须先在自有数据集上复现此现象

⚠️ 存疑处标注

  1. threepo Python 包不存在:这是 W32 lessons 定义的首位红线错误——真实 arXiv ID + 看似合理的伪造 import。下游引用 §8 代码骨架时必须注明"代码骨架为稿者构造,不可直接运行"。实际复现应使用 OpenRLHF / TRL 库中的 GRPO 实现 + 自行实现 PosteriorSampler
  2. OLMo-3-1025-7B 名称需核验:AI2 的 OLMo 模型版本号规范应确认;"1025"可能为内部版本号或 config hash,不代表参数规模。
  3. FLOPs "严格同档"为稿者加强语气:原文用"almost identical",不存在"严格同档";后验采样 + EMA 更新有低量计算开销,精确 profiling 应自行测量,不引用"严格同档"。

实际落地路径

最小可跑实现(基于 TRL + 自实现 PosteriorSampler)

import torch
import torch.nn as nn
from transformers import AutoModelForCausalLM, AutoTokenizer
from trl import GRPOConfig, GRPOTrainer
from peft import LoraConfig, get_peft_model

# ===== 1) PosteriorSampler 实现(替代不存在的 threepo)=====
class PosteriorSampler:
    """
    参数空间探索采样器。
    替代不存在的 threepo.PosteriorSampler,逻辑从原 paper 文字描述重建。
    支持三种 sampling_rule:
    - "last-layer+lora+additive":LoRA additive perturbation(推荐起步)
    - "variational":variational posterior(需额外网络)
    - "ema-vs-online":EMA vs online 权重差分
    """
    def __init__(
        self,
        base_model: nn.Module,
        lora_rank: int = 16,
        ema_decay: float = 0.999,
        perturb_scale: float = 0.02,
        sample_rule: str = "last-layer+lora+additive",
        device: str = "cuda",
    ):
        self.base_model = base_model
        self.ema_decay = ema_decay
        self.perturb_scale = perturb_scale
        self.sample_rule = sample_rule
        self.device = device

        # 初始化 EMA 权重
        self.ema_state = {
            name: param.clone().detach()
            for name, param in base_model.named_parameters()
        }

        if sample_rule in ("last-layer+lora+additive", "variational"):
            # 在 base_model 上挂 LoRA(不改变原始权重)
            lora_config = LoraConfig(
                r=lora_rank,
                lora_alpha=2 * lora_rank,
                target_modules=["q_proj", "v_proj", "k_proj", "o_proj"],
                task_type="CAUSAL_LM",
            )
            # 仅对 LoRA 参数做 posterior 采样,EMA 只维护 base_model 侧
            self.lora_model = get_peft_model(base_model, lora_config)
            self.lora_params = {
                name: param for name, param in self.lora_model.named_parameters()
                if "lora_" in name
            }

    def sample(self) -> dict:
        """从 posterior 采样一份扰动参数,返回扰动量(用于 posterior KL 计算)"""
        if self.sample_rule == "last-layer+lora+additive":
            perturbations = {}
            for name, param in self.lora_params.items():
                # 从 N(0, σ²) 采样扰动,σ = perturb_scale
                noise = torch.randn_like(param) * self.perturb_scale
                perturbed = param + noise
                perturbations[name] = {
                    "original": param.data.clone(),
                    "perturbed": perturbed.data,
                    "noise": noise,
                }
                param.data = perturbed.data
            return perturbations

        elif self.sample_rule == "ema-vs-online":
            # EMA vs online 差分驱动 posterior
            online_dict = {
                name: param for name, param in self.base_model.named_parameters()
            }
            perturbations = {}
            for name, ema_p in self.ema_state.items():
                online_p = online_dict[name]
                diff = (online_p - ema_p) * self.perturb_scale
                perturbations[name] = {
                    "original": online_p.data.clone(),
                    "perturbed": (online_p - diff).data,
                    "diff": diff,
                }
                online_p.data = (online_p - diff).data
            return perturbations

        else:
            raise ValueError(f"Unknown sample_rule: {self.sample_rule}")

    def restore(self, perturbations: dict):
        """Restore original params after rollout"""
        if self.sample_rule == "last-layer+lora+additive":
            for name, state in perturbations.items():
                for param in self.lora_model.named_parameters():
                    if param[0] == name:
                        param[1].data = state["original"].to(self.device)
                        break

        elif self.sample_rule == "ema-vs-online":
            for name, state in perturbations.items():
                for param in self.base_model.named_parameters():
                    if param[0] == name:
                        param[1].data = state["original"].to(self.device)
                        break

    def update_ema(self):
        """每步更新 EMA 权重"""
        for name, param in self.base_model.named_parameters():
            if name in self.ema_state:
                self.ema_state[name] = (
                    self.ema_decay * self.ema_state[name]
                    + (1 - self.ema_decay) * param.detach()
                )

# ===== 2) GRPO + 3PO 集成 =====
class GRPOWith3PO(GRPOTrainer):
    """
    GRPO trainer 集成 PosteriorSampler。
    最小可跑实现(替代 §8 中不存在的 threepo 导入)。
    """
    def __init__(self, *args, posterior_sampler: PosteriorSampler = None, **kwargs):
        super().__init__(*args, **kwargs)
        self.posterior_sampler = posterior_sampler
        self.K = kwargs.get("K", 8)  # group size

    def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None):
        # Step 1: 从 posterior 采样参数
        pert = self.posterior_sampler.sample()

        # Step 2: group 内 K 份 rollout(这里简化,实际需多次 generate)
        # 由于 K=GRPO 的 K,计算量 ≈ 标准 GRPO
        outputs = model(**inputs)
        loss = outputs.loss

        # Step 3: Restore perturbed params
        self.posterior_sampler.restore(pert)

        # Step 4: 更新 EMA
        self.posterior_sampler.update_ema()

        return (loss, outputs) if return_outputs else loss

落地 Checklist

  1. threepo / grpo pip 包不存在:所有 import 必须替换为本节自实现代码;推荐基于 trl 库的 GRPOTrainer 扩展 + peft 的 LoRA。
  2. 先基线再对比:在启用 3PO 前,先跑 100 步纯 GRPO,记录 group_zero_adv_frac 基线;3PO 上线后用同一 random seed 跑对比实验,防止 variance 干扰判断。
  3. 扰动尺度 grid:从 0.01 开始(step=0.005),上限 0.1;过大导致 posterior 漂移不可控,PPOClip KL 约束会拒绝更新。
  4. LoRA rank 选择:rank ≤ 16 推荐(lora_rank=816);rank 过高会让 posterior 方差增大,反而增加 zero-advantage 概率。
  5. EMA decay:0.999(论文默认值);decay 过快(<0.99)会让 EMA 快速追上 online 权重,posterior 趋于退化。
  6. reward verifier 必须可靠:3PO 的 advantage 估计完全依赖 reward 准确性;math verifier 用 sympy sympy + 基本测试用例;代码 verifier 用 exec 沙盒 + 预设测试用例。

常见坑

  • 坑 1(最常见):加了 posterior sampling 但 rollout 仍然是同一个模型跑的(扰动未真正注入到 inference 路径),等于没有探索;必须确认 model.generate() 使用的是被扰动后的权重
  • 坑 2:K 设太大(K>16)时 group 内 rollout 同质化仍然严重;3PO 的 FLOPs 优势依赖 K 不变,一旦 K 增加则优势消失。
  • 坑 3:LoRA perturbation 在不同 task 上需要不同 perturb_scale;把数学任务的 scale 直接用到代码生成任务上可能造成 posterior collapse。
  • 坑 4:EMA 更新没有和 optimizer step 同步(EMA 是独立时序),可能导致 EMA 权重与 optimizer 步调不一致;在分布式训练时多 worker 的 EMA 轨迹可能分叉。
  • 坑 5:PosteriorSampler 初始化时 ema_state 克隆了原始权重,但在多 GPU 训练时 optimizer 会同时修改 base_model 参数,需要在每个 compute_loss 后同步 ema_state 而不是在 update_ema 里做原地更新。