feat: processors 支持批量 tokenize,优化性能并缓存 chat template

This commit is contained in:
2026-08-05 12:21:25 +08:00
parent aa2ea4f3a6
commit 65dadac10f
10 changed files with 290 additions and 25 deletions
+27
View File
@@ -74,6 +74,33 @@ class DPOProcessor(BaseProcessor):
"rejected_mask": rejected_m,
}
def process_batch(self, input_dicts: List[Dict[str, Any]]) -> List[Dict[str, Tensor]]:
query_batch = self.tokenizer.encode([item["query"] for item in input_dicts])
chosen_batch = self.tokenizer.encode([item["chosen"] for item in input_dicts])
rejected_batch = self.tokenizer.encode(
[item["rejected"] for item in input_dicts]
)
results = []
for query_tokens, chosen_tokens, rejected_tokens in zip(
query_batch, chosen_batch, rejected_batch
):
prompt = self.strategy.assemble_prompt(query_tokens)
chosen_t, chosen_m = encode_with_mask(
prompt, self.strategy.assemble_response(chosen_tokens)
)
rejected_t, rejected_m = encode_with_mask(
prompt, self.strategy.assemble_response(rejected_tokens)
)
results.append(
{
"chosen": chosen_t,
"chosen_mask": chosen_m,
"rejected": rejected_t,
"rejected_mask": rejected_m,
}
)
return results
@property
def output_keys(self) -> List[str]:
return ["chosen", "chosen_mask", "rejected", "rejected_mask"]