Skip to content

Support Top-k MOPD loss #10079

Description

@mccatec

Checklist / 检查清单

  • I have searched existing issues, and this is a new feature request. / 我已经搜索过现有的 issues,确认这是一个新的 Feature Request。

Feature Request Description / Feature Request 描述

Feature Request: Support Top-K OPD-RL Distillation Loss for MOPD

Add support for the differentiable top-k OPD-RL distillation objective introduced in MOPD (arXiv:2606.30406, Eq. 5).

The implementation should compute a per-token, top-k-truncated reverse KL between the current policy and teacher distributions over the teacher's top-k token support, with the bias-correction term required by MOPD.

A straightforward reverse KL computed only over the teacher's top-k tokens is biased because probability mass outside the truncated support is discarded. In particular, the naively truncated objective is not necessarily minimized when the student's probabilities match the teacher's probabilities on the retained top-k tokens.

Renormalizing the truncated student and teacher distributions also changes the original full-vocabulary probabilities and therefore does not preserve the intended distillation objective.

MOPD addresses this by using the generalized KL divergence over the unnormalized top-k probabilities:

KL_t =\sum_{v \in \mathrm{TopK}}\left[p_s(v)\left(\log p_s(v) - \log p_t(v)\right) - p_s(v) + p_t(v) \right]

where both p_s and p_t are the raw temperature-1 full-vocabulary probabilities at the teacher's top-k token indices, without top-k renormalization.

The additional p_t(v) - p_s(v) term removes the truncation-induced bias. Each token-level summand becomes the Bregman divergence associated with (x logx), making it non-negative and minimized exactly when p_s(v) = p_t(v) for every retained teacher top-k token.

Proposed Behavior:

Provide a per-token loss function with inputs equivalent to:

  • teacher_topk_logprobs: teacher log probabilities at the teacher's top-k token indices, shape [B, T, K]
  • policy_topk_logps: current policy log probabilities evaluated at the same token indices, shape [B, T, K]
  • completion_mask: response-token mask, shape [B, T]

The function should return a [B, T] tensor containing the top-k OPD-RL divergence for each completion token.

The policy log probabilities must remain attached to the computation graph so that gradients flow through the student's probabilities. Teacher probabilities are fixed targets.

Pull Request / Pull Request 信息

No response

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    enhancementNew feature or request

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions