myLightningOPD / slime /utils /mask_utils.py
ayh015's picture
Upload folder using huggingface_hub
6011e08 verified
Raw
History Blame Contribute Delete
8 kB
# 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