refactor: 重构打包模块,新增 BFD/FFD/Greedy 三种 bin-packing 算法,默认 BFD
- 将 pipeline/packing.py 拆分为 packing/ 子包 (base/stream/binpack) - 新增 BfdPacker(默认)/FfDPacker/GreedyPacker,移除 StreamingPacker - 超长序列直接截断至 pack_size - group_size 语义改为"每 N 个 chunk 合并为一块",默认 1000 - 新增 AutoTokenizer.token_to_id(),修复 ChatML 中 hacky 的 nl_id 获取 - pad_value 默认改为 2(pad_token_id),position_ids pad=0, loss_mask pad=False - 新增 position_ids 打包后归零一致性测试 - scripts/cache_h5.py 新增 --pack-algo 参数
This commit is contained in:
@@ -0,0 +1,84 @@
|
||||
"""Sequence packing algorithms for LLM training data.
|
||||
|
||||
Available packers:
|
||||
- BfdPacker: Best-Fit Decreasing, samples never split (default)
|
||||
- FfDPacker: First-Fit Decreasing, samples never split
|
||||
- GreedyPacker: First-fit in input order, samples never split
|
||||
"""
|
||||
|
||||
from typing import Dict, List, Optional, Union
|
||||
|
||||
import torch
|
||||
|
||||
from pipeline.packing.base import BasePacker
|
||||
from pipeline.packing.binpack import GreedyPacker, FfDPacker, BfdPacker
|
||||
|
||||
|
||||
def pack_tensors(
|
||||
tensors: Dict[str, List[torch.Tensor]],
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool] = 0,
|
||||
dtypes: Optional[Dict[str, torch.dtype]] = None,
|
||||
pad_values: Optional[Dict[str, Union[int, bool]]] = None,
|
||||
algo: Optional[Union[str, BasePacker]] = None,
|
||||
) -> Dict[str, List[torch.Tensor]]:
|
||||
"""Pack multiple named tensor groups in parallel.
|
||||
|
||||
Each group is packed independently with its own packer instance.
|
||||
|
||||
Args:
|
||||
tensors: Dict mapping key names to lists of 1D tensors.
|
||||
pack_size: Fixed chunk length.
|
||||
pad_value: Default padding value, used for keys not in pad_values.
|
||||
dtypes: Optional per-key dtype declarations.
|
||||
pad_values: Optional per-key padding values (e.g. pad_token_id for
|
||||
'sequence', False for 'loss_mask', 0 for 'position_ids').
|
||||
algo: Packing algorithm to use. Can be 'bfd' (default),
|
||||
'ffd', 'greedy', or a BasePacker instance.
|
||||
|
||||
Returns:
|
||||
Dict mapping key names to lists of packed tensors.
|
||||
"""
|
||||
if dtypes is None:
|
||||
dtypes = {}
|
||||
if pad_values is None:
|
||||
pad_values = {}
|
||||
|
||||
output: Dict[str, List[torch.Tensor]] = {}
|
||||
for key, seqs in tensors.items():
|
||||
key_pad = pad_values.get(key, pad_value)
|
||||
actual_packer = _resolve_algo(algo, pack_size, key_pad)
|
||||
dtype = dtypes.get(key)
|
||||
if dtype is not None:
|
||||
actual_packer.dtype = dtype
|
||||
output[key] = actual_packer.pack(seqs)
|
||||
return output
|
||||
|
||||
|
||||
def _resolve_algo(
|
||||
algo: Optional[Union[str, BasePacker]],
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool],
|
||||
) -> BasePacker:
|
||||
if algo is None or algo == "bfd":
|
||||
return BfdPacker(pack_size, pad_value)
|
||||
if isinstance(algo, BasePacker):
|
||||
cls = type(algo)
|
||||
return cls(pack_size, pad_value)
|
||||
if algo == "ffd":
|
||||
return FfDPacker(pack_size, pad_value)
|
||||
if algo == "greedy":
|
||||
return GreedyPacker(pack_size, pad_value)
|
||||
raise ValueError(
|
||||
f"Unknown packing algorithm: {algo}. "
|
||||
f"Choose from: bfd, ffd, greedy"
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"BasePacker",
|
||||
"BfdPacker",
|
||||
"FfDPacker",
|
||||
"GreedyPacker",
|
||||
"pack_tensors",
|
||||
]
|
||||
@@ -0,0 +1,49 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
|
||||
class BasePacker(ABC):
|
||||
"""Abstract base class for sequence packing algorithms.
|
||||
|
||||
All packers must implement pack() and reset().
|
||||
pack() takes a list of 1D tensors and returns a list of packed fixed-size tensors.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool] = 0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
self.pack_size = pack_size
|
||||
self.pad_value = pad_value
|
||||
self.dtype = dtype
|
||||
|
||||
@abstractmethod
|
||||
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
|
||||
"""Pack sequences into fixed-size chunks."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def reset(self) -> None:
|
||||
"""Reset packer state for instance reuse."""
|
||||
...
|
||||
|
||||
def _validate_and_normalize(self, sequences: List[Tensor]) -> List[Tensor]:
|
||||
"""Validate 1D tensors and unify dtype."""
|
||||
if self.dtype is None and sequences:
|
||||
self.dtype = sequences[0].dtype
|
||||
|
||||
normalized: List[Tensor] = []
|
||||
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}"
|
||||
)
|
||||
if seq.dtype != self.dtype:
|
||||
seq = seq.to(self.dtype)
|
||||
normalized.append(seq)
|
||||
return normalized
|
||||
@@ -0,0 +1,174 @@
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from pipeline.packing.base import BasePacker
|
||||
from pipeline.utils import error_handler
|
||||
|
||||
|
||||
def _truncate(tokens: List, max_len: int) -> List:
|
||||
return tokens[:max_len]
|
||||
|
||||
|
||||
def _pad_bin(bin_list: List, target_len: int, pad_value: Union[int, bool], dtype: torch.dtype) -> Tensor:
|
||||
bin_list.extend([pad_value] * (target_len - len(bin_list)))
|
||||
return torch.tensor(bin_list, dtype=dtype)
|
||||
|
||||
|
||||
class GreedyPacker(BasePacker):
|
||||
"""Greedy first-fit packer (no sorting).
|
||||
|
||||
Sequences are packed in input order into the first bin with enough space.
|
||||
Overlong sequences (> pack_size) are truncated to pack_size.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool] = 0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
super().__init__(pack_size, pad_value, dtype)
|
||||
self._bins: List[List] = []
|
||||
|
||||
def reset(self) -> None:
|
||||
self._bins = []
|
||||
|
||||
@error_handler()
|
||||
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
|
||||
if not sequences:
|
||||
return []
|
||||
|
||||
normalized = self._validate_and_normalize(sequences)
|
||||
self._bins = []
|
||||
pack_size = self.pack_size
|
||||
pad_value = self.pad_value
|
||||
|
||||
for seq in normalized:
|
||||
seq_len = int(seq.shape[0])
|
||||
if seq_len > pack_size:
|
||||
self._bins.append(_truncate(seq.tolist(), pack_size))
|
||||
continue
|
||||
placed = False
|
||||
for bin_list in self._bins:
|
||||
if len(bin_list) + seq_len <= pack_size:
|
||||
bin_list.extend(seq.tolist())
|
||||
placed = True
|
||||
break
|
||||
if not placed:
|
||||
self._bins.append(list(seq.tolist()))
|
||||
|
||||
packages: List[Tensor] = []
|
||||
for bin_list in self._bins:
|
||||
packages.append(_pad_bin(bin_list, pack_size, pad_value, self.dtype))
|
||||
|
||||
return packages
|
||||
|
||||
|
||||
class FfDPacker(BasePacker):
|
||||
"""First-Fit Decreasing (FFD) bin-packing packer.
|
||||
|
||||
Sequences are sorted by descending length, then packed into the first
|
||||
bin with enough space. Overlong sequences are truncated to pack_size.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool] = 0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
super().__init__(pack_size, pad_value, dtype)
|
||||
self._bins: List[List] = []
|
||||
|
||||
def reset(self) -> None:
|
||||
self._bins = []
|
||||
|
||||
@error_handler()
|
||||
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
|
||||
if not sequences:
|
||||
return []
|
||||
|
||||
normalized = self._validate_and_normalize(sequences)
|
||||
self._bins = []
|
||||
pack_size = self.pack_size
|
||||
pad_value = self.pad_value
|
||||
|
||||
indexed = [(int(s.shape[0]), s) for s in normalized]
|
||||
indexed.sort(key=lambda x: x[0], reverse=True)
|
||||
|
||||
for seq_len, seq in indexed:
|
||||
if seq_len > pack_size:
|
||||
self._bins.append(_truncate(seq.tolist(), pack_size))
|
||||
continue
|
||||
placed = False
|
||||
for bin_list in self._bins:
|
||||
if len(bin_list) + seq_len <= pack_size:
|
||||
bin_list.extend(seq.tolist())
|
||||
placed = True
|
||||
break
|
||||
if not placed:
|
||||
self._bins.append(list(seq.tolist()))
|
||||
|
||||
packages: List[Tensor] = []
|
||||
for bin_list in self._bins:
|
||||
packages.append(_pad_bin(bin_list, pack_size, pad_value, self.dtype))
|
||||
|
||||
return packages
|
||||
|
||||
|
||||
class BfdPacker(BasePacker):
|
||||
"""Best-Fit Decreasing (BFD) bin-packing packer.
|
||||
|
||||
Sequences are sorted by descending length, then packed into the bin
|
||||
that minimizes remaining space (tightest fit).
|
||||
Overlong sequences are truncated to pack_size.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pack_size: int,
|
||||
pad_value: Union[int, bool] = 0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
super().__init__(pack_size, pad_value, dtype)
|
||||
self._bins: List[List] = []
|
||||
|
||||
def reset(self) -> None:
|
||||
self._bins = []
|
||||
|
||||
@error_handler()
|
||||
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
|
||||
if not sequences:
|
||||
return []
|
||||
|
||||
normalized = self._validate_and_normalize(sequences)
|
||||
self._bins = []
|
||||
pack_size = self.pack_size
|
||||
pad_value = self.pad_value
|
||||
|
||||
indexed = [(int(s.shape[0]), s) for s in normalized]
|
||||
indexed.sort(key=lambda x: x[0], reverse=True)
|
||||
|
||||
for seq_len, seq in indexed:
|
||||
if seq_len > pack_size:
|
||||
self._bins.append(_truncate(seq.tolist(), pack_size))
|
||||
continue
|
||||
best_idx = -1
|
||||
best_remain = pack_size + 1
|
||||
for i, bin_list in enumerate(self._bins):
|
||||
remain = pack_size - len(bin_list)
|
||||
if seq_len <= remain < best_remain:
|
||||
best_remain = remain
|
||||
best_idx = i
|
||||
if best_idx >= 0:
|
||||
self._bins[best_idx].extend(seq.tolist())
|
||||
else:
|
||||
self._bins.append(list(seq.tolist()))
|
||||
|
||||
packages: List[Tensor] = []
|
||||
for bin_list in self._bins:
|
||||
packages.append(_pad_bin(bin_list, pack_size, pad_value, self.dtype))
|
||||
|
||||
return packages
|
||||
Reference in New Issue
Block a user