refactor(dump_pt_file): 重构数据处理流程并优化代码结构

This commit is contained in:
2025-07-21 20:27:07 +08:00
parent f9ce67e03c
commit 7f45c701b4
2 changed files with 14 additions and 10 deletions
+14 -4
View File
@@ -1,9 +1,18 @@
from utils import dump_pkl_files, fetch_files from utils import dump_pkl_files, fetch_files
from tokenizer import BpeTokenizer from tokenizer import BpeTokenizer
import torch
import os import os
def processor(intput_str: str): def get_processor(tokenizer: BpeTokenizer):
return f"{intput_str} <eos>" def processor(intput_dict: dict) -> dict:
segment = intput_dict["text"]
ids = tokenizer.encode(f"{segment} <eos>")
t_ids = torch.tensor(ids, dtype=torch.int32)
return {'sequence': t_ids}
return processor
if __name__ == "__main__": if __name__ == "__main__":
tokenizer = BpeTokenizer("tokenizer.json") tokenizer = BpeTokenizer("tokenizer.json")
@@ -13,9 +22,10 @@ if __name__ == "__main__":
os.path.join("dataset", "english-wiki"), os.path.join("dataset", "english-wiki"),
os.path.join("dataset", "chinese-wiki"), os.path.join("dataset", "chinese-wiki"),
] ]
base_out_dir = "pkl_output" base_out_dir = "pkl_output"
files = [] files = []
for dir_path in base_dir: for dir_path in base_dir:
files.extend(fetch_files(dir_path)) files.extend(fetch_files(dir_path))
processor = get_processor(tokenizer)
dump_pkl_files(tokenizer, files, base_out_dir, processor) dump_pkl_files(files, base_out_dir, processor, ["text"])
-6
View File
@@ -25,12 +25,6 @@ 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 = []
sequences.sort(key=lambda x: x.numel(), reverse=True) sequences.sort(key=lambda x: x.numel(), reverse=True)