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(