refactor(dump_pt_file): 重构数据处理流程并优化代码结构
This commit is contained in:
+14
-4
@@ -1,9 +1,18 @@
|
|||||||
from utils import dump_pkl_files, fetch_files
|
from utils import dump_pkl_files, fetch_files
|
||||||
from tokenizer import BpeTokenizer
|
from tokenizer import BpeTokenizer
|
||||||
|
import torch
|
||||||
import os
|
import os
|
||||||
|
|
||||||
def processor(intput_str: str):
|
def get_processor(tokenizer: BpeTokenizer):
|
||||||
return f"{intput_str} <eos>"
|
def processor(intput_dict: dict) -> dict:
|
||||||
|
segment = intput_dict["text"]
|
||||||
|
ids = tokenizer.encode(f"{segment} <eos>")
|
||||||
|
t_ids = torch.tensor(ids, dtype=torch.int32)
|
||||||
|
|
||||||
|
return {'sequence': t_ids}
|
||||||
|
|
||||||
|
return processor
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
tokenizer = BpeTokenizer("tokenizer.json")
|
tokenizer = BpeTokenizer("tokenizer.json")
|
||||||
@@ -13,9 +22,10 @@ if __name__ == "__main__":
|
|||||||
os.path.join("dataset", "english-wiki"),
|
os.path.join("dataset", "english-wiki"),
|
||||||
os.path.join("dataset", "chinese-wiki"),
|
os.path.join("dataset", "chinese-wiki"),
|
||||||
]
|
]
|
||||||
|
|
||||||
base_out_dir = "pkl_output"
|
base_out_dir = "pkl_output"
|
||||||
files = []
|
files = []
|
||||||
for dir_path in base_dir:
|
for dir_path in base_dir:
|
||||||
files.extend(fetch_files(dir_path))
|
files.extend(fetch_files(dir_path))
|
||||||
|
processor = get_processor(tokenizer)
|
||||||
dump_pkl_files(tokenizer, files, base_out_dir, processor)
|
dump_pkl_files(files, base_out_dir, processor, ["text"])
|
||||||
@@ -25,12 +25,6 @@ def comprehensive_normalization(text):
|
|||||||
pattern = re.compile('|'.join(re.escape(k) for k in replacements))
|
pattern = re.compile('|'.join(re.escape(k) for k in replacements))
|
||||||
return pattern.sub(lambda m: replacements[m.group()], text)
|
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]:
|
def pack_sequences(sequences: List[Tensor], pack_size: int, pad_value: int) -> List[Tensor]:
|
||||||
packages = []
|
packages = []
|
||||||
sequences.sort(key=lambda x: x.numel(), reverse=True)
|
sequences.sort(key=lambda x: x.numel(), reverse=True)
|
||||||
|
|||||||
Reference in New Issue
Block a user