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
|