diff --git a/dump_pt_file.py b/dump_pt_file.py index e45794b..f8dce4e 100644 --- a/dump_pt_file.py +++ b/dump_pt_file.py @@ -1,16 +1,10 @@ -from modules.utils import dump_pkl_files, fetch_files, get_pt_processor +from modules.utils import dump_pkl_files, fetch_files, fetch_folders, get_pt_processor from modules.tokenizer import BpeTokenizer -import os if __name__ == "__main__": tokenizer = BpeTokenizer("tokenizer.json") - base_dir = [ - os.path.join("dataset", "chinese-c4"), - os.path.join("dataset", "english-fineweb"), - os.path.join("dataset", "english-wiki"), - os.path.join("dataset", "chinese-wiki"), - ] + base_dir = fetch_folders("dataset") base_out_dir = "pkl_output" files = [] diff --git a/dump_sft_file.py b/dump_sft_file.py index 145b623..cb82cd5 100644 --- a/dump_sft_file.py +++ b/dump_sft_file.py @@ -1,19 +1,14 @@ -from modules.utils import dump_pkl_files, fetch_files, get_sft_processor +from modules.utils import dump_pkl_files, fetch_files, fetch_folders, get_sft_processor from modules.tokenizer import BpeTokenizer -import os - if __name__ == "__main__": tokenizer = BpeTokenizer("tokenizer.json") - base_dir = [ - os.path.join("dataset", "Ling-Coder-SFT"), - os.path.join("dataset", "chinese-instruct") - ] + 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 d62eb5f..529531f 100644 --- a/modules/utils.py +++ b/modules/utils.py @@ -14,6 +14,16 @@ import re def fetch_files(directory): return [os.path.join(root, f) for root, _, files in os.walk(directory) for f in files] + + +def fetch_folders(root_dir, filter_func=None): + folders = [] + for root, dirs, _ in os.walk(root_dir): + for dir_name in dirs: + folder_path = os.path.join(root, dir_name) + if filter_func is None or filter_func(folder_path): + folders.append(folder_path) + return folders def comprehensive_normalization(text):