refactor(utils): 重构数据处理流程并添加 DPO 数据缓存支持

This commit is contained in:
2025-08-10 11:27:02 +08:00
parent 3966481bfe
commit ce263985d9
3 changed files with 16 additions and 28 deletions
-14
View File
@@ -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"])
-14
View File
@@ -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"])
+16
View File
@@ -148,6 +148,22 @@ def get_sft_processor(tokenizer: BpeTokenizer):
return processor 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( def process_dataset(
dataset_dict: DatasetDict, dataset_dict: DatasetDict,