Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
77 changes: 77 additions & 0 deletions tests/utils/test_packing.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
patch_hybrid_linear_attention_varlen,
)

import copy
import inspect
import logging
from contextlib import ExitStack
Expand Down Expand Up @@ -160,6 +161,80 @@ def test_mask_packed_sequence_boundaries_across_multiple_rows():
assert torch.any(flat != -100)


def test_enable_padding_free_metadata_does_not_mutate_examples():
# The wrapper derives seq_lengths from input_ids when the dataset does not
# carry them. It used to cache that back onto the example, so a collator
# wrapper wrote a new column into the caller's rows; the sibling
# enable_sample_packing collects the same lengths without writing back.
collator = _PaddingFreeCollator()
trainer = SimpleNamespace(
data_collator = collator,
args = SimpleNamespace(remove_unused_columns = True),
)
enable_padding_free_metadata(_DummyModel(), trainer)

examples = [
{"input_ids": [1, 2, 3], "labels": [1, 2, 3]},
{"input_ids": [4, 5], "labels": [4, 5]},
]
before = copy.deepcopy(examples)

batch = trainer.data_collator.torch_call(examples)

assert examples == before, "collator wrapper mutated the caller's examples"
assert torch.equal(batch["packed_seq_lengths"], torch.tensor([3, 2], dtype = torch.int32))

# explicit seq_lengths are still honoured, and still not written back
explicit = [{"input_ids": [1, 2, 3], "seq_lengths": [2, 1]}]
explicit_before = copy.deepcopy(explicit)
batch = trainer.data_collator.torch_call(explicit)
assert explicit == explicit_before
assert torch.equal(batch["packed_seq_lengths"], torch.tensor([2, 1], dtype = torch.int32))


def test_enable_padding_free_metadata_still_hands_derived_lengths_to_the_collator():
# Not writing to the caller's row must not mean the wrapped collator stops
# seeing the lengths. TRL's padding-free collator decides whether to use
# seq_lengths at all from examples[0], then reads it off every row, so a row
# that reaches it without the key (or with a null one) changes the
# position_ids it builds, or makes it sum(None).
collator = _PaddingFreeCollator()
trainer = SimpleNamespace(
data_collator = collator,
args = SimpleNamespace(remove_unused_columns = True),
)
enable_padding_free_metadata(_DummyModel(), trainer)

examples = [
{"input_ids": [1, 2, 3], "labels": [1, 2, 3]},
{"input_ids": [4, 5], "labels": [4, 5]},
]
before = copy.deepcopy(examples)

trainer.data_collator.torch_call(examples)

assert examples == before, "collator wrapper mutated the caller's examples"
assert [row["seq_lengths"] for row in collator.seen] == [[3], [2]]
assert [row["labels"] for row in collator.seen] == [[1, 2, 3], [4, 5]]

# a row carrying the column with no value counts as missing here, so it has
# to be replaced rather than passed through
nulled = [{"input_ids": [1, 2], "seq_lengths": None}]
nulled_before = copy.deepcopy(nulled)

trainer.data_collator.torch_call(nulled)

assert nulled == nulled_before
assert [row["seq_lengths"] for row in collator.seen] == [[2]]

# nothing was derived, so the caller's own list goes straight through
explicit = [{"input_ids": [1, 2, 3], "seq_lengths": [2, 1]}]

trainer.data_collator.torch_call(explicit)

assert collator.seen is explicit


def test_configure_sample_packing():
config = SimpleNamespace()
configure_sample_packing(config)
Expand Down Expand Up @@ -978,9 +1053,11 @@ def __init__(self):
self.padding_free = True
self.return_position_ids = False
self.calls = 0
self.seen = None

def torch_call(self, examples):
self.calls += 1
self.seen = examples
return {
"input_ids": torch.tensor([[0]], dtype = torch.long),
"examples_seen": self.calls,
Expand Down
13 changes: 10 additions & 3 deletions unsloth/utils/packing.py
Original file line number Diff line number Diff line change
Expand Up @@ -252,18 +252,25 @@ def enable_padding_free_metadata(model, trainer):

def torch_call_with_padding_free_metadata(examples: Sequence[dict]):
seq_lengths: list[int] = []
collated = examples
if examples and isinstance(examples[0], dict):
for example in examples:
for index, example in enumerate(examples):
lengths = example.get("seq_lengths")
if lengths is None:
ids = example.get("input_ids")
if ids is None:
continue
lengths = [len(ids)]

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P2 Badge 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 👍 / 👎.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

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.

example["seq_lengths"] = lengths
# The wrapped collator decides whether to use seq_lengths at
# all from the first example and then reads it off every row,
# so the derived lengths still have to reach it. Put them on a
# shallow copy instead of the caller's own row.
if collated is examples:
collated = list(examples)
collated[index] = {**example, "seq_lengths": lengths}
seq_lengths.extend(lengths)

batch = original_torch_call(examples)
batch = original_torch_call(collated)
if seq_lengths:
# Labels left alone for the same reason as enable_sample_packing: num_items_in_batch is counted off
# this batch and the zoo's discount of the boundary targets is idempotent.
Expand Down
Loading