Stop enable_padding_free_metadata writing seq_lengths into the caller's examples - #8049
vineethsaivs wants to merge 1 commit into
Conversation
There was a problem hiding this comment.
💡 Codex Review
Here are some automated review suggestions for this pull request.
Reviewed commit: b88595cb64
ℹ️ 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".
| @@ -206,7 +206,6 @@ def torch_call_with_padding_free_metadata(examples: Sequence[dict]): | |||
| if ids is None: | |||
| continue | |||
| lengths = [len(ids)] | |||
There was a problem hiding this comment.
Preserve derived lengths for the wrapped collator
When a padding-free batch has a seq_lengths key whose value is None (for example, a nullable/custom metadata column), this fallback now only fixes the local packed_seq_lengths tensor and still passes the original None through to original_torch_call. The TRL padding-free collator checks for the key on the first example and then consumes example["seq_lengths"], so these rows can now error or produce inconsistent position_ids even though this wrapper already derived a valid fallback. Please keep caller data immutable by passing a copied/normalized examples list to the underlying collator instead of dropping the derived metadata.
Useful? React with 👍 / 👎.
There was a problem hiding this comment.
That case is already covered. example.get("seq_lengths") returns None both when the key is absent and when its value is None, so a null column takes the same derive branch:
lengths = example.get("seq_lengths") # None, either way
if lengths is None:
ids = example.get("input_ids")
...
lengths = [len(ids)]
collated[index] = {**example, "seq_lengths": lengths}{**example, "seq_lengths": lengths} overwrites the null, so the wrapped collator sees the derived list and never the None.
The one row that still reaches it untouched is a row with no input_ids at all, which hits the continue. That is unchanged from before this PR, where the write-back was also skipped on that path.
|
You are right, and the reason is bigger than a nullable column, so thanks for pushing on it. Fixed in I had framed the write-back as caching, which was wrong. It was load-bearing. TRL's Two things break, not one. Your null case reaches The packed row's second document starts at position 2 instead of 0, so it attends across the document boundary. That is silent, and it is exactly what padding-free metadata exists to prevent. Your suggested shape is what I did: pass a normalized copy to the wrapped collator, built only when something was actually derived, so the common path still hands the caller's own list straight through with no per-row copies. Ran it three ways against |
|
All three red checks are pre-existing, and I checked rather than assumed. |
|
Addressed already, flagging it so the thread is not left looking open. This was raised against
|
|
Confirmed the write-back is still there in unsloth/utils/packing.py and that the follow-up commit keeps the derived lengths reaching the wrapped collator via a copy. The three Core jobs are green on main now, so could you push a refresh and check they pass here before I review? |
30d8d4f to
ed30fc6
Compare
|
Rebased on current main. Both CPU source-harness regressions and pre-commit pass; the expired CI logs are unavailable, so this push requests fresh Core runs. |
…'s examples Put the derived lengths on a shallow copy, so the wrapped collator still sees them without the caller's own rows being mutated.
0c1e631 to
841d32f
Compare
The padding-free collator wrapper adds
seq_lengthsto the caller's examples while deriving metadata. Pass copies containing the derived lengths to the wrapped collator, leaving the original examples untouched.Rebased onto current main. The two CPU source-harness regressions pass for missing, explicit and null lengths, and changed-file pre-commit passes. Full package tests cannot import on this host and remain for upstream CI. The old CI logs have expired, so this refresh requests a fresh run.