From e844680025bea6ff35106e29fbf42e6924d6de0c Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Mon, 29 Sep 2025 14:13:46 +0800 Subject: [PATCH] =?UTF-8?q?fix(utils):=20=E9=87=8D=E5=91=BD=E5=90=8D?= =?UTF-8?q?=E5=8F=98=E9=87=8F=E5=B9=B6=E7=BB=9F=E4=B8=80=E5=BC=A0=E9=87=8F?= =?UTF-8?q?=E8=BD=AC=E6=8D=A2=E9=80=BB=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- modules/utils.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/modules/utils.py b/modules/utils.py index e946e0d..619e40a 100644 --- a/modules/utils.py +++ b/modules/utils.py @@ -121,10 +121,10 @@ def dump_pkl_files( def get_pt_processor(tokenizer: BpeTokenizer): def processor(intput_dict: dict) -> dict: segment = intput_dict["text"] - ids = tokenizer.encode(f"{segment} ") - t_ids = torch.tensor(ids, dtype=torch.int32) + tokens = tokenizer.encode(f"{segment}") + tokens = torch.tensor(tokens, dtype=torch.int32) - return {'sequence': t_ids} + return {'sequence': tokens} return processor @@ -133,6 +133,7 @@ def get_sft_processor(tokenizer: BpeTokenizer): def processor(input_dict: dict): query, response = input_dict["query"], input_dict["response"] tokens = tokenizer.encode(f"<|user|> {query} <|system|> {response}\n") + tokens = torch.tensor(tokens, dtype=torch.int32) return {"sequence": tokens} @@ -169,7 +170,7 @@ def cache_files(tokenizer, files, base_out_dir, cache_type): keys = ["sequence"] elif cache_type == "sft": processor = get_sft_processor(tokenizer) - keys = ["sequence", "mask"] + keys = ["sequence"] elif cache_type == "dpo": processor = get_dpo_processor(tokenizer) keys = ["chosen", "chosen_mask", "rejected", "rejected_mask"]