Skip to content

[DistillationTrainer refactor] Emit prompt_ids / prompt_mask / completion_ids#6487

Open
qgallouedec wants to merge 4 commits into
19-loss-consumes-completion-maskfrom
20-emit-prompt-completion-ids
Open

[DistillationTrainer refactor] Emit prompt_ids / prompt_mask / completion_ids#6487
qgallouedec wants to merge 4 commits into
19-loss-consumes-completion-maskfrom
20-emit-prompt-completion-ids

Conversation

@qgallouedec

@qgallouedec qgallouedec commented Jul 21, 2026

Copy link
Copy Markdown
Member

Item 20 (Group B) — stacked on PR 19. Part of #6449.

Third step of the layout migration (add → switch → delete). Emit GRPO's keys alongside the existing ones; nothing consumes them yet (item 21 switches consumers, item 22 deletes the old helpers).

  • Both generation callers and the collator now emit prompt_ids (= prompts), prompt_mask (= prompt_attention_mask), and completion_ids (the completion region of input_ids; empty for the prompt-only collator).
  • Buffer-side emission is throwaway (replaced by GRPO generation at item 26).

New test asserts cat(prompt_ids, completion_ids) == input_ids and that the prompt keys mirror the existing tensors.

Verified: 53 tests pass (distillation + server); ruff clean.


Note

Low Risk
Additive tensor aliases on batch dicts with no consumer changes; behavior is covered by a new test asserting key parity with existing tensors.

Overview
Adds GRPO-aligned batch keys to DistillationTrainer while keeping prompts / prompt_attention_mask unchanged—part of a layout migration (add → switch consumers → remove old keys).

The collator and both on-policy generation paths (model.generate and vLLM buffer updates) now include prompt_ids (same as prompts), prompt_mask (same as prompt_attention_mask), and completion_ids (the completion slice of input_ids; for prompt-only collation this slice is empty). Nothing in the trainer reads these keys yet.

A new integration test checks the keys are present on the batch passed to compute_loss and that torch.cat([prompt_ids, completion_ids], dim=1) matches input_ids.

Reviewed by Cursor Bugbot for commit 77712c0. Bugbot is set up for automated code reviews on this repo. Configure here.

@bot-ci-comment

Copy link
Copy Markdown

The docs for this PR live here. All of your documentation changes will be reflected on that endpoint. The docs are available until 30 days after the last update.

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

💡 Codex Review

Here are some automated review suggestions for this pull request.

Reviewed commit: b38280e33a

ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

Comment thread trl/experimental/distillation/distillation_trainer.py
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant