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 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
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 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
View File
@@ -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,