refactor(utils): 简化SFT处理器逻辑

This commit is contained in:
2025-09-28 15:50:10 +08:00
parent 44be73d255
commit c3ebc4f681
+3 -11
View File
@@ -132,17 +132,9 @@ def get_pt_processor(tokenizer: BpeTokenizer):
def get_sft_processor(tokenizer: BpeTokenizer): def get_sft_processor(tokenizer: BpeTokenizer):
def processor(input_dict: dict): def processor(input_dict: dict):
query, response = input_dict["query"], input_dict["response"] query, response = input_dict["query"], input_dict["response"]
prefix_seg = f"<|user|> {query} <|system|> <bos>" tokens = tokenizer.encode(f"<|user|> {query} <|system|> <bos>{response}<eos>\n")
suffix_seg = f"{response}<eos>\n"
prefix_ids = tokenizer.encode(prefix_seg) return {"sequence": tokens}
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}
return processor return processor