# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 from transformers import AutoTokenizer def get_response_lengths(loss_masks: list[list[int]]) -> list[int]: return [mask.count(1) if 1 in mask else 0 for mask in loss_masks] class MultiTurnLossMaskGenerator: def __init__(self, tokenizer: AutoTokenizer, tokenizer_type: str = "qwen"): self.tokenizer = tokenizer self.system_message_length, self.gen_token_length = self.get_system_message_length() self.tokenizer_type = tokenizer_type def get_response_lengths(self, loss_masks: list[list[int]]) -> list[int]: return get_response_lengths(loss_masks) def find_all_sublist_indices(self, main_list, sublist): sublist_len = len(sublist) indices = [] for i in range(len(main_list) - sublist_len + 1): if main_list[i : i + sublist_len] == sublist: indices.append(i) return indices def get_system_message_length(self) -> tuple[int, int]: test_string = "FOR TESTING ONLY" test_messages = [ {"role": "user", "content": test_string}, {"role": "user", "content": test_string}, ] raw_token_ids = self.tokenizer(test_string, add_special_tokens=False)["input_ids"] chat_template_token = self.tokenizer.apply_chat_template( test_messages, add_special_tokens=False, tokenize=False ) chat_template_token_ids = self.tokenizer(chat_template_token, add_special_tokens=False)["input_ids"] idx_1, idx_2 = self.find_all_sublist_indices(chat_template_token_ids, raw_token_ids) end_interval = len(chat_template_token_ids) - len(raw_token_ids) - idx_2 gen_token_length = len( self.tokenizer.apply_chat_template( test_messages, add_special_tokens=False, tokenize=True, add_generation_prompt=True ) ) - len(chat_template_token_ids) system_message_length = idx_1 - ((idx_2 - idx_1) - end_interval - len(raw_token_ids)) return system_message_length, gen_token_length def gen_multi_turn_loss_mask_qwen( self, messages: list[dict], tools: list[dict] = None ) -> tuple[list[int], list[int]]: all_loss_masks = [] all_token_ids = [] for i, message in enumerate(messages): if i == 0: message_ids = self.tokenizer.apply_chat_template([message], tokenize=True, tools=tools) else: message_ids = self.tokenizer.apply_chat_template([message], tokenize=True) if message["role"] != "system" and i > 0: message_ids = message_ids[self.system_message_length :] if message["role"] == "assistant": loss_mask = [0] * self.gen_token_length + [1] * (len(message_ids) - self.gen_token_length) else: loss_mask = [0] * len(message_ids) if message.get("step_loss_mask", 1) != 1: loss_mask = [0] * len(message_ids) all_loss_masks.extend(loss_mask) all_token_ids.extend(message_ids) return all_token_ids, all_loss_masks def gen_multi_turn_loss_mask_qwen3( self, messages: list[dict], tools: list[dict] = None ) -> tuple[list[int], list[int]]: all_loss_masks = [] all_token_ids = [] prefix_message = {"role": "user", "content": "FOR CALCULATING LOSS MASK ONLY"} prefix_token_ids = self.tokenizer.apply_chat_template([prefix_message], tokenize=True) for i, message in enumerate(messages): if i == 0: tailed_message_ids = self.tokenizer.apply_chat_template( [message, prefix_message], tokenize=True, tools=tools ) message_ids = tailed_message_ids[: -len(prefix_token_ids)] else: prefixed_message_ids = self.tokenizer.apply_chat_template([prefix_message, message], tokenize=True) message_ids = prefixed_message_ids[len(prefix_token_ids) :] if message["role"] != "system" and i > 0: message_ids = message_ids[self.system_message_length :] if message["role"] == "assistant": loss_mask = [0] * self.gen_token_length + [1] * (len(message_ids) - self.gen_token_length) else: loss_mask = [0] * len(message_ids) if message.get("step_loss_mask", 1) != 1: loss_mask = [0] * len(message_ids) all_loss_masks.extend(loss_mask) all_token_ids.extend(message_ids) return all_token_ids, all_loss_masks def gen_multi_turn_loss_mask_distill_qwen( self, messages: list[dict], tools: list[dict] = None ) -> tuple[list[int], list[int]]: prompt = self.tokenizer.apply_chat_template( messages[:1], tokenize=False, add_generation_prompt=True, tools=tools ) response = messages[-1]["content"] prompt_tokens = self.tokenizer(prompt, add_special_tokens=False)["input_ids"] response_tokens = self.tokenizer(response, add_special_tokens=False)["input_ids"] response_length = len(response_tokens) token_ids = prompt_tokens + response_tokens loss_mask = [0] * len(prompt_tokens) + [1] * response_length if messages[-1].get("step_loss_mask", 1) != 1: loss_mask = [0] * len(token_ids) return token_ids, loss_mask def get_loss_mask(self, messages: list[dict], tools: list[dict] = None) -> tuple[list[int], list[int]]: if self.tokenizer_type == "qwen": if "<|Assistant|>" in self.tokenizer.get_added_vocab(): return self.gen_multi_turn_loss_mask_distill_qwen(messages, tools) return self.gen_multi_turn_loss_mask_qwen(messages, tools) elif self.tokenizer_type == "qwen3": return self.gen_multi_turn_loss_mask_qwen3(messages, tools) elif self.tokenizer_type == "distill_qwen": return self.gen_multi_turn_loss_mask_distill_qwen(messages, tools) else: raise ValueError(f"Unsupported tokenizer type: {self.tokenizer_type}") def get_loss_mask_with_multimodal_alignment( self, messages: list[dict], input_ids: list[int], tools: list[dict] = None ) -> tuple[list[int], list[int]]: text = [] for msg in messages: if isinstance(msg.get("content"), list): text_parts = [] for item in msg["content"]: if isinstance(item, dict) and item.get("type") == "text": text_parts.append(item.get("text", "")) elif isinstance(item, str): text_parts.append(item) text.append({"role": msg["role"], "content": " ".join(text_parts)}) else: text.append(msg) _, loss_mask_text = self.get_loss_mask(text, tools=tools) diff = len(input_ids) - len(loss_mask_text) assert diff >= 0, ( f"input_ids (length={len(input_ids)}) is shorter than text loss_mask (length={len(loss_mask_text)}) " f"Please check if processor and tokenizer tokenization are consistent." ) loss_mask = [0] * diff + loss_mask_text return input_ids, loss_mask def get_text_from_loss_mask(self, token_ids: list[int], loss_masks: list[int]) -> list[str]: selected_texts = [] current_tokens = [] for idx, mask in enumerate(loss_masks): if mask == 1: current_tokens.append(token_ids[idx]) elif current_tokens: selected_texts.append(self.tokenizer.decode(current_tokens)) current_tokens = [] if current_tokens: selected_texts.append(self.tokenizer.decode(current_tokens)) return selected_texts