diff --git a/sft_file.py b/sft_file.py index 66579ef..6038797 100644 --- a/sft_file.py +++ b/sft_file.py @@ -5,7 +5,8 @@ import os if __name__ == "__main__": tokenizer = BpeTokenizer("tokenizer.json") 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" files = [] diff --git a/utils.py b/utils.py index 392a507..0a390e5 100644 --- a/utils.py +++ b/utils.py @@ -27,22 +27,19 @@ 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 dump_pkl_files( - tokenizer: BpeTokenizer, + tokenizer: BpeTokenizer, files: List[str], - base_out_dir: str, - encoder: Callable[[str], str]=None, - key: str='text', - packing_size: int=None -): + base_out_dir: str, + process_func: Callable[[dict], str], + packing_size: int = -1 +): def process_line(line: str) -> Tensor: - line = json.loads(line)[key] - processed_line = encoder(line) if encoder else line - ids = tokenizer.encode(processed_line) - arrow = torch.tensor(ids, dtype=torch.int32) - return arrow - + dict_line = json.loads(line) + tokens = process_func(dict_line) + ids = tokenizer.encode(tokens) + return torch.tensor(ids, dtype=torch.int32) + for file_path in files: out_file_name = os.path.basename(file_path).replace(".jsonl", ".pkl") out_file_path = os.path.join(base_out_dir, out_file_name) @@ -55,40 +52,38 @@ def dump_pkl_files( for line in tqdm(lines, desc=f"Processing {file_name}", leave=False): arrow = process_line(line) arrows.append(arrow) - - if packing_size is None: - with open(out_file_path, "wb") as f: - package_tensor = torch.cat(arrows) - pkl.dump(package_tensor, f) - else: - arrows.sort(key=lambda x: x.numel(), reverse=True) - packages = [] - cur_size = 0 - cur_tensor = torch.tensor([]) - - for i in tqdm(range(0, len(arrows)), desc=f"Packing {file_name}", leave=False): - cur_ids = arrows[i] - if cur_ids.numel() <= packing_size: - if cur_ids.numel() + cur_tensor.numel() <= packing_size: - cur_size += cur_ids.numel() - cur_tensor = torch.cat([cur_tensor, cur_ids]) + if packing_size > 0: + with open(out_file_path, "wb") as f: + package_tensor = torch.cat(arrows) + pkl.dump(package_tensor, f) + else: + arrows.sort(key=lambda x: x.numel(), reverse=True) + packages = [] + cur_size = 0 + cur_tensor = torch.tensor([]) + + for i in tqdm(range(0, len(arrows)), desc=f"Packing {file_name}", leave=False): + cur_ids = arrows[i] + if cur_ids.numel() <= packing_size: + if cur_ids.numel() + cur_tensor.numel() <= packing_size: + cur_size += cur_ids.numel() + cur_tensor = torch.cat([cur_tensor, cur_ids]) + else: + cur_tensor = F.pad( + cur_tensor, + (0, packing_size - cur_tensor.numel()), + 'constant', + tokenizer.pad_id + ) + packages.append(cur_tensor) + cur_tensor = cur_ids else: - cur_tensor = F.pad( - cur_tensor, - (0, packing_size - cur_tensor.numel()), - 'constant', - tokenizer.pad_id - ) - packages.append(cur_tensor) - cur_tensor = cur_ids - else: - packages.append(cur_ids[:packing_size]) - - - with open(out_file_path, "wb") as f: - package_tensor = torch.cat(packages) - pkl.dump(package_tensor, f) + packages.append(cur_ids[:packing_size]) + + with open(out_file_path, "wb") as f: + package_tensor = torch.cat(packages) + pkl.dump(package_tensor, f) def process_dataset(