fix: 修复 pipeline 模块中的打包逻辑缺陷并完善测试覆盖
This commit is contained in:
+7
-8
@@ -38,21 +38,20 @@ def cache_jsonl(
|
||||
for file_path in files:
|
||||
file_name = Path(file_path).stem
|
||||
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
lines = f.readlines()
|
||||
|
||||
arrows = []
|
||||
for line in tqdm(lines, desc=f"Processing {file_name}", leave=False):
|
||||
arrow = processor.process(json.loads(line))
|
||||
if arrow is not None:
|
||||
arrows.append(arrow)
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
for line in tqdm(f, desc=f"Processing {file_name}", leave=False):
|
||||
arrow = processor.process(json.loads(line))
|
||||
if arrow is not None:
|
||||
arrows.append(arrow)
|
||||
|
||||
package = {key: [a[key] for a in arrows] for key in processor.output_keys}
|
||||
|
||||
output = {}
|
||||
for key in processor.output_keys:
|
||||
if pack_size > 0:
|
||||
output[key] = SequencePacker(pack_size, pad_value).pack(package[key])
|
||||
packer = SequencePacker(pack_size, pad_value) # 每个键独立实例
|
||||
output[key] = packer.pack(package[key])
|
||||
else:
|
||||
output[key] = package[key]
|
||||
|
||||
|
||||
+19
-12
@@ -28,9 +28,9 @@ class IOHandler:
|
||||
return folders
|
||||
|
||||
@staticmethod
|
||||
def save_h5(file_path: str, file_name: str, tensor_group: Dict[str, List[Tensor]]) -> None:
|
||||
os.makedirs(file_path, exist_ok=True)
|
||||
full_path = os.path.join(file_path, f"{file_name}.h5")
|
||||
def save_h5(output_dir: str, file_name: str, tensor_group: Dict[str, List[Tensor]]) -> None:
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
full_path = os.path.join(output_dir, f"{file_name}.h5")
|
||||
with h5py.File(full_path, 'w') as f:
|
||||
for key, tensors in tensor_group.items():
|
||||
grp = f.create_group(key)
|
||||
@@ -38,19 +38,26 @@ class IOHandler:
|
||||
grp.create_dataset(f'data_{idx}', data=tensor.cpu().numpy())
|
||||
|
||||
@staticmethod
|
||||
def load_h5(file_path: str, share_memory: bool = True) -> Dict[str, List[Tensor]]:
|
||||
def load_h5(file_path: str, share_memory=True) -> Dict[str, List[Tensor]]:
|
||||
tensor_group: Dict[str, List[Tensor]] = {}
|
||||
|
||||
root_path = Path(file_path)
|
||||
h5_files = list(root_path.rglob("*.h5")) + list(root_path.rglob("*.hdf5"))
|
||||
|
||||
for h5_file in h5_files:
|
||||
with h5py.File(h5_file, 'r') as f:
|
||||
for key in f.keys():
|
||||
grp = f[key]
|
||||
tensors = [
|
||||
(torch.from_numpy(dset[:]).share_memory_() if share_memory
|
||||
else torch.from_numpy(dset[:]))
|
||||
for dset_name in grp.keys()
|
||||
for dset in [grp[dset_name]]
|
||||
]
|
||||
tensor_group.setdefault(key, []).extend(tensors)
|
||||
return tensor_group
|
||||
dsets = []
|
||||
for dset_name in grp.keys():
|
||||
dset = grp[dset_name]
|
||||
tensor = torch.from_numpy(dset[:])
|
||||
if share_memory:
|
||||
tensor = tensor.share_memory_()
|
||||
dsets.append(tensor)
|
||||
|
||||
if tensor_group.get(key) is None:
|
||||
tensor_group[key] = []
|
||||
tensor_group[key].extend(dsets)
|
||||
|
||||
return tensor_group
|
||||
+68
-16
@@ -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()
|
||||
+24
-2
@@ -61,8 +61,30 @@ class DPOProcessor(BaseProcessor):
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def process(self, input_dict: dict) -> dict:
|
||||
# TODO: 实现 DPO 处理逻辑
|
||||
return None
|
||||
query = input_dict["query"]
|
||||
chosen_response = input_dict["chosen"]
|
||||
rejected_response = input_dict["rejected"]
|
||||
|
||||
q = self.tokenizer.encode(
|
||||
f"<|im_start|>user\n{query}<|im_end|>\n<|im_start|>assistant\n"
|
||||
)
|
||||
|
||||
chosen = self.tokenizer.encode(f"{chosen_response}<|im_end|>\n<eos>")
|
||||
chosen_tokens = torch.tensor(q + chosen, dtype=torch.int32)
|
||||
chosen_mask = torch.zeros_like(chosen_tokens, dtype=torch.bool)
|
||||
chosen_mask[len(q):] = True
|
||||
|
||||
rejected = self.tokenizer.encode(f"{rejected_response}<|im_end|>\n<eos>")
|
||||
rejected_tokens = torch.tensor(q + rejected, dtype=torch.int32)
|
||||
rejected_mask = torch.zeros_like(rejected_tokens, dtype=torch.bool)
|
||||
rejected_mask[len(q):] = True
|
||||
|
||||
return {
|
||||
"chosen": chosen_tokens,
|
||||
"chosen_mask": chosen_mask,
|
||||
"rejected": rejected_tokens,
|
||||
"rejected_mask": rejected_mask,
|
||||
}
|
||||
|
||||
@property
|
||||
def output_keys(self) -> List[str]:
|
||||
|
||||
Reference in New Issue
Block a user