refactor(sft_file, utils): 重构数据处理流程并支持新数据集

This commit is contained in:
2025-07-21 15:11:57 +08:00
parent 56b64ab8c3
commit b6c02c04b2
2 changed files with 42 additions and 46 deletions
+2 -1
View File
@@ -5,7 +5,8 @@ import os
if __name__ == "__main__": if __name__ == "__main__":
tokenizer = BpeTokenizer("tokenizer.json") tokenizer = BpeTokenizer("tokenizer.json")
base_dir = [ base_dir = [
os.path.join("dataset", "belle-sft"), # os.path.join("dataset", "belle-sft"),
os.path.join("dataset", "chinese-instruct")
] ]
base_out_dir = "pkl_output" base_out_dir = "pkl_output"
files = [] files = []
+7 -12
View File
@@ -27,21 +27,18 @@ 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 dump_pkl_files( def dump_pkl_files(
tokenizer: BpeTokenizer, tokenizer: BpeTokenizer,
files: List[str], files: List[str],
base_out_dir: str, base_out_dir: str,
encoder: Callable[[str], str]=None, process_func: Callable[[dict], str],
key: str='text', packing_size: int = -1
packing_size: int=None
): ):
def process_line(line: str) -> Tensor: def process_line(line: str) -> Tensor:
line = json.loads(line)[key] dict_line = json.loads(line)
processed_line = encoder(line) if encoder else line tokens = process_func(dict_line)
ids = tokenizer.encode(processed_line) ids = tokenizer.encode(tokens)
arrow = torch.tensor(ids, dtype=torch.int32) return torch.tensor(ids, dtype=torch.int32)
return arrow
for file_path in files: for file_path in files:
out_file_name = os.path.basename(file_path).replace(".jsonl", ".pkl") out_file_name = os.path.basename(file_path).replace(".jsonl", ".pkl")
@@ -55,8 +52,7 @@ def dump_pkl_files(
for line in tqdm(lines, desc=f"Processing {file_name}", leave=False): for line in tqdm(lines, desc=f"Processing {file_name}", leave=False):
arrow = process_line(line) arrow = process_line(line)
arrows.append(arrow) arrows.append(arrow)
if packing_size > 0:
if packing_size is None:
with open(out_file_path, "wb") as f: with open(out_file_path, "wb") as f:
package_tensor = torch.cat(arrows) package_tensor = torch.cat(arrows)
pkl.dump(package_tensor, f) pkl.dump(package_tensor, f)
@@ -90,7 +86,6 @@ def dump_pkl_files(
pkl.dump(package_tensor, f) pkl.dump(package_tensor, f)
def process_dataset( def process_dataset(
dataset_dict: DatasetDict, dataset_dict: DatasetDict,
output_subdir: str, output_subdir: str,