refactor(utils): 重构数据处理逻辑
This commit is contained in:
+2
-13
@@ -1,18 +1,7 @@
|
|||||||
from modules.utils import dump_pkl_files, fetch_files
|
from modules.utils import dump_pkl_files, fetch_files, get_pt_processor
|
||||||
from modules.tokenizer import BpeTokenizer
|
from modules.tokenizer import BpeTokenizer
|
||||||
import torch
|
|
||||||
import os
|
import os
|
||||||
|
|
||||||
def get_processor(tokenizer: BpeTokenizer):
|
|
||||||
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")
|
||||||
@@ -27,5 +16,5 @@ if __name__ == "__main__":
|
|||||||
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)
|
processor = get_pt_processor(tokenizer)
|
||||||
dump_pkl_files(files, base_out_dir, processor, ["text"])
|
dump_pkl_files(files, base_out_dir, processor, ["text"])
|
||||||
+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
|
from modules.tokenizer import BpeTokenizer
|
||||||
import torch
|
|
||||||
import os
|
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__":
|
if __name__ == "__main__":
|
||||||
tokenizer = BpeTokenizer("tokenizer.json")
|
tokenizer = BpeTokenizer("tokenizer.json")
|
||||||
@@ -32,6 +14,6 @@ if __name__ == "__main__":
|
|||||||
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)
|
processor = get_sft_processor(tokenizer)
|
||||||
|
|
||||||
dump_pkl_files(files, base_out_dir,processor, ["sequence", "mask"])
|
dump_pkl_files(files, base_out_dir,processor, ["sequence", "mask"])
|
||||||
+34
-2
@@ -1,5 +1,6 @@
|
|||||||
from typing import Dict, List, Callable, Union
|
from typing import Dict, List, Callable, Union
|
||||||
from datasets import DatasetDict
|
from datasets import DatasetDict
|
||||||
|
from .tokenizer import BpeTokenizer
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
@@ -14,6 +15,7 @@ def fetch_files(directory):
|
|||||||
return [os.path.join(root, f)
|
return [os.path.join(root, f)
|
||||||
for root, _, files in os.walk(directory) for f in files]
|
for root, _, files in os.walk(directory) for f in files]
|
||||||
|
|
||||||
|
|
||||||
def comprehensive_normalization(text):
|
def comprehensive_normalization(text):
|
||||||
replacements = {
|
replacements = {
|
||||||
'\u2018': "'", '\u2019': "'", '\u0060': "'",
|
'\u2018': "'", '\u2019': "'", '\u0060': "'",
|
||||||
@@ -25,6 +27,7 @@ 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 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)
|
||||||
@@ -63,6 +66,7 @@ def pack_sequences(sequences: List[Tensor], pack_size: int, pad_value: int) -> L
|
|||||||
|
|
||||||
return packages
|
return packages
|
||||||
|
|
||||||
|
|
||||||
def dump_pkl_files(
|
def dump_pkl_files(
|
||||||
files: List[str],
|
files: List[str],
|
||||||
base_out_dir: str,
|
base_out_dir: str,
|
||||||
@@ -104,8 +108,36 @@ def dump_pkl_files(
|
|||||||
|
|
||||||
with open(out_file_path, "wb") as f:
|
with open(out_file_path, "wb") as f:
|
||||||
pkl.dump(output_package, f)
|
pkl.dump(output_package, f)
|
||||||
|
|
||||||
|
|
||||||
|
def get_pt_processor(tokenizer: BpeTokenizer):
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
def get_sft_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
|
||||||
|
|
||||||
|
|
||||||
def process_dataset(
|
def process_dataset(
|
||||||
dataset_dict: DatasetDict,
|
dataset_dict: DatasetDict,
|
||||||
|
|||||||
Reference in New Issue
Block a user