From c3ebc4f681c0006fcec2da791b5552b0ab596dad Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sun, 28 Sep 2025 15:50:10 +0800 Subject: [PATCH] =?UTF-8?q?refactor(utils):=20=E7=AE=80=E5=8C=96SFT?= =?UTF-8?q?=E5=A4=84=E7=90=86=E5=99=A8=E9=80=BB=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- modules/utils.py | 14 +++----------- 1 file changed, 3 insertions(+), 11 deletions(-) diff --git a/modules/utils.py b/modules/utils.py index 31ab8ab..e946e0d 100644 --- a/modules/utils.py +++ b/modules/utils.py @@ -132,17 +132,9 @@ def get_pt_processor(tokenizer: BpeTokenizer): 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|> " - suffix_seg = f"{response}\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} + tokens = tokenizer.encode(f"<|user|> {query} <|system|> {response}\n") + + return {"sequence": tokens} return processor