Skip to content
Closed
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
65 changes: 65 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,68 @@ def test_mask_packed_sequence_boundaries_across_multiple_rows():
assert torch.any(flat != -100)


def test_enable_padding_free_metadata_does_not_mutate_examples():
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 = [{"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():
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]]

# seq_lengths=None counts as missing: TRL would sum(None)
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]]

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 +1041,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
10 changes: 7 additions & 3 deletions unsloth/utils/packing.py
Original file line number Diff line number Diff line change
Expand Up @@ -255,18 +255,22 @@ 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)]
example["seq_lengths"] = lengths
# TRL's collator keys seq_lengths off examples[0] and reads every row: pass a copy, not the caller's 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