fix: 修复 pipeline 模块中的打包逻辑缺陷并完善测试覆盖

This commit is contained in:
2026-03-30 12:22:37 +08:00
parent 07ef471aa6
commit 71887bb4bb
11 changed files with 1186 additions and 39 deletions
+68 -16
View File
@@ -1,35 +1,87 @@
import logging
from typing import List
import torch
from torch import Tensor
logger = logging.getLogger(__name__)
class SequencePacker:
"""序列打包(bin-packing"""
def __init__(self, pack_size: int, pad_value: int = 0):
def __init__(self, pack_size: int, pad_value: int = 0, dtype=torch.int32):
self.pack_size = pack_size
self.pad_value = pad_value
self.dtype = dtype
self._reset()
def _reset(self) -> None:
"""Reset internal state for instance reuse."""
self._current_pack = torch.full(
(self.pack_size,), self.pad_value, dtype=self.dtype
)
self._current_pos = 0
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
"""
Pack sequences into fixed-size packages.
Args:
sequences: List of input tensors
Returns:
List of packed tensors, each with length equal to pack_size
"""
# Input validation
if not sequences:
return []
for i, seq in enumerate(sequences):
if seq.dim() != 1:
raise ValueError(
f"Expected 1D tensor at index {i}, got {seq.dim()}D tensor with shape {seq.shape}"
)
# Check dtype compatibility and warn if mismatched
if seq.dtype != self.dtype:
logger.warning(
f"Input tensor dtype {seq.dtype} does not match packer dtype {self.dtype}, "
f"will be converted. This may affect packing efficiency."
)
packages = []
sequences.sort(key=lambda x: x.numel(), reverse=True)
# Sort by length in descending order to improve packing efficiency
# Use sorted() to avoid modifying the input list
sorted_sequences = sorted(sequences, key=lambda x: x.numel(), reverse=True)
current_pack = torch.full((self.pack_size,), self.pad_value, dtype=torch.int32)
current_pos = 0
for tensor in sequences:
tensor = tensor[:self.pack_size] if tensor.numel() > self.pack_size else tensor
for tensor in sorted_sequences:
# Truncate sequences that exceed pack_size
if tensor.numel() > self.pack_size:
logger.warning(
f"Sequence length {tensor.numel()} exceeds pack_size {self.pack_size}, truncating"
)
tensor = tensor[: self.pack_size]
tensor_size = tensor.numel()
if current_pos + tensor_size > self.pack_size:
packages.append(current_pack)
current_pack = torch.full((self.pack_size,), self.pad_value, dtype=torch.int32)
current_pos = 0
# Current package is full, create a new one
if self._current_pos + tensor_size > self.pack_size:
packages.append(self._current_pack)
self._current_pack = torch.full(
(self.pack_size,), self.pad_value, dtype=self.dtype
)
self._current_pos = 0
current_pack[current_pos:current_pos + tensor_size] = tensor
current_pos += tensor_size
# Place tensor in current package (remaining positions stay as pad_value)
self._current_pack[self._current_pos : self._current_pos + tensor_size] = (
tensor
)
self._current_pos += tensor_size
if current_pos > 0:
packages.append(current_pack)
# Handle the last package
if self._current_pos > 0:
packages.append(self._current_pack)
self._current_pack = None
self._current_pos = 0
return packages
def reset(self) -> None:
"""Reset packer state for reuse. More efficient than creating a new instance."""
self._reset()