From fcd2b66b5c90c4c785cd9ea74d7e8fb27e7cda4d Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sun, 10 Aug 2025 14:22:48 +0800 Subject: [PATCH] =?UTF-8?q?feat(utils):=20=E6=B7=BB=E5=8A=A0=20DPO=20?= =?UTF-8?q?=E8=AE=AD=E7=BB=83=E6=A8=A1=E5=BC=8F=E7=9A=84=E5=A4=84=E7=90=86?= =?UTF-8?q?=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- modules/utils.py | 26 +++++++++++++++++++++++++- 1 file changed, 25 insertions(+), 1 deletion(-) diff --git a/modules/utils.py b/modules/utils.py index e6bb50b..caab0a9 100644 --- a/modules/utils.py +++ b/modules/utils.py @@ -148,6 +148,29 @@ def get_sft_processor(tokenizer: BpeTokenizer): return processor +def get_dpo_processor(tokenizer: BpeTokenizer): + def processor(input_dict: dict): + prompt, chosen, rejected = input_dict["prompt"], input_dict["chosen"], input_dict["rejected"] + prefix_seg = f"<|user|> {prompt} <|system|> " + chosen_seg = f"{chosen}\n" + rejected_seg = f"{rejected}\n" + prefix_ids = tokenizer.encode(prefix_seg) + chosen_ids = tokenizer.encode(chosen_seg) + rejected_ids = tokenizer.encode(rejected_seg) + + chosen_seq = torch.tensor(prefix_ids + chosen_ids, dtype=torch.int32) + chosen_mask = torch.zeros_like(chosen_seq, dtype=torch.bool) + chosen_mask[len(prefix_ids):] = True + + rejected_seq = torch.tensor(prefix_ids + rejected_ids, dtype=torch.int32) + resjected_mask = torch.zeros_like(rejected_seq, dtype=torch.bool) + resjected_mask[len(prefix_ids):] = True + + return {"chosen": chosen_seq, "chosen_mask": chosen_mask, "rejected": rejected_seq, "rejected_mask": resjected_mask} + + return processor + + def cache_files(tokenizer, files, base_out_dir, cache_type): processor = None keys = [] @@ -158,7 +181,8 @@ def cache_files(tokenizer, files, base_out_dir, cache_type): processor = get_sft_processor(tokenizer) keys = ["query", "response"] elif cache_type == "dpo": - keys = ["query", "response"] + processor = get_dpo_processor(tokenizer) + keys = ["chosen", "chosen_mask", "rejected", "rejected_mask"] else: raise ValueError("Invalid cache type")