perf: 优化 processors、cache、packing 模块性能并简化 README

This commit is contained in:
2026-03-30 17:50:26 +08:00
parent 7baa3ea0c3
commit f67bad0d8b
5 changed files with 129 additions and 220 deletions
+21 -14
View File
@@ -2,7 +2,7 @@
import json
import os
import logging
from typing import List
from typing import List, Dict
from pathlib import Path
from tqdm import tqdm
@@ -39,33 +39,40 @@ def cache_jsonl(
"""
os.makedirs(output_dir, exist_ok=True)
output_files: List[str] = []
# Cache output_keys to avoid repeated attribute access
output_keys = processor.output_keys
for file_path in files:
file_name = Path(file_path).stem
arrows = []
# Pre-allocate lists for each output key
arrows: Dict[str, List] = {key: [] for key in output_keys}
# Read and process all lines
with open(file_path, "r", encoding="utf-8") as f:
for line_num, line in enumerate(tqdm(f, desc=f"Processing {file_name}", leave=False), start=1):
try:
arrow = processor.process(json.loads(line))
result = processor.process(json.loads(line))
if result is not None:
# Batch append: add each key's tensor to corresponding list
for key in output_keys:
arrows[key].append(result[key])
except json.JSONDecodeError as e:
logger.warning(f"JSON decode error in {file_path} line {line_num}: {e}. Skipping line.")
continue
except Exception as e:
logger.warning(f"Unexpected error processing line {line_num} in {file_path}: {e}. Skipping line.")
continue
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:
packer = SequencePacker(pack_size, pad_value) # independent instance per key
output[key] = packer.pack(package[key])
else:
output[key] = package[key]
# Convert lists to tensors once per key
if pack_size > 0:
output = {}
for key in output_keys:
packer = SequencePacker(pack_size, pad_value)
output[key] = packer.pack(arrows[key])
else:
# No packing: directly use the arrow tensors
output = arrows
IOHandler.save_h5(output_dir, file_name, output)
h5_path = os.path.join(output_dir, f"{file_name}.h5")
+1
View File
@@ -34,6 +34,7 @@ class IOHandler:
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)
+45 -26
View File
@@ -1,5 +1,5 @@
import logging
from typing import List
from typing import List, Optional
import torch
from torch import Tensor
@@ -14,14 +14,23 @@ class SequencePacker:
self.pack_size = pack_size
self.pad_value = pad_value
self.dtype = dtype
# Pre-allocate buffer for better performance
self._buffer: Optional[Tensor] = None
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
)
# Reuse buffer instead of creating new tensors
if self._buffer is None or self._buffer.shape[0] != self.pack_size:
self._buffer = torch.full(
(self.pack_size,), self.pad_value, dtype=self.dtype
)
else:
self._buffer.fill_(self.pad_value)
self._current_pos = 0
self._packages: List[Tensor] = []
# Backward compatibility: maintain _current_pack reference
self._current_pack = self._buffer
@error_handler()
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
@@ -37,53 +46,63 @@ class SequencePacker:
# Input validation
if not sequences:
return []
# Validate and cache tensor sizes in one pass
tensor_sizes = []
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
tensor_sizes.append(seq.numel())
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 = []
# 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)
# Reset state for new packing
self._packages = []
self._reset()
# Combine sequences with their sizes for sorting
indexed_seqs = list(zip(sequences, tensor_sizes))
# Sort by size descending (First-Fit Decreasing algorithm)
indexed_seqs.sort(key=lambda x: x[1], reverse=True)
for tensor in sorted_sequences:
for tensor, tensor_size in indexed_seqs:
# Truncate sequences that exceed pack_size
if tensor.numel() > self.pack_size:
if tensor_size > self.pack_size:
logger.warning(
f"Sequence length {tensor.numel()} exceeds pack_size {self.pack_size}, truncating"
f"Sequence length {tensor_size} exceeds pack_size {self.pack_size}, truncating"
)
tensor_size = self.pack_size
tensor = tensor[: self.pack_size]
tensor_size = tensor.numel()
# 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
)
# Finish current package (pad to pack_size)
package = self._buffer.clone()
self._packages.append(package)
# Reset buffer for reuse
self._buffer.fill_(self.pad_value)
self._current_pos = 0
# Place tensor in current package (remaining positions stay as pad_value)
self._current_pack[self._current_pos : self._current_pos + tensor_size] = (
tensor
)
# Place tensor in current package
self._buffer[self._current_pos : self._current_pos + tensor_size] = tensor
self._current_pos += tensor_size
# Handle the last package
# Handle the last package (pad to pack_size)
if self._current_pos > 0:
packages.append(self._current_pack)
self._current_pack = None
self._current_pos = 0
package = self._buffer.clone()
self._packages.append(package)
return packages
# Clear buffer and reset state for backward compatibility
self._buffer = None
self._current_pack = None
self._current_pos = 0
return self._packages
def reset(self) -> None:
"""Reset packer state for reuse. More efficient than creating a new instance."""
+16 -7
View File
@@ -41,14 +41,18 @@ class SFTProcessor(BaseProcessor):
self.tokenizer = tokenizer
def process(self, input_dict: Dict[str, Any]) -> Dict[str, Tensor]:
query, response = input_dict["query"], input_dict["response"]
query = input_dict["query"]
response = input_dict["response"]
q = self.tokenizer.encode(
f"<|im_start|>user\n{query}<|im_end|>\n<|im_start|>assistant\n"
)
a = self.tokenizer.encode(f"{response}<|im_end|>\n<eos>")
q_len = len(q)
tokens = torch.tensor(q + a, dtype=torch.int32)
loss_mask = torch.zeros_like(tokens, dtype=torch.bool)
loss_mask[len(q):] = True
loss_mask = torch.zeros(q_len + len(a), dtype=torch.bool)
loss_mask[q_len:] = True
return {"sequence": tokens, "loss_mask": loss_mask}
@property
@@ -72,14 +76,19 @@ class DPOProcessor(BaseProcessor):
)
chosen = self.tokenizer.encode(f"{chosen_response}<|im_end|>\n<eos>")
q_len = len(q)
chosen_len = len(chosen)
chosen_tokens = torch.tensor(q + chosen, dtype=torch.int32)
chosen_mask = torch.zeros_like(chosen_tokens, dtype=torch.bool)
chosen_mask[len(q):] = True
chosen_mask = torch.zeros(q_len + chosen_len, dtype=torch.bool)
chosen_mask[q_len:] = True
rejected = self.tokenizer.encode(f"{rejected_response}<|im_end|>\n<eos>")
rejected_len = len(rejected)
rejected_tokens = torch.tensor(q + rejected, dtype=torch.int32)
rejected_mask = torch.zeros_like(rejected_tokens, dtype=torch.bool)
rejected_mask[len(q):] = True
rejected_mask = torch.zeros(q_len + rejected_len, dtype=torch.bool)
rejected_mask[q_len:] = True
return {
"chosen": chosen_tokens,