Checklist / 检查清单
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
Checklist / 检查清单
Bug Description / Bug 描述
For a separate local teacher with exact full-vocabulary logits
(
gkd_logits_topk is None),_compute_teacher_logits_localeagerly materializesone
[S, V]teacher logits tensor per rollout micro-batch at batchpreparation, 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 areretained at once, OOMing long teacher responses (observed in practice).
How to Reproduce / 如何复现
Change
_compute_teacher_output_local, asingle-micro-batch helper that forwards one local teacher microbatch and
returns its
TeacherOutput._compute_teacher_logitsskips the eager materialization whengkd_logits_topk is None(full-vocabulary mode): the encoded batches keepteacher_model_inputsinstead ofteacher_output.forward_stepcomputes the teacher output just in time for the micro-batch itis about to train on, so only one
[S, V]logits tensor is alive at atime. Fails closed with
RuntimeErrorif a batch has neither an eagerteacher_outputnor usableteacher_model_inputs.[S, K]) remain eager — they are small; thedeferral targets the exact full-vocabulary path only.
_on_train_step_batch).Additional Information / 补充信息
No response