Skip to content

fix(gkd): defer full-vocabulary local teacher logits to forward_step #10097

Description

@wzkk123

Checklist / 检查清单

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

Bug Description / Bug 描述

For a separate local teacher with exact full-vocabulary logits
(gkd_logits_topk is None), _compute_teacher_logits_local eagerly materializes
one [S, V] teacher logits tensor per rollout micro-batch at batch
preparation
, and all of them stay alive for the whole generation cycle. With
64 queries × n=4 (256 rollouts), up to 256 full [S, 151k] logits tensors are
retained at once, OOMing long teacher responses (observed in practice).

How to Reproduce / 如何复现

Change

  • Split the per-batch eager loop into _compute_teacher_output_local, a
    single-micro-batch helper that forwards one local teacher microbatch and
    returns its TeacherOutput.
  • _compute_teacher_logits skips the eager materialization when
    gkd_logits_topk is None (full-vocabulary mode): the encoded batches keep
    teacher_model_inputs instead of teacher_output.
  • forward_step computes the teacher output just in time for the micro-batch it
    is about to train on, so only one [S, V] logits tensor is alive at a
    time. Fails closed with RuntimeError if a batch has neither an eager
    teacher_output nor usable teacher_model_inputs.
  • Compressed top-k logprobs ([S, K]) remain eager — they are small; the
    deferral targets the exact full-vocabulary path only.
  • Self-distillation behavior is unchanged (still recomputed per train step via
    _on_train_step_batch).

Additional Information / 补充信息

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

    bugSomething isn't working

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions