fix: 修复打包策略问题
This commit is contained in:
+51
-19
@@ -15,16 +15,13 @@ class TestSequencePacker:
|
||||
torch.tensor([6, 7, 8, 9], dtype=torch.int32),
|
||||
]
|
||||
packages = packer.pack(sequences)
|
||||
assert len(packages) >= 1
|
||||
assert len(packages) == 1
|
||||
for pkg in packages:
|
||||
assert pkg.shape == (10,)
|
||||
|
||||
# Verify all original values are present
|
||||
all_values = []
|
||||
for pkg in packages:
|
||||
all_values.extend(pkg[pkg != 0].tolist())
|
||||
for val in [1, 2, 3, 4, 5, 6, 7, 8, 9]:
|
||||
assert val in all_values
|
||||
# Verify all original values are present in order
|
||||
assert packages[0][:9].tolist() == [1, 2, 3, 4, 5, 6, 7, 8, 9]
|
||||
assert packages[0][9] == 0 # padding
|
||||
|
||||
def test_empty_list_input(self):
|
||||
packer = SequencePacker(pack_size=10)
|
||||
@@ -37,12 +34,13 @@ class TestSequencePacker:
|
||||
assert packages[0][:3].tolist() == [1, 2, 3]
|
||||
assert packages[0][3:].tolist() == [-1] * 7
|
||||
|
||||
def test_truncate_long_sequence(self, caplog):
|
||||
def test_long_sequence_split_across_chunks(self):
|
||||
"""Sequences longer than pack_size are split across multiple chunks."""
|
||||
packer = SequencePacker(pack_size=5, pad_value=0)
|
||||
packages = packer.pack([torch.tensor([1, 2, 3, 4, 5, 6, 7, 8], dtype=torch.int32)])
|
||||
assert len(packages) == 1
|
||||
assert len(packages) == 2
|
||||
assert packages[0].tolist() == [1, 2, 3, 4, 5]
|
||||
assert "truncating" in caplog.text.lower() or "exceeds" in caplog.text.lower()
|
||||
assert packages[1].tolist() == [6, 7, 8, 0, 0]
|
||||
|
||||
def test_padding_value(self):
|
||||
packer = SequencePacker(pack_size=8, pad_value=99)
|
||||
@@ -104,24 +102,58 @@ class TestSequencePacker:
|
||||
assert packages[1].tolist() == [11] + [-1] * 9
|
||||
|
||||
def test_cross_group_ordering(self):
|
||||
"""Tensor groups with identical per-item lengths are sorted identically."""
|
||||
packer = SequencePacker(pack_size=10, pad_value=0)
|
||||
# sequences: lengths [3, 1, 4] -> after sort desc: [4, 3, 1]
|
||||
"""Separate packers for different dtypes produce identical chunk boundaries."""
|
||||
seq_packer = SequencePacker(pack_size=10, pad_value=0, dtype=torch.int32)
|
||||
mask_packer = SequencePacker(pack_size=10, pad_value=False, dtype=torch.bool)
|
||||
# sequences: lengths [3, 1, 4]
|
||||
seqs = [
|
||||
torch.tensor([1, 2, 3], dtype=torch.int32),
|
||||
torch.tensor([10], dtype=torch.int32),
|
||||
torch.tensor([4, 5, 6, 7], dtype=torch.int32),
|
||||
]
|
||||
masks = [
|
||||
torch.tensor([True, True, True], dtype=torch.bool),
|
||||
torch.tensor([True], dtype=torch.bool),
|
||||
torch.tensor([True, True, True, True], dtype=torch.bool),
|
||||
torch.tensor([False, False, True], dtype=torch.bool),
|
||||
torch.tensor([False], dtype=torch.bool),
|
||||
torch.tensor([False, False, False, True], dtype=torch.bool),
|
||||
]
|
||||
packed_seqs = packer.pack(seqs)
|
||||
packer.reset()
|
||||
packed_masks = packer.pack(masks)
|
||||
packed_seqs = seq_packer.pack(seqs)
|
||||
packed_masks = mask_packer.pack(masks)
|
||||
|
||||
# Verify mask packer uses bool dtype
|
||||
assert packed_masks[0].dtype == torch.bool
|
||||
# Both groups should produce the same number of packages
|
||||
assert len(packed_seqs) == len(packed_masks)
|
||||
|
||||
def test_stream_split_across_chunks(self):
|
||||
"""Sequences are split across chunks in streaming mode."""
|
||||
packer = SequencePacker(pack_size=5, pad_value=0)
|
||||
packages = packer.pack([
|
||||
torch.tensor([1, 2, 3], dtype=torch.int32),
|
||||
torch.tensor([4, 5, 6, 7, 8], dtype=torch.int32),
|
||||
])
|
||||
assert len(packages) == 2
|
||||
# First chunk: [1, 2, 3, 4, 5] — first seq + part of second
|
||||
assert packages[0].tolist() == [1, 2, 3, 4, 5]
|
||||
# Second chunk: [6, 7, 8, 0, 0] — rest of second + padding
|
||||
assert packages[1].tolist() == [6, 7, 8, 0, 0]
|
||||
|
||||
def test_reset_method(self):
|
||||
packer = SequencePacker(pack_size=10, pad_value=0)
|
||||
seqs = [torch.tensor([1, 2, 3], dtype=torch.int32)]
|
||||
packer.pack(seqs)
|
||||
assert len(packer._packages) == 1
|
||||
packer.reset()
|
||||
assert len(packer._packages) == 0
|
||||
assert packer._pos == 0
|
||||
assert packer._buffer == []
|
||||
|
||||
def test_no_sorting_needed(self):
|
||||
"""Streaming concat preserves input order, no sorting."""
|
||||
packer = SequencePacker(pack_size=4, pad_value=-1)
|
||||
# short then long (fits in 2 chunks)
|
||||
packages = packer.pack([
|
||||
torch.tensor([1], dtype=torch.int32),
|
||||
torch.tensor([2, 3, 4, 5, 6, 7], dtype=torch.int32),
|
||||
])
|
||||
assert packages[0].tolist() == [1, 2, 3, 4]
|
||||
assert packages[1].tolist() == [5, 6, 7, -1]
|
||||
|
||||
Reference in New Issue
Block a user