refactor(utils): 重构数据处理逻辑
This commit is contained in:
+2
-20
@@ -1,25 +1,7 @@
|
||||
from modules.utils import dump_pkl_files, fetch_files
|
||||
from modules.utils import dump_pkl_files, fetch_files, get_sft_processor
|
||||
from modules.tokenizer import BpeTokenizer
|
||||
import torch
|
||||
import os
|
||||
|
||||
def get_processor(tokenizer: BpeTokenizer):
|
||||
def processor(input_dict: dict):
|
||||
query, response = input_dict["query"], input_dict["response"]
|
||||
prefix_seg = f"<|user|> {query} <|system|> <bos>"
|
||||
suffix_seg = f"{response}<eos>\n"
|
||||
prefix_ids = tokenizer.encode(prefix_seg)
|
||||
suffix_ids = tokenizer.encode(suffix_seg)
|
||||
|
||||
tokens = prefix_ids + suffix_ids
|
||||
tokens = torch.tensor(tokens, dtype=torch.int32)
|
||||
masks = torch.zeros_like(tokens, dtype=torch.bool)
|
||||
masks[len(prefix_ids):] = True
|
||||
|
||||
return {"sequence": tokens, "mask": masks}
|
||||
|
||||
return processor
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
tokenizer = BpeTokenizer("tokenizer.json")
|
||||
@@ -32,6 +14,6 @@ if __name__ == "__main__":
|
||||
for dir_path in base_dir:
|
||||
files.extend(fetch_files(dir_path))
|
||||
|
||||
processor = get_processor(tokenizer)
|
||||
processor = get_sft_processor(tokenizer)
|
||||
|
||||
dump_pkl_files(files, base_out_dir,processor, ["sequence", "mask"])
|
||||
Reference in New Issue
Block a user