Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions slime/backends/megatron_utils/model_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,8 @@ def wrapped_model_provider(
if args.megatron_to_hf_mode == "bridge":
from megatron.bridge import AutoBridge

import slime_plugins.megatron_bridge # noqa: F401 # register custom bridges

bridge = AutoBridge.from_hf_pretrained(args.hf_checkpoint, trust_remote_code=True)
provider = bridge.to_megatron_provider(load_weights=False)
# TODO: we should not manually set this...
Expand All @@ -91,6 +93,8 @@ def wrapped_model_provider(
provider.expert_model_parallel_size = args.expert_model_parallel_size
provider.expert_tensor_parallel_size = args.expert_tensor_parallel_size
provider.sequence_parallel = args.sequence_parallel
provider.context_parallel_size = args.context_parallel_size
provider.variable_seq_lengths = args.variable_seq_lengths
if getattr(args, "decoder_first_pipeline_num_layers", None) is not None:
provider.num_layers_in_first_pipeline_stage = args.decoder_first_pipeline_num_layers
if getattr(args, "decoder_last_pipeline_num_layers", None) is not None:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,8 @@ def _streaming_quantized():

def _process_conversion_tasks(vanilla_conversion_tasks, new_weight_dict):
def _handle_one(task):
if task is None:
return None
if task.param_weight is None:
return task

Expand Down
25 changes: 19 additions & 6 deletions slime/utils/mask_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,11 @@ def get_system_message_length(self) -> tuple[int, int]:
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
test_messages,
add_special_tokens=False,
tokenize=True,
add_generation_prompt=True,
return_dict=False,
)
) - len(chat_template_token_ids)

Expand All @@ -53,9 +57,11 @@ def gen_multi_turn_loss_mask_qwen(

for i, message in enumerate(messages):
if i == 0:
message_ids = self.tokenizer.apply_chat_template([message], tokenize=True, tools=tools)
message_ids = self.tokenizer.apply_chat_template(
[message], tokenize=True, tools=tools, return_dict=False
)
else:
message_ids = self.tokenizer.apply_chat_template([message], tokenize=True)
message_ids = self.tokenizer.apply_chat_template([message], tokenize=True, return_dict=False)

if message["role"] != "system" and i > 0:
message_ids = message_ids[self.system_message_length :]
Expand All @@ -80,16 +86,23 @@ def gen_multi_turn_loss_mask_qwen3(
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)
prefix_token_ids = self.tokenizer.apply_chat_template([prefix_message], tokenize=True, return_dict=False)

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, prefix_message],
tokenize=True,
tools=tools,
return_dict=False,
)
message_ids = tailed_message_ids[: -len(prefix_token_ids)]
else:
prefixed_message_ids = self.tokenizer.apply_chat_template([prefix_message, message], tokenize=True)
prefixed_message_ids = self.tokenizer.apply_chat_template(
[prefix_message, message],
tokenize=True,
return_dict=False,
)
message_ids = prefixed_message_ids[len(prefix_token_ids) :]

if message["role"] != "system" and i > 0:
Expand Down
1 change: 1 addition & 0 deletions slime_plugins/megatron_bridge/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
import slime_plugins.megatron_bridge.glm4v_moe # noqa: F401 # register GLM-4.6V bridge
Loading