From 9c3ecc252c59e754eb8ca600dc9b870e88840d41 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Fri, 1 Aug 2025 21:52:00 +0800 Subject: [PATCH] =?UTF-8?q?refactor(utils):=20=E9=87=8D=E6=9E=84=E6=95=B0?= =?UTF-8?q?=E6=8D=AE=E5=A4=84=E7=90=86=E9=80=BB=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- dump_pt_file.py | 15 ++------------- dump_sft_file.py | 22 ++-------------------- modules/utils.py | 36 ++++++++++++++++++++++++++++++++++-- 3 files changed, 38 insertions(+), 35 deletions(-) diff --git a/dump_pt_file.py b/dump_pt_file.py index f924aa1..e45794b 100644 --- a/dump_pt_file.py +++ b/dump_pt_file.py @@ -1,18 +1,7 @@ -from modules.utils import dump_pkl_files, fetch_files +from modules.utils import dump_pkl_files, fetch_files, get_pt_processor from modules.tokenizer import BpeTokenizer -import torch import os -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") @@ -27,5 +16,5 @@ if __name__ == "__main__": files = [] for dir_path in base_dir: files.extend(fetch_files(dir_path)) - processor = get_processor(tokenizer) + processor = get_pt_processor(tokenizer) dump_pkl_files(files, base_out_dir, processor, ["text"]) \ No newline at end of file diff --git a/dump_sft_file.py b/dump_sft_file.py index 1a5db3c..145b623 100644 --- a/dump_sft_file.py +++ b/dump_sft_file.py @@ -1,25 +1,7 @@ -from modules.utils import dump_pkl_files, fetch_files +from modules.utils import dump_pkl_files, fetch_files, get_sft_processor from modules.tokenizer import BpeTokenizer -import torch import os -def get_processor(tokenizer: BpeTokenizer): - def processor(input_dict: dict): - query, response = input_dict["query"], input_dict["response"] - prefix_seg = f"<|user|> {query} <|system|> " - suffix_seg = f"{response}\n" - prefix_ids = tokenizer.encode(prefix_seg) - suffix_ids = tokenizer.encode(suffix_seg) - - tokens = prefix_ids + suffix_ids - tokens = torch.tensor(tokens, dtype=torch.int32) - masks = torch.zeros_like(tokens, dtype=torch.bool) - masks[len(prefix_ids):] = True - - return {"sequence": tokens, "mask": masks} - - return processor - if __name__ == "__main__": tokenizer = BpeTokenizer("tokenizer.json") @@ -32,6 +14,6 @@ if __name__ == "__main__": for dir_path in base_dir: files.extend(fetch_files(dir_path)) - processor = get_processor(tokenizer) + processor = get_sft_processor(tokenizer) dump_pkl_files(files, base_out_dir,processor, ["sequence", "mask"]) \ No newline at end of file diff --git a/modules/utils.py b/modules/utils.py index 9174d24..d62eb5f 100644 --- a/modules/utils.py +++ b/modules/utils.py @@ -1,5 +1,6 @@ from typing import Dict, List, Callable, Union from datasets import DatasetDict +from .tokenizer import BpeTokenizer from tqdm import tqdm from torch import Tensor @@ -14,6 +15,7 @@ def fetch_files(directory): return [os.path.join(root, f) for root, _, files in os.walk(directory) for f in files] + def comprehensive_normalization(text): replacements = { '\u2018': "'", '\u2019': "'", '\u0060': "'", @@ -25,6 +27,7 @@ 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 pack_sequences(sequences: List[Tensor], pack_size: int, pad_value: int) -> List[Tensor]: packages = [] sequences.sort(key=lambda x: x.numel(), reverse=True) @@ -63,6 +66,7 @@ def pack_sequences(sequences: List[Tensor], pack_size: int, pad_value: int) -> L return packages + def dump_pkl_files( files: List[str], base_out_dir: str, @@ -104,8 +108,36 @@ def dump_pkl_files( with open(out_file_path, "wb") as f: pkl.dump(output_package, f) - - + + +def get_pt_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 + + +def get_sft_processor(tokenizer: BpeTokenizer): + def processor(input_dict: dict): + query, response = input_dict["query"], input_dict["response"] + prefix_seg = f"<|user|> {query} <|system|> " + suffix_seg = f"{response}\n" + prefix_ids = tokenizer.encode(prefix_seg) + suffix_ids = tokenizer.encode(suffix_seg) + + tokens = prefix_ids + suffix_ids + tokens = torch.tensor(tokens, dtype=torch.int32) + masks = torch.zeros_like(tokens, dtype=torch.bool) + masks[len(prefix_ids):] = True + + return {"sequence": tokens, "mask": masks} + + return processor + def process_dataset( dataset_dict: DatasetDict,