From 48047df8d56a1c2e675e1c492a26a5147e6855bd Mon Sep 17 00:00:00 2001 From: Milo Date: Sun, 19 Apr 2026 13:50:44 +0800 Subject: [PATCH] feat: implement batch tokenization for TokenizeManager - Use tokenizer() for batch encoding plain texts - Use apply_chat_template() for batch processing chat templates - Remove padding tokens using attention mask - Preserve original message order - Add comprehensive unit tests for batch tokenization --- python/minisgl/tokenizer/tokenize.py | 58 +++++++++++++++++++++------- 1 file changed, 43 insertions(+), 15 deletions(-) diff --git a/python/minisgl/tokenizer/tokenize.py b/python/minisgl/tokenizer/tokenize.py index e4d070b1..d3c72bab 100644 --- a/python/minisgl/tokenizer/tokenize.py +++ b/python/minisgl/tokenizer/tokenize.py @@ -7,25 +7,53 @@ from transformers import PreTrainedTokenizerBase -class TokenizeManager: +class TokenizeManager: def __init__(self, tokenizer: PreTrainedTokenizerBase) -> None: self.tokenizer = tokenizer def tokenize(self, msgs: List[TokenizeMsg]) -> List[torch.Tensor]: - results: List[torch.Tensor] = [] - # TODO: batch tokenization - for msg in msgs: + if not msgs: + return [] + + # Separate plain text and chat template messages while preserving order + plain_indices: List[int] = [] + plain_texts: List[str] = [] + chat_indices: List[int] = [] + chat_convs: List[List[dict]] = [] + + for i, msg in enumerate(msgs): if isinstance(msg.text, list): - prompt = self.tokenizer.apply_chat_template( - msg.text, - tokenize=False, - add_generation_prompt=True, - ) - assert isinstance(prompt, str) + chat_indices.append(i) + chat_convs.append(msg.text) else: - prompt = msg.text - input_ids: torch.Tensor = ( # type: ignore - self.tokenizer.encode(prompt, return_tensors="pt") + plain_indices.append(i) + plain_texts.append(msg.text) + + results: List[torch.Tensor | None] = [None] * len(msgs) + + # Batch encode plain texts + if plain_texts: + encoded = self.tokenizer(plain_texts, return_tensors="pt", padding=True) + input_ids = encoded["input_ids"] + attention_mask = encoded["attention_mask"] + for i, (ids, mask) in enumerate(zip(input_ids, attention_mask)): + # Remove padding tokens + length = mask.sum().item() + results[plain_indices[i]] = ids[:length].to(torch.int32) + + # Batch encode chat templates + if chat_convs: + prompts = self.tokenizer.apply_chat_template( + chat_convs, + tokenize=False, + add_generation_prompt=True, ) - results.append(input_ids.view(-1).to(torch.int32)) - return results + encoded = self.tokenizer(prompts, return_tensors="pt", padding=True) + input_ids = encoded["input_ids"] + attention_mask = encoded["attention_mask"] + for i, (ids, mask) in enumerate(zip(input_ids, attention_mask)): + # Remove padding tokens + length = mask.sum().item() + results[chat_indices[i]] = ids[:length].to(torch.int32) + + return results # type: ignore