From 5ea1c6355db51d3a42efe19f3b0a0e4aa4567367 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Tue, 22 Jul 2025 10:03:06 +0800 Subject: [PATCH] =?UTF-8?q?refactor(data):=20=E9=87=8D=E6=9E=84=E6=95=B0?= =?UTF-8?q?=E6=8D=AE=E5=A4=84=E7=90=86=E6=B5=81=E7=A8=8B=E4=BB=A5=E9=80=82?= =?UTF-8?q?=E5=BA=94=E6=96=B0=E9=9C=80=E6=B1=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- chinese-instruct.py | 2 +- dump_sft_file.py | 4 ++-- utils.py | 5 ++++- 3 files changed, 7 insertions(+), 4 deletions(-) diff --git a/chinese-instruct.py b/chinese-instruct.py index d3048db..8a2d613 100644 --- a/chinese-instruct.py +++ b/chinese-instruct.py @@ -12,7 +12,7 @@ def replace_seg(query:str, response:str) -> str: query = query.replace(old, new) response = response.replace(old, new) - return {"query": query, "response": response} + return query, response def process_func(input_dict: dict): query = input_dict["prompt"] if input_dict["prompt"] else "" diff --git a/dump_sft_file.py b/dump_sft_file.py index 70b2670..6a2ca7a 100644 --- a/dump_sft_file.py +++ b/dump_sft_file.py @@ -16,7 +16,7 @@ def get_processor(tokenizer: BpeTokenizer): masks = torch.zeros_like(tokens) masks[:len(prefix_ids)] = 1 - return tokens, masks + return {"sequence": tokens, "mask": masks} return processor @@ -34,4 +34,4 @@ if __name__ == "__main__": processor = get_processor(tokenizer) - dump_pkl_files(files, base_out_dir,) \ No newline at end of file + dump_pkl_files(files, base_out_dir,processor, ["sequence", "mask"]) \ No newline at end of file diff --git a/utils.py b/utils.py index 0f7b11d..5244f13 100644 --- a/utils.py +++ b/utils.py @@ -79,12 +79,15 @@ def dump_pkl_files( os.makedirs(os.path.dirname(out_file_path), exist_ok=True) arrows: Dict[str, List[Tensor]] = {} + for key in output_keys: + arrows[key] = [] with open(file_path, "r") as f: lines = f.readlines() for line in tqdm(lines, desc=f"Processing {file_name}", leave=False): - arrow = process_func(line) + line_dict = json.loads(line) + arrow = process_func(line_dict) for key in output_keys: arrows[key].extend(arrow[key])