186 lines
7.8 KiB
Python
186 lines
7.8 KiB
Python
# 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
|