refactor(utils): 重构数据处理逻辑

This commit is contained in:
2025-08-01 21:52:00 +08:00
parent 5a1dc4ac38
commit 9c3ecc252c
3 changed files with 38 additions and 35 deletions
+2 -13
View File
@@ -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
import torch
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__":
tokenizer = BpeTokenizer("tokenizer.json")
@@ -27,5 +16,5 @@ if __name__ == "__main__":
files = []
for dir_path in base_dir:
files.extend(fetch_files(dir_path))
processor = get_processor(tokenizer)
processor = get_pt_processor(tokenizer)
dump_pkl_files(files, base_out_dir, processor, ["text"])
+2 -20
View File
@@ -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"])
+32
View File
@@ -1,5 +1,6 @@
from typing import Dict, List, Callable, Union
from datasets import DatasetDict
from .tokenizer import BpeTokenizer
from tqdm import tqdm
from torch import Tensor
@@ -14,6 +15,7 @@ def fetch_files(directory):
return [os.path.join(root, f)
for root, _, files in os.walk(directory) for f in files]
def comprehensive_normalization(text):
replacements = {
'\u2018': "'", '\u2019': "'", '\u0060': "'",
@@ -25,6 +27,7 @@ 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 pack_sequences(sequences: List[Tensor], pack_size: int, pad_value: int) -> List[Tensor]:
packages = []
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
def dump_pkl_files(
files: List[str],
base_out_dir: str,
@@ -106,6 +110,34 @@ def dump_pkl_files(
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(
dataset_dict: DatasetDict,