diff --git a/dump_pt_file.py b/dump_pt_file.py index 1822141..c6e17d5 100644 --- a/dump_pt_file.py +++ b/dump_pt_file.py @@ -1,9 +1,18 @@ from utils import dump_pkl_files, fetch_files from tokenizer import BpeTokenizer +import torch import os -def processor(intput_str: str): - return f"{intput_str} " +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") @@ -13,9 +22,10 @@ if __name__ == "__main__": os.path.join("dataset", "english-wiki"), os.path.join("dataset", "chinese-wiki"), ] + base_out_dir = "pkl_output" files = [] for dir_path in base_dir: files.extend(fetch_files(dir_path)) - - dump_pkl_files(tokenizer, files, base_out_dir, processor) \ No newline at end of file + processor = get_processor(tokenizer) + dump_pkl_files(files, base_out_dir, processor, ["text"]) \ No newline at end of file diff --git a/utils.py b/utils.py index f73dff2..97418e9 100644 --- a/utils.py +++ b/utils.py @@ -25,12 +25,6 @@ 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 pt_processor(text): - text = comprehensive_normalization(text) - text = text.lower() - return 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)