From 7f45c701b4b017c647abe66ce2e5a351c45151ea Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Mon, 21 Jul 2025 20:27:07 +0800 Subject: [PATCH] =?UTF-8?q?refactor(dump=5Fpt=5Ffile):=20=E9=87=8D?= =?UTF-8?q?=E6=9E=84=E6=95=B0=E6=8D=AE=E5=A4=84=E7=90=86=E6=B5=81=E7=A8=8B?= =?UTF-8?q?=E5=B9=B6=E4=BC=98=E5=8C=96=E4=BB=A3=E7=A0=81=E7=BB=93=E6=9E=84?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- dump_pt_file.py | 18 ++++++++++++++---- utils.py | 6 ------ 2 files changed, 14 insertions(+), 10 deletions(-) diff --git a/dump_pt_file.py b/dump_pt_file.py index 1822141..c6e17d5 100644 --- a/dump_pt_file.py +++ b/dump_pt_file.py @@ -1,9 +1,18 @@ from utils import dump_pkl_files, fetch_files from tokenizer import BpeTokenizer +import torch import os -def processor(intput_str: str): - return f"{intput_str} " +def get_processor(tokenizer: BpeTokenizer): + def processor(intput_dict: dict) -> dict: + segment = intput_dict["text"] + ids = tokenizer.encode(f"{segment} ") + t_ids = torch.tensor(ids, dtype=torch.int32) + + return {'sequence': t_ids} + + return processor + if __name__ == "__main__": tokenizer = BpeTokenizer("tokenizer.json") @@ -13,9 +22,10 @@ if __name__ == "__main__": os.path.join("dataset", "english-wiki"), os.path.join("dataset", "chinese-wiki"), ] + base_out_dir = "pkl_output" files = [] for dir_path in base_dir: files.extend(fetch_files(dir_path)) - - dump_pkl_files(tokenizer, files, base_out_dir, processor) \ No newline at end of file + processor = get_processor(tokenizer) + dump_pkl_files(files, base_out_dir, processor, ["text"]) \ No newline at end of file diff --git a/utils.py b/utils.py index f73dff2..97418e9 100644 --- a/utils.py +++ b/utils.py @@ -25,12 +25,6 @@ def comprehensive_normalization(text): pattern = re.compile('|'.join(re.escape(k) for k in replacements)) 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]: packages = [] sequences.sort(key=lambda x: x.numel(), reverse=True)