From ce263985d958961b8b786b8c02acf83f086cae1b Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sun, 10 Aug 2025 11:27:02 +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=E6=B5=81=E7=A8=8B=E5=B9=B6=E6=B7=BB?= =?UTF-8?q?=E5=8A=A0=20DPO=20=E6=95=B0=E6=8D=AE=E7=BC=93=E5=AD=98=E6=94=AF?= =?UTF-8?q?=E6=8C=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- dump_pt_file.py | 14 -------------- dump_sft_file.py | 14 -------------- modules/utils.py | 16 ++++++++++++++++ 3 files changed, 16 insertions(+), 28 deletions(-) delete mode 100644 dump_pt_file.py delete mode 100644 dump_sft_file.py diff --git a/dump_pt_file.py b/dump_pt_file.py deleted file mode 100644 index f8dce4e..0000000 --- a/dump_pt_file.py +++ /dev/null @@ -1,14 +0,0 @@ -from modules.utils import dump_pkl_files, fetch_files, fetch_folders, get_pt_processor -from modules.tokenizer import BpeTokenizer - - -if __name__ == "__main__": - tokenizer = BpeTokenizer("tokenizer.json") - base_dir = fetch_folders("dataset") - - base_out_dir = "pkl_output" - files = [] - for dir_path in base_dir: - files.extend(fetch_files(dir_path)) - 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 deleted file mode 100644 index cb82cd5..0000000 --- a/dump_sft_file.py +++ /dev/null @@ -1,14 +0,0 @@ -from modules.utils import dump_pkl_files, fetch_files, fetch_folders, get_sft_processor -from modules.tokenizer import BpeTokenizer - -if __name__ == "__main__": - tokenizer = BpeTokenizer("tokenizer.json") - base_dir = fetch_folders("dataset") - - base_out_dir = "pkl_output" - files = [] - for dir_path in base_dir: - files.extend(fetch_files(dir_path)) - - 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 529531f..e6bb50b 100644 --- a/modules/utils.py +++ b/modules/utils.py @@ -147,6 +147,22 @@ def get_sft_processor(tokenizer: BpeTokenizer): return {"sequence": tokens, "mask": masks} return processor + +def cache_files(tokenizer, files, base_out_dir, cache_type): + processor = None + keys = [] + if cache_type == "pt": + processor = get_pt_processor(tokenizer) + keys = ["text"] + elif cache_type == "sft": + processor = get_sft_processor(tokenizer) + keys = ["query", "response"] + elif cache_type == "dpo": + keys = ["query", "response"] + else: + raise ValueError("Invalid cache type") + + dump_pkl_files(files, base_out_dir, processor, keys) def process_dataset(