feat(utils): 添加序列打包功能并优化数据处理流程

This commit is contained in:
2025-07-21 19:17:04 +08:00
parent 6d52bae1ca
commit 70485d7c01
+16 -7
View File
@@ -1,10 +1,8 @@
from typing import Dict, List, Callable, Union from typing import Dict, List, Callable, Union
from datasets import DatasetDict from datasets import DatasetDict
from tokenizer import BpeTokenizer
from tqdm import tqdm from tqdm import tqdm
from torch import Tensor from torch import Tensor
import torch.nn.functional as F
import pickle as pkl import pickle as pkl
import torch import torch
import json import json
@@ -27,6 +25,11 @@ def comprehensive_normalization(text):
pattern = re.compile('|'.join(re.escape(k) for k in replacements)) pattern = re.compile('|'.join(re.escape(k) for k in replacements))
return pattern.sub(lambda m: replacements[m.group()], text) return pattern.sub(lambda m: replacements[m.group()], text)
def pt_processor(text):
text = comprehensive_normalization(text)
text = text.lower()
return text
def pack_sequences(sequences: List[Tensor], pack_size: int, pad_value: int) -> List[Tensor]: def pack_sequences(sequences: List[Tensor], pack_size: int, pad_value: int) -> List[Tensor]:
packages = [] packages = []
@@ -71,16 +74,18 @@ def dump_pkl_files(
base_out_dir: str, base_out_dir: str,
process_func: Callable[[dict], dict], process_func: Callable[[dict], dict],
output_keys: List[str], output_keys: List[str],
packing_size: int = -1 packing_size: int = -1,
pad_value: int = 0
): ):
for file_path in files: for file_path in files:
out_file_name = os.path.basename(file_path).replace(".jsonl", ".pkl") out_file_name = os.path.basename(file_path).replace(".jsonl", ".pkl")
out_file_path = os.path.join(base_out_dir, out_file_name) out_file_path = os.path.join(base_out_dir, out_file_name)
file_name = os.path.basename(file_path) file_name = os.path.basename(file_path)
arrows: Dict[str, List[Tensor]] = {}
os.makedirs(os.path.dirname(out_file_path), exist_ok=True) os.makedirs(os.path.dirname(out_file_path), exist_ok=True)
arrows: Dict[str, List[Tensor]] = {}
with open(file_path, "r") as f: with open(file_path, "r") as f:
lines = f.readlines() lines = f.readlines()
@@ -89,10 +94,14 @@ def dump_pkl_files(
for key in output_keys: for key in output_keys:
arrows[key].extend(arrow[key]) arrows[key].extend(arrow[key])
output_package = {} output_package: Dict[str, Tensor] = {}
for key in output_keys: for key in output_keys:
tensor = torch.cat(arrows[key]) print(f"Packaging key: '{key}'")
output_package[key] = tensor if packing_size > 0:
arrows[key] = pack_sequences(arrows[key], packing_size, pad_value)
sequence = torch.cat(arrows[key])
output_package[key] = sequence
with open(out_file_path, "w") as f: with open(out_file_path, "w") as f:
pkl.dump(output_package, f) pkl.dump(output_package, f)