fix: 修复打包策略问题
This commit is contained in:
+44
-74
@@ -9,64 +9,55 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
class SequencePacker:
|
class SequencePacker:
|
||||||
"""
|
"""
|
||||||
Packs variable-length sequences into fixed-size tensors, suitable for
|
Stream-concatenation packer for LLM training sequences.
|
||||||
concatenating unequal-length training samples into uniform shapes
|
|
||||||
for DataLoader / model training.
|
|
||||||
|
|
||||||
Algorithm (Sorted Greedy Fill, based on First-Fit Decreasing heuristic):
|
Algorithm (streaming concat):
|
||||||
|
|
||||||
Input: sequences = [A(len=5), B(len=2), C(len=3)], pack_size = 8
|
Input: sequences = [A(len=3), B(len=5), C(len=2)], pack_size = 6
|
||||||
|
|
||||||
1. Validate & Normalize
|
1. Validate & Normalize
|
||||||
- Check 1D dimension, unify dtype, truncate overlong sequences with warning
|
- Check 1D dimension, unify dtype, warn on overlong sequences
|
||||||
- Result: [(A,5), (B,2), (C,3)]
|
- Result: [A, B, C]
|
||||||
|
|
||||||
2. Sort by length descending (FFD)
|
2. Stream into buffer, slice off full chunks
|
||||||
- Result: [(A,5), (C,3), (B,2)]
|
- buffer += A(3) -> [a1 a2 a3], pos=3
|
||||||
|
- buffer += B(5) -> [a1 a2 a3 b1 b2 b3 b4 b5], pos=8
|
||||||
|
pos >= 6 -> flush [a1 a2 a3 b1 b2 b3], buffer=[b4 b5], pos=2
|
||||||
|
- buffer += C(2) -> [b4 b5 c1 c2], pos=4
|
||||||
|
loop ends -> flush tail [b4 b5 c1 c2 PAD PAD]
|
||||||
|
|
||||||
3. Greedy fill: write into a pre-allocated buffer sequentially, flush when full
|
Output: [[a1 a2 a3 b1 b2 b3], [b4 b5 c1 c2 PAD PAD]]
|
||||||
- Write A(5) -> buffer = [A A A A A _ _ _], pos=5
|
|
||||||
- Write C(3) -> pos+3=8 <= 8 -> buffer = [A A A A A C C C], pos=8
|
|
||||||
- Buffer full -> flush as package[0], reset buffer & pos=0
|
|
||||||
- Write B(2) -> buffer = [B B _ _ _ _ _ _], pos=2
|
|
||||||
- Loop ends -> flush tail -> package[1] = [B B 0 0 0 0 0 0]
|
|
||||||
|
|
||||||
Output: [package[0], package[1]]
|
Samples may be split across chunks — this is intentional and standard
|
||||||
|
practice in LLM training (TRL, Megatron-LM, etc.).
|
||||||
|
|
||||||
Cross-group consistency:
|
Cross-group consistency:
|
||||||
When packing different key groups (e.g. sequences and loss_masks)
|
Different tensor groups (e.g. input_ids, loss_masks) packed with
|
||||||
with separate pack() calls, tensors at the same index always have
|
separate packer instances on samples with matching lengths produce
|
||||||
identical lengths, so the descending sort produces the exact same
|
identical chunk boundaries. Element-level correspondence is preserved.
|
||||||
ordering. Element-level correspondence across groups is preserved.
|
|
||||||
|
|
||||||
Performance:
|
|
||||||
- Pre-allocated buffer reused via fill_() to avoid repeated tensor creation
|
|
||||||
- Attributes cached as local variables inside the loop to reduce lookup overhead
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, pack_size: int, pad_value: int = 0, dtype: torch.dtype = None):
|
def __init__(self, pack_size: int, pad_value: int = 0, dtype: torch.dtype = torch.int32):
|
||||||
self.pack_size = pack_size
|
self.pack_size = pack_size
|
||||||
self.pad_value = pad_value
|
self.pad_value = pad_value
|
||||||
self.dtype = dtype # None = follow input dtype
|
self.dtype = dtype
|
||||||
self._buffer: Tensor | None = None
|
self._buffer: List[int] = []
|
||||||
self._pos = 0
|
self._pos: int = 0
|
||||||
self._packages: List[Tensor] = []
|
self._packages: List[Tensor] = []
|
||||||
|
|
||||||
def reset(self) -> None:
|
def reset(self) -> None:
|
||||||
"""Reset packer state for instance reuse, unlocking dtype."""
|
"""Reset packer state for instance reuse."""
|
||||||
self.dtype = None
|
self._buffer = []
|
||||||
self._buffer = None
|
|
||||||
self._pos = 0
|
self._pos = 0
|
||||||
self._packages = []
|
self._packages = []
|
||||||
|
|
||||||
@error_handler()
|
@error_handler()
|
||||||
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
|
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
|
||||||
"""
|
"""
|
||||||
Pack sequences into fixed-size packages using First-Fit Decreasing.
|
Pack sequences via streaming concatenation into fixed-size chunks.
|
||||||
|
|
||||||
Sequences are sorted by length descending to minimize wasted padding.
|
Sequences are concatenated in order and sliced at pack_size boundaries.
|
||||||
All tensor groups (e.g. sequences, loss_masks) with matching per-item
|
The final chunk is padded with pad_value.
|
||||||
lengths produce identical ordering, so cross-group correspondence is preserved.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
sequences: List of 1D input tensors.
|
sequences: List of 1D input tensors.
|
||||||
@@ -77,54 +68,33 @@ class SequencePacker:
|
|||||||
if not sequences:
|
if not sequences:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
# --- validate & normalize in a single pass ---
|
# --- validate & normalize ---
|
||||||
normalized: list[tuple[Tensor, int]] = []
|
normalized: List[Tensor] = []
|
||||||
target_dtype = self.dtype if self.dtype is not None else sequences[0].dtype
|
|
||||||
for i, seq in enumerate(sequences):
|
for i, seq in enumerate(sequences):
|
||||||
if seq.dim() != 1:
|
if seq.dim() != 1:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Expected 1D tensor at index {i}, got {seq.dim()}D tensor with shape {seq.shape}"
|
f"Expected 1D tensor at index {i}, got {seq.dim()}D tensor with shape {seq.shape}"
|
||||||
)
|
)
|
||||||
if seq.dtype != target_dtype:
|
if seq.dtype != self.dtype:
|
||||||
seq = seq.to(target_dtype)
|
seq = seq.to(self.dtype)
|
||||||
length = seq.numel()
|
normalized.append(seq)
|
||||||
if length > self.pack_size:
|
|
||||||
seq = seq[: self.pack_size]
|
|
||||||
length = self.pack_size
|
|
||||||
normalized.append((seq, length))
|
|
||||||
|
|
||||||
# --- reset internal state ---
|
# --- stream into buffer, slice off full chunks ---
|
||||||
buf = self._buffer
|
self._buffer = []
|
||||||
if buf is None or buf.dtype != target_dtype:
|
|
||||||
buf = torch.full((self.pack_size,), self.pad_value, dtype=target_dtype)
|
|
||||||
self._buffer = buf
|
|
||||||
buf.fill_(self.pad_value)
|
|
||||||
self._pos = 0
|
|
||||||
self._packages = []
|
self._packages = []
|
||||||
|
|
||||||
# --- sort by length descending (FFD heuristic) ---
|
|
||||||
normalized.sort(key=lambda x: x[1], reverse=True)
|
|
||||||
|
|
||||||
# --- greedy fill ---
|
|
||||||
buf = self._buffer
|
|
||||||
pos = self._pos
|
|
||||||
packages = self._packages
|
|
||||||
pack_size = self.pack_size
|
pack_size = self.pack_size
|
||||||
pad_value = self.pad_value
|
buf = self._buffer
|
||||||
|
|
||||||
for tensor, length in normalized:
|
for seq in normalized:
|
||||||
if pos + length > pack_size:
|
buf.extend(seq.tolist())
|
||||||
# flush current package
|
while len(buf) >= pack_size:
|
||||||
packages.append(buf.clone())
|
self._packages.append(torch.tensor(buf[:pack_size], dtype=self.dtype))
|
||||||
buf.fill_(pad_value)
|
buf = buf[pack_size:]
|
||||||
pos = 0
|
|
||||||
buf[pos : pos + length] = tensor
|
|
||||||
pos += length
|
|
||||||
|
|
||||||
# flush the last (possibly partial) package
|
# flush tail with padding
|
||||||
if pos > 0:
|
if buf:
|
||||||
packages.append(buf.clone())
|
padded = buf + [self.pad_value] * (pack_size - len(buf))
|
||||||
|
self._packages.append(torch.tensor(padded, dtype=self.dtype))
|
||||||
|
|
||||||
# write back state
|
self._pos = len(buf)
|
||||||
self._pos = pos
|
|
||||||
return self._packages
|
return self._packages
|
||||||
|
|||||||
+51
-19
@@ -15,16 +15,13 @@ class TestSequencePacker:
|
|||||||
torch.tensor([6, 7, 8, 9], dtype=torch.int32),
|
torch.tensor([6, 7, 8, 9], dtype=torch.int32),
|
||||||
]
|
]
|
||||||
packages = packer.pack(sequences)
|
packages = packer.pack(sequences)
|
||||||
assert len(packages) >= 1
|
assert len(packages) == 1
|
||||||
for pkg in packages:
|
for pkg in packages:
|
||||||
assert pkg.shape == (10,)
|
assert pkg.shape == (10,)
|
||||||
|
|
||||||
# Verify all original values are present
|
# Verify all original values are present in order
|
||||||
all_values = []
|
assert packages[0][:9].tolist() == [1, 2, 3, 4, 5, 6, 7, 8, 9]
|
||||||
for pkg in packages:
|
assert packages[0][9] == 0 # padding
|
||||||
all_values.extend(pkg[pkg != 0].tolist())
|
|
||||||
for val in [1, 2, 3, 4, 5, 6, 7, 8, 9]:
|
|
||||||
assert val in all_values
|
|
||||||
|
|
||||||
def test_empty_list_input(self):
|
def test_empty_list_input(self):
|
||||||
packer = SequencePacker(pack_size=10)
|
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, 2, 3]
|
||||||
assert packages[0][3:].tolist() == [-1] * 7
|
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)
|
packer = SequencePacker(pack_size=5, pad_value=0)
|
||||||
packages = packer.pack([torch.tensor([1, 2, 3, 4, 5, 6, 7, 8], dtype=torch.int32)])
|
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 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):
|
def test_padding_value(self):
|
||||||
packer = SequencePacker(pack_size=8, pad_value=99)
|
packer = SequencePacker(pack_size=8, pad_value=99)
|
||||||
@@ -104,24 +102,58 @@ class TestSequencePacker:
|
|||||||
assert packages[1].tolist() == [11] + [-1] * 9
|
assert packages[1].tolist() == [11] + [-1] * 9
|
||||||
|
|
||||||
def test_cross_group_ordering(self):
|
def test_cross_group_ordering(self):
|
||||||
"""Tensor groups with identical per-item lengths are sorted identically."""
|
"""Separate packers for different dtypes produce identical chunk boundaries."""
|
||||||
packer = SequencePacker(pack_size=10, pad_value=0)
|
seq_packer = SequencePacker(pack_size=10, pad_value=0, dtype=torch.int32)
|
||||||
# sequences: lengths [3, 1, 4] -> after sort desc: [4, 3, 1]
|
mask_packer = SequencePacker(pack_size=10, pad_value=False, dtype=torch.bool)
|
||||||
|
# sequences: lengths [3, 1, 4]
|
||||||
seqs = [
|
seqs = [
|
||||||
torch.tensor([1, 2, 3], dtype=torch.int32),
|
torch.tensor([1, 2, 3], dtype=torch.int32),
|
||||||
torch.tensor([10], dtype=torch.int32),
|
torch.tensor([10], dtype=torch.int32),
|
||||||
torch.tensor([4, 5, 6, 7], dtype=torch.int32),
|
torch.tensor([4, 5, 6, 7], dtype=torch.int32),
|
||||||
]
|
]
|
||||||
masks = [
|
masks = [
|
||||||
torch.tensor([True, True, True], dtype=torch.bool),
|
torch.tensor([False, False, True], dtype=torch.bool),
|
||||||
torch.tensor([True], dtype=torch.bool),
|
torch.tensor([False], dtype=torch.bool),
|
||||||
torch.tensor([True, True, True, True], dtype=torch.bool),
|
torch.tensor([False, False, False, True], dtype=torch.bool),
|
||||||
]
|
]
|
||||||
packed_seqs = packer.pack(seqs)
|
packed_seqs = seq_packer.pack(seqs)
|
||||||
packer.reset()
|
packed_masks = mask_packer.pack(masks)
|
||||||
packed_masks = packer.pack(masks)
|
|
||||||
|
|
||||||
# Verify mask packer uses bool dtype
|
# Verify mask packer uses bool dtype
|
||||||
assert packed_masks[0].dtype == torch.bool
|
assert packed_masks[0].dtype == torch.bool
|
||||||
# Both groups should produce the same number of packages
|
# Both groups should produce the same number of packages
|
||||||
assert len(packed_seqs) == len(packed_masks)
|
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