diff --git a/src/maxtext/checkpoint_conversion/utils/hf_model_configs.py b/src/maxtext/checkpoint_conversion/utils/hf_model_configs.py index 89abd56d4c..5dc3315d77 100644 --- a/src/maxtext/checkpoint_conversion/utils/hf_model_configs.py +++ b/src/maxtext/checkpoint_conversion/utils/hf_model_configs.py @@ -1883,6 +1883,56 @@ def __init__(self, **kwargs): } qwen3_vl_30b_a3b_config = PTConfig(**qwen3_vl_30b_a3b_dict) +cosmos3_nano_reasoner_dict = { + "architectures": ["Cosmos3OmniForConditionalGeneration"], + "model_type": "cosmos3_omni", + "text_config": { + "attention_bias": False, + "attention_dropout": 0.0, + "bos_token_id": 151643, + "dtype": "bfloat16", + "eos_token_id": 151645, + "head_dim": 128, + "hidden_act": "silu", + "hidden_size": 4096, + "initializer_range": 0.02, + "intermediate_size": 12288, + "max_position_embeddings": 262144, + "model_type": "qwen3_vl_text", + "num_attention_heads": 32, + "num_hidden_layers": 36, + "num_key_value_heads": 8, + "pad_token_id": None, + "rms_norm_eps": 1e-06, + "rope_parameters": { + "mrope_interleaved": True, + "mrope_section": [24, 20, 20], + "rope_theta": 5000000, + "rope_type": "default", + }, + "tie_word_embeddings": True, + "use_cache": True, + "vocab_size": 151936, + }, + "vision_config": { + "deepstack_visual_indexes": [8, 16, 24], + "depth": 27, + "hidden_act": "gelu_pytorch_tanh", + "hidden_size": 1152, + "in_channels": 3, + "initializer_range": 0.02, + "intermediate_size": 4304, + "model_type": "qwen3_vl_vision", + "num_heads": 16, + "num_position_embeddings": 2304, + "out_hidden_size": 4096, + "patch_size": 16, + "spatial_merge_size": 2, + "temporal_patch_size": 2, + }, +} +cosmos3_nano_reasoner_config = PTConfig(**cosmos3_nano_reasoner_dict) + # {maxtext model name: hf model config} HF_MODEL_CONFIGS = { @@ -1913,6 +1963,7 @@ def __init__(self, **kwargs): "qwen3-vl-2b": qwen3_vl_2b_config, "qwen3-vl-4b": qwen3_vl_4b_config, "qwen3-vl-30b-a3b": qwen3_vl_30b_a3b_config, + "cosmos3-nano-reasoner": cosmos3_nano_reasoner_config, "llama3.1-8b": llama31_8b_config, "llama3.1-8b-Instruct": llama31_8b_config, "llama3.1-70b": llama31_70b_config, diff --git a/src/maxtext/checkpoint_conversion/utils/param_mapping.py b/src/maxtext/checkpoint_conversion/utils/param_mapping.py index 94e96173f1..39a9d58236 100644 --- a/src/maxtext/checkpoint_conversion/utils/param_mapping.py +++ b/src/maxtext/checkpoint_conversion/utils/param_mapping.py @@ -3695,69 +3695,117 @@ def replace_prefix(val): mapping[composite_key] = f"model.language_model.layers.{i}.mlp.experts.gate_up_proj" mapping[key_wo] = f"model.language_model.layers.{i}.mlp.experts.down_proj" - vision_config = config["vision_config"] - n_vision_layers = vision_config["depth"] + if maxtext_config.use_multimodal: + vision_config = config["vision_config"] + n_vision_layers = vision_config["depth"] - mapping["params-vision_encoder-Qwen3VLVisionEncoder_0-patch_embed-proj-kernel"] = "model.visual.patch_embed.proj.weight" - mapping["params-vision_encoder-Qwen3VLVisionEncoder_0-patch_embed-proj-bias"] = "model.visual.patch_embed.proj.bias" + mapping["params-vision_encoder-Qwen3VLVisionEncoder_0-patch_embed-proj-kernel"] = ( + "model.visual.patch_embed.proj.weight" + ) + mapping["params-vision_encoder-Qwen3VLVisionEncoder_0-patch_embed-proj-bias"] = "model.visual.patch_embed.proj.bias" - mapping["params-vision_encoder-Qwen3VLVisionEncoder_0-pos_embed_interpolate-pos_embed"] = ( - "model.visual.pos_embed.weight" - ) + mapping["params-vision_encoder-Qwen3VLVisionEncoder_0-pos_embed_interpolate-pos_embed"] = ( + "model.visual.pos_embed.weight" + ) - for i in range(n_vision_layers): - prefix = f"params-vision_encoder-Qwen3VLVisionEncoder_0-blocks_{i}" - hf_prefix = f"model.visual.blocks.{i}" + for i in range(n_vision_layers): + prefix = f"params-vision_encoder-Qwen3VLVisionEncoder_0-blocks_{i}" + hf_prefix = f"model.visual.blocks.{i}" - mapping[f"{prefix}-ln1-scale"] = f"{hf_prefix}.norm1.weight" - mapping[f"{prefix}-ln1-bias"] = f"{hf_prefix}.norm1.bias" - mapping[f"{prefix}-ln2-scale"] = f"{hf_prefix}.norm2.weight" - mapping[f"{prefix}-ln2-bias"] = f"{hf_prefix}.norm2.bias" + mapping[f"{prefix}-ln1-scale"] = f"{hf_prefix}.norm1.weight" + mapping[f"{prefix}-ln1-bias"] = f"{hf_prefix}.norm1.bias" + mapping[f"{prefix}-ln2-scale"] = f"{hf_prefix}.norm2.weight" + mapping[f"{prefix}-ln2-bias"] = f"{hf_prefix}.norm2.bias" - mapping[ - ( - f"{prefix}-attn-attn-query-kernel", - f"{prefix}-attn-attn-key-kernel", - f"{prefix}-attn-attn-value-kernel", - ) - ] = f"{hf_prefix}.attn.qkv.weight" - mapping[ - ( - f"{prefix}-attn-attn-query-bias", - f"{prefix}-attn-attn-key-bias", - f"{prefix}-attn-attn-value-bias", - ) - ] = f"{hf_prefix}.attn.qkv.bias" - mapping[f"{prefix}-attn-attn-out-kernel"] = f"{hf_prefix}.attn.proj.weight" - mapping[f"{prefix}-attn-attn-out-bias"] = f"{hf_prefix}.attn.proj.bias" + mapping[ + ( + f"{prefix}-attn-attn-query-kernel", + f"{prefix}-attn-attn-key-kernel", + f"{prefix}-attn-attn-value-kernel", + ) + ] = f"{hf_prefix}.attn.qkv.weight" + mapping[ + ( + f"{prefix}-attn-attn-query-bias", + f"{prefix}-attn-attn-key-bias", + f"{prefix}-attn-attn-value-bias", + ) + ] = f"{hf_prefix}.attn.qkv.bias" + mapping[f"{prefix}-attn-attn-out-kernel"] = f"{hf_prefix}.attn.proj.weight" + mapping[f"{prefix}-attn-attn-out-bias"] = f"{hf_prefix}.attn.proj.bias" - mapping[f"{prefix}-mlp-kernel"] = f"{hf_prefix}.mlp.linear_fc1.weight" - mapping[f"{prefix}-mlp-bias"] = f"{hf_prefix}.mlp.linear_fc1.bias" - mapping[f"{prefix}-mlp_out-kernel"] = f"{hf_prefix}.mlp.linear_fc2.weight" - mapping[f"{prefix}-mlp_out-bias"] = f"{hf_prefix}.mlp.linear_fc2.bias" + mapping[f"{prefix}-mlp-kernel"] = f"{hf_prefix}.mlp.linear_fc1.weight" + mapping[f"{prefix}-mlp-bias"] = f"{hf_prefix}.mlp.linear_fc1.bias" + mapping[f"{prefix}-mlp_out-kernel"] = f"{hf_prefix}.mlp.linear_fc2.weight" + mapping[f"{prefix}-mlp_out-bias"] = f"{hf_prefix}.mlp.linear_fc2.bias" - deepstack_indexes = vision_config.get("deepstack_visual_indexes", [5, 11, 17]) - for merger_idx, _ in enumerate(deepstack_indexes): - prefix = f"params-vision_encoder-Qwen3VLVisionEncoder_0-merger_{merger_idx}" - hf_prefix = f"model.visual.deepstack_merger_list.{merger_idx}" - - mapping[f"{prefix}-ln_q-scale"] = f"{hf_prefix}.norm.weight" - mapping[f"{prefix}-ln_q-bias"] = f"{hf_prefix}.norm.bias" - mapping[f"{prefix}-mlp_0-kernel"] = f"{hf_prefix}.linear_fc1.weight" - mapping[f"{prefix}-mlp_0-bias"] = f"{hf_prefix}.linear_fc1.bias" - mapping[f"{prefix}-mlp_2-kernel"] = f"{hf_prefix}.linear_fc2.weight" - mapping[f"{prefix}-mlp_2-bias"] = f"{hf_prefix}.linear_fc2.bias" - - mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-ln_q-scale"] = "model.visual.merger.norm.weight" - mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-ln_q-bias"] = "model.visual.merger.norm.bias" - mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-mlp_0-kernel"] = "model.visual.merger.linear_fc1.weight" - mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-mlp_0-bias"] = "model.visual.merger.linear_fc1.bias" - mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-mlp_2-kernel"] = "model.visual.merger.linear_fc2.weight" - mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-mlp_2-bias"] = "model.visual.merger.linear_fc2.bias" + deepstack_indexes = vision_config.get("deepstack_visual_indexes", [5, 11, 17]) + for merger_idx, _ in enumerate(deepstack_indexes): + prefix = f"params-vision_encoder-Qwen3VLVisionEncoder_0-merger_{merger_idx}" + hf_prefix = f"model.visual.deepstack_merger_list.{merger_idx}" + + mapping[f"{prefix}-ln_q-scale"] = f"{hf_prefix}.norm.weight" + mapping[f"{prefix}-ln_q-bias"] = f"{hf_prefix}.norm.bias" + mapping[f"{prefix}-mlp_0-kernel"] = f"{hf_prefix}.linear_fc1.weight" + mapping[f"{prefix}-mlp_0-bias"] = f"{hf_prefix}.linear_fc1.bias" + mapping[f"{prefix}-mlp_2-kernel"] = f"{hf_prefix}.linear_fc2.weight" + mapping[f"{prefix}-mlp_2-bias"] = f"{hf_prefix}.linear_fc2.bias" + + mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-ln_q-scale"] = "model.visual.merger.norm.weight" + mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-ln_q-bias"] = "model.visual.merger.norm.bias" + mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-mlp_0-kernel"] = ( + "model.visual.merger.linear_fc1.weight" + ) + mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-mlp_0-bias"] = "model.visual.merger.linear_fc1.bias" + mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-mlp_2-kernel"] = ( + "model.visual.merger.linear_fc2.weight" + ) + mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-mlp_2-bias"] = "model.visual.merger.linear_fc2.bias" + + return mapping + + +def COSMOS3_MAXTEXT_TO_HF_PARAM_MAPPING(config, maxtext_config, scan_layers=False): + """Returns mapping from MaxText to HuggingFace Cosmos3-Nano Reasoner weight paths.""" + # 1. Reuse QWEN3_VL mapping + qwen3_vl_mapping = QWEN3_VL_MAXTEXT_TO_HF_PARAM_MAPPING(config, maxtext_config, scan_layers) + + mapping = {} + + def translate_hf_key(val): + if isinstance(val, list): + return [translate_hf_key(v) for v in val] + elif isinstance(val, str): + # Strip prefixes used by Qwen3-VL in HF + if val.startswith("model.language_model."): + val = val[len("model.language_model.") :] + elif val.startswith("model.visual."): + val = val[len("model.visual.") :] + elif val.startswith("model."): + val = val[len("model.") :] + + # Apply Cosmos3 specific attention naming + val = val.replace("self_attn.q_proj", "self_attn.to_q") + val = val.replace("self_attn.k_proj", "self_attn.to_k") + val = val.replace("self_attn.v_proj", "self_attn.to_v") + val = val.replace("self_attn.o_proj", "self_attn.to_out") + val = val.replace("self_attn.q_norm", "self_attn.norm_q") + val = val.replace("self_attn.k_norm", "self_attn.norm_k") + return val + return val + + for key, value in qwen3_vl_mapping.items(): + mapping[key] = translate_hf_key(value) return mapping +def COSMOS3_MAXTEXT_TO_HF_PARAM_HOOK_FN(config, maxtext_config, scan_layers=False, saving_to_hf=False): + """Creates parameter transformation functions for Cosmos3-Nano Reasoner.""" + # Hooks operate on MaxText Parameter paths, which are identical to Qwen3-VL + return QWEN3_VL_MAXTEXT_TO_HF_PARAM_HOOK_FN(config, maxtext_config, scan_layers, saving_to_hf) + + def QWEN3_VL_MAXTEXT_TO_HF_PARAM_HOOK_FN(config, maxtext_config, scan_layers=False, saving_to_hf=False): """Creates parameter transformation functions for Qwen3-VL.""" mapping = {} @@ -3793,99 +3841,100 @@ def process_wi_0_wi_1_fused(input_tensor, target_shape=None): mapping[composite_key] = process_wi_0_wi_1_fused mapping[f"params-decoder-layers_{i}-moe_block-wo"] = None - vision_config = config["vision_config"] - n_vision_layers = vision_config["depth"] - hidden_size = vision_config["hidden_size"] + if maxtext_config.use_multimodal: + vision_config = config["vision_config"] + n_vision_layers = vision_config["depth"] + hidden_size = vision_config["hidden_size"] - def reshape_kernel_vision(input_tensor, target_shape): - """Reshape kernel for vision layers.""" - if saving_to_hf: - flipped_target_shape = np.flip(np.array(target_shape)) - return input_tensor.reshape(flipped_target_shape).T - else: - return input_tensor.T.reshape(target_shape) + def reshape_kernel_vision(input_tensor, target_shape): + """Reshape kernel for vision layers.""" + if saving_to_hf: + flipped_target_shape = np.flip(np.array(target_shape)) + return input_tensor.reshape(flipped_target_shape).T + else: + return input_tensor.T.reshape(target_shape) - def reshape_conv3d_patch_embed(input_tensor, target_shape): - """Reshape 3D conv patch embedding weight.""" - if saving_to_hf: - return input_tensor.transpose(4, 3, 0, 1, 2) - else: - return input_tensor.transpose(2, 3, 4, 1, 0) + def reshape_conv3d_patch_embed(input_tensor, target_shape): + """Reshape 3D conv patch embedding weight.""" + if saving_to_hf: + return input_tensor.transpose(4, 3, 0, 1, 2) + else: + return input_tensor.transpose(2, 3, 4, 1, 0) - def process_qkv_vision(input_tensor, target_shape=None): - """Handles composite_mt_key: maxtext (query, key, value) <-> hf (qkv).""" - if saving_to_hf: - q, k, v = input_tensor - q_hf = q.reshape(hidden_size, hidden_size).T - k_hf = k.reshape(hidden_size, hidden_size).T - v_hf = v.reshape(hidden_size, hidden_size).T - return np.concatenate([q_hf, k_hf, v_hf], axis=0) - else: - q_hf = input_tensor[:hidden_size, :] - k_hf = input_tensor[hidden_size : 2 * hidden_size, :] - v_hf = input_tensor[2 * hidden_size :, :] - q_mt = q_hf.T.reshape(target_shape[0]) # pyrefly: ignore[unsupported-operation] - k_mt = k_hf.T.reshape(target_shape[1]) # pyrefly: ignore[unsupported-operation] - v_mt = v_hf.T.reshape(target_shape[2]) # pyrefly: ignore[unsupported-operation] - return np.stack([q_mt, k_mt, v_mt], axis=-1) - - def process_qkv_bias_vision(input_tensor, target_shape=None): - """Handles composite_mt_key: maxtext (query_bias, key_bias, value_bias) <-> hf (qkv_bias).""" - if saving_to_hf: - qb, kb, vb = input_tensor - qb_hf = qb.reshape(hidden_size) - kb_hf = kb.reshape(hidden_size) - vb_hf = vb.reshape(hidden_size) - return np.concatenate([qb_hf, kb_hf, vb_hf], axis=0) - else: - qb_hf = input_tensor[:hidden_size] - kb_hf = input_tensor[hidden_size : 2 * hidden_size] - vb_hf = input_tensor[2 * hidden_size :] - qb_mt = qb_hf.reshape(target_shape[0]) # pyrefly: ignore[unsupported-operation] - kb_mt = kb_hf.reshape(target_shape[1]) # pyrefly: ignore[unsupported-operation] - vb_mt = vb_hf.reshape(target_shape[2]) # pyrefly: ignore[unsupported-operation] - return np.stack([qb_mt, kb_mt, vb_mt], axis=-1) + def process_qkv_vision(input_tensor, target_shape=None): + """Handles composite_mt_key: maxtext (query, key, value) <-> hf (qkv).""" + if saving_to_hf: + q, k, v = input_tensor + q_hf = q.reshape(hidden_size, hidden_size).T + k_hf = k.reshape(hidden_size, hidden_size).T + v_hf = v.reshape(hidden_size, hidden_size).T + return np.concatenate([q_hf, k_hf, v_hf], axis=0) + else: + q_hf = input_tensor[:hidden_size, :] + k_hf = input_tensor[hidden_size : 2 * hidden_size, :] + v_hf = input_tensor[2 * hidden_size :, :] + q_mt = q_hf.T.reshape(target_shape[0]) # pyrefly: ignore[unsupported-operation] + k_mt = k_hf.T.reshape(target_shape[1]) # pyrefly: ignore[unsupported-operation] + v_mt = v_hf.T.reshape(target_shape[2]) # pyrefly: ignore[unsupported-operation] + return np.stack([q_mt, k_mt, v_mt], axis=-1) + + def process_qkv_bias_vision(input_tensor, target_shape=None): + """Handles composite_mt_key: maxtext (query_bias, key_bias, value_bias) <-> hf (qkv_bias).""" + if saving_to_hf: + qb, kb, vb = input_tensor + qb_hf = qb.reshape(hidden_size) + kb_hf = kb.reshape(hidden_size) + vb_hf = vb.reshape(hidden_size) + return np.concatenate([qb_hf, kb_hf, vb_hf], axis=0) + else: + qb_hf = input_tensor[:hidden_size] + kb_hf = input_tensor[hidden_size : 2 * hidden_size] + vb_hf = input_tensor[2 * hidden_size :] + qb_mt = qb_hf.reshape(target_shape[0]) # pyrefly: ignore[unsupported-operation] + kb_mt = kb_hf.reshape(target_shape[1]) # pyrefly: ignore[unsupported-operation] + vb_mt = vb_hf.reshape(target_shape[2]) # pyrefly: ignore[unsupported-operation] + return np.stack([qb_mt, kb_mt, vb_mt], axis=-1) - def reshape_vision_attn_out(input_tensor, target_shape): - """Reshape vision attention output projection.""" - if saving_to_hf: - return input_tensor.reshape(hidden_size, hidden_size).T - else: - return input_tensor.T.reshape(target_shape) + def reshape_vision_attn_out(input_tensor, target_shape): + """Reshape vision attention output projection.""" + if saving_to_hf: + return input_tensor.reshape(hidden_size, hidden_size).T + else: + return input_tensor.T.reshape(target_shape) - mapping["params-vision_encoder-Qwen3VLVisionEncoder_0-patch_embed-proj-kernel"] = reshape_conv3d_patch_embed + mapping["params-vision_encoder-Qwen3VLVisionEncoder_0-patch_embed-proj-kernel"] = reshape_conv3d_patch_embed - for i in range(n_vision_layers): - prefix = f"params-vision_encoder-Qwen3VLVisionEncoder_0-blocks_{i}" + for i in range(n_vision_layers): + prefix = f"params-vision_encoder-Qwen3VLVisionEncoder_0-blocks_{i}" - mapping[ - ( - f"{prefix}-attn-attn-query-kernel", - f"{prefix}-attn-attn-key-kernel", - f"{prefix}-attn-attn-value-kernel", - ) - ] = process_qkv_vision - mapping[ - ( - f"{prefix}-attn-attn-query-bias", - f"{prefix}-attn-attn-key-bias", - f"{prefix}-attn-attn-value-bias", - ) - ] = process_qkv_bias_vision + mapping[ + ( + f"{prefix}-attn-attn-query-kernel", + f"{prefix}-attn-attn-key-kernel", + f"{prefix}-attn-attn-value-kernel", + ) + ] = process_qkv_vision + mapping[ + ( + f"{prefix}-attn-attn-query-bias", + f"{prefix}-attn-attn-key-bias", + f"{prefix}-attn-attn-value-bias", + ) + ] = process_qkv_bias_vision - mapping[f"{prefix}-attn-attn-out-kernel"] = reshape_vision_attn_out + mapping[f"{prefix}-attn-attn-out-kernel"] = reshape_vision_attn_out - mapping[f"{prefix}-mlp-kernel"] = reshape_kernel_vision - mapping[f"{prefix}-mlp_out-kernel"] = reshape_kernel_vision + mapping[f"{prefix}-mlp-kernel"] = reshape_kernel_vision + mapping[f"{prefix}-mlp_out-kernel"] = reshape_kernel_vision - deepstack_indexes = vision_config.get("deepstack_visual_indexes", [5, 11, 17]) - for merger_idx, _ in enumerate(deepstack_indexes): - prefix = f"params-vision_encoder-Qwen3VLVisionEncoder_0-merger_{merger_idx}" - mapping[f"{prefix}-mlp_0-kernel"] = reshape_kernel_vision - mapping[f"{prefix}-mlp_2-kernel"] = reshape_kernel_vision + deepstack_indexes = vision_config.get("deepstack_visual_indexes", [5, 11, 17]) + for merger_idx, _ in enumerate(deepstack_indexes): + prefix = f"params-vision_encoder-Qwen3VLVisionEncoder_0-merger_{merger_idx}" + mapping[f"{prefix}-mlp_0-kernel"] = reshape_kernel_vision + mapping[f"{prefix}-mlp_2-kernel"] = reshape_kernel_vision - mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-mlp_0-kernel"] = reshape_kernel_vision - mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-mlp_2-kernel"] = reshape_kernel_vision + mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-mlp_0-kernel"] = reshape_kernel_vision + mapping["params-vision_encoder-Qwen3VLVisionProjector_0-merger-mlp_2-kernel"] = reshape_kernel_vision return mapping @@ -4245,6 +4294,7 @@ def mhc_concat_scale(input_tensors, target_shape=None): "qwen3-vl-2b": QWEN3_VL_MAXTEXT_TO_HF_PARAM_MAPPING, "qwen3-vl-4b": QWEN3_VL_MAXTEXT_TO_HF_PARAM_MAPPING, "qwen3-vl-30b-a3b": QWEN3_VL_MAXTEXT_TO_HF_PARAM_MAPPING, + "cosmos3-nano-reasoner": COSMOS3_MAXTEXT_TO_HF_PARAM_MAPPING, "llama3.1-8b": LLAMA31_MAXTEXT_TO_HF_PARAM_MAPPING, "llama3.1-8b-Instruct": LLAMA31_MAXTEXT_TO_HF_PARAM_MAPPING, "llama3.1-70b": LLAMA31_MAXTEXT_TO_HF_PARAM_MAPPING, @@ -4299,6 +4349,7 @@ def mhc_concat_scale(input_tensors, target_shape=None): "qwen3-vl-2b": QWEN3_VL_MAXTEXT_TO_HF_PARAM_HOOK_FN, "qwen3-vl-4b": QWEN3_VL_MAXTEXT_TO_HF_PARAM_HOOK_FN, "qwen3-vl-30b-a3b": QWEN3_VL_MAXTEXT_TO_HF_PARAM_HOOK_FN, + "cosmos3-nano-reasoner": COSMOS3_MAXTEXT_TO_HF_PARAM_HOOK_FN, "llama3.1-8b": LLAMA31_MAXTEXT_TO_HF_PARAM_HOOK_FN, "llama3.1-8b-Instruct": LLAMA31_MAXTEXT_TO_HF_PARAM_HOOK_FN, "llama3.1-70b": LLAMA31_MAXTEXT_TO_HF_PARAM_HOOK_FN, diff --git a/src/maxtext/checkpoint_conversion/utils/utils.py b/src/maxtext/checkpoint_conversion/utils/utils.py index 449b2d2194..479fa6fbf1 100644 --- a/src/maxtext/checkpoint_conversion/utils/utils.py +++ b/src/maxtext/checkpoint_conversion/utils/utils.py @@ -1113,7 +1113,7 @@ def load_hf_dict_from_safetensors(model_id_or_path, token, revision, framework=" revision=revision, ) # load safetensors - ckpt_paths = sorted(pathlib.Path(local_path).glob("[!.]*.safetensors")) + ckpt_paths = sorted(pathlib.Path(local_path).rglob("[!.]*.safetensors")) hf_state_dict = {} max_logging.log(f"Loading {len(ckpt_paths)} checkpoints") for ckpt_path in tqdm(ckpt_paths, total=len(ckpt_paths)): diff --git a/src/maxtext/configs/models/cosmos3-nano-reasoner.yml b/src/maxtext/configs/models/cosmos3-nano-reasoner.yml new file mode 100644 index 0000000000..dc740bd9d5 --- /dev/null +++ b/src/maxtext/configs/models/cosmos3-nano-reasoner.yml @@ -0,0 +1,42 @@ +# Model config for NVIDIA Cosmos3-Nano Reasoner (Text-only mode) + +# Core Architectural Parameters +decoder_block: "qwen3" +base_emb_dim: 4096 +base_mlp_dim: 12288 +base_num_query_heads: 32 +base_num_kv_heads: 8 +base_num_decoder_layers: 36 +head_dim: 128 +mlp_activations: ["silu", "linear"] +vocab_size: 151936 +normalization_layer_epsilon: 1.0e-6 +use_qk_norm: true +logits_via_embedding: false + +# RoPE Settings +rope_max_timescale: 5000000 + +# General Model Settings +enable_dropout: false + +# Multimodal Settings +use_multimodal: true +use_mrope: true +mrope_section: [24, 20, 20] + +# Vision Encoder Configuration +image_size_for_vit: 768 +hidden_size_for_vit: 1152 +intermediate_size_for_vit: 4304 +num_attention_heads_for_vit: 16 +num_hidden_layers_for_vit: 27 +num_channels_for_vit: 3 +patch_size_for_vit: 16 +temporal_patch_size_for_vit: 2 +spatial_merge_size_for_vit: 2 +out_hidden_size_for_vit: 4096 +num_position_embeddings_for_vit: 2304 +deepstack_visual_indexes_for_vit: [8, 16, 24] + +vision_encoder_block: "qwen3_vl" diff --git a/src/maxtext/configs/types.py b/src/maxtext/configs/types.py index 03e5db969b..ac94f1fe24 100644 --- a/src/maxtext/configs/types.py +++ b/src/maxtext/configs/types.py @@ -274,6 +274,7 @@ class ProfilerType(str, Enum): "qwen3-vl-2b", "qwen3-vl-4b", "qwen3-vl-30b-a3b", + "cosmos3-nano-reasoner", "qwen3-next-80b-a3b", "qwen3-omni-30b-a3b", "qwen3-custom-30b-a3b", @@ -3851,6 +3852,7 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de "qwen3-vl-30b-a3b", "qwen3.5-35b-a3b", "qwen3.5-397b-a17b", + "cosmos3-nano-reasoner", ) if self.model_name not in valid_mm_models and self.model_name != "default": raise ValueError(f"Multimodal is only supported for {valid_mm_models}, not {self.model_name}") diff --git a/src/maxtext/layers/attentions.py b/src/maxtext/layers/attentions.py index f24d2b48d7..7590c13b2b 100644 --- a/src/maxtext/layers/attentions.py +++ b/src/maxtext/layers/attentions.py @@ -894,7 +894,7 @@ def init_rotary_embedding(self): rope_type = self.rope_type rope_use_scale = self.config.rope_use_scale if self.is_vision: - if self.config.model_name.startswith("qwen3"): + if self.config.model_name.startswith("qwen3") or self.config.model_name.startswith("cosmos3"): rotary_embedding = Qwen3OmniMoeVisionRotaryEmbedding( hidden_size=self.config.hidden_size_for_vit, num_attention_heads=self.config.num_attention_heads_for_vit, diff --git a/src/maxtext/layers/nnx_decoders.py b/src/maxtext/layers/nnx_decoders.py index 1f300af067..14ab8b4fe0 100644 --- a/src/maxtext/layers/nnx_decoders.py +++ b/src/maxtext/layers/nnx_decoders.py @@ -1361,6 +1361,7 @@ def _apply_embedding( "qwen3-vl-30b-a3b", "qwen3.5-35b-a3b", "qwen3.5-397b-a17b", + "cosmos3-nano-reasoner", }: y = mm_utils.merge_mm_embeddings( text_embeddings=y, @@ -1379,6 +1380,7 @@ def _apply_embedding( "qwen3-vl-30b-a3b", "qwen3.5-35b-a3b", "qwen3.5-397b-a17b", + "cosmos3-nano-reasoner", }: y = mm_utils.merge_mm_embeddings( text_embeddings=y, diff --git a/src/maxtext/multimodal/processor.py b/src/maxtext/multimodal/processor.py index 025d8ab40b..bd8619c275 100644 --- a/src/maxtext/multimodal/processor.py +++ b/src/maxtext/multimodal/processor.py @@ -37,6 +37,8 @@ "qwen3-vl-30b-a3b": ("qwen3_vl", "qwen3_moe"), "qwen3.5-35b-a3b": ("qwen3_5", "qwen3_5"), "qwen3.5-397b-a17b": ("qwen3_5", "qwen3_5"), + # Cosmos + "cosmos3-nano-reasoner": ("qwen3_vl", "qwen3"), } @@ -100,7 +102,7 @@ def preprocess_mm_data(config): images = [mm_utils.load_image_from_path(p) for p in config.image_path.split(",")] processor_outputs = preprocess_mm_data_llama4(images) - elif vision_block in ["qwen3_omni", "qwen3_vl", "qwen3_5"]: + elif vision_block in ["qwen3_omni", "qwen3_vl", "qwen3_5", "cosmos3-nano-reasoner"]: from maxtext.multimodal.processor_qwen3_omni import preprocess_mm_data_qwen3_omni # pylint: disable=import-outside-toplevel processor_outputs = preprocess_mm_data_qwen3_omni(config) @@ -127,7 +129,7 @@ def preprocess_image_for_training(image, config): from maxtext.multimodal.processor_llama4 import preprocess_mm_data_llama4 # pylint: disable=import-outside-toplevel return preprocess_mm_data_llama4(image) - elif vision_block in ["qwen3_omni", "qwen3_vl", "qwen3_5"]: + elif vision_block in ["qwen3_omni", "qwen3_vl", "qwen3_5", "cosmos3-nano-reasoner"]: from maxtext.multimodal.processor_qwen3_omni import preprocess_mm_data_qwen3_omni_for_training # pylint: disable=import-outside-toplevel return preprocess_mm_data_qwen3_omni_for_training(image, config) @@ -151,7 +153,7 @@ def get_image_offsets(config, processor_output: mm_utils.PreprocessorOutput | No from maxtext.multimodal.processor_llama4 import get_image_offsets_llama4 # pylint: disable=import-outside-toplevel return get_image_offsets_llama4(processor_output) - elif vision_block in ["qwen3_omni", "qwen3_vl", "qwen3_5"]: + elif vision_block in ["qwen3_omni", "qwen3_vl", "qwen3_5", "cosmos3-nano-reasoner"]: from maxtext.multimodal.processor_qwen3_omni import get_mm_offsets_qwen3_omni # pylint: disable=import-outside-toplevel return get_mm_offsets_qwen3_omni(config, processor_output) @@ -210,7 +212,7 @@ def reformat_response(response, model_name): elif decoder_block in ["gemma4", "gemma4_small"]: formatted_response = f"{response}" return formatted_response - elif decoder_block in ["qwen3", "qwen3_moe", "qwen3_5"]: + elif decoder_block in ["qwen3", "qwen3_moe", "qwen3_5", "cosmos3-nano-reasoner"]: formatted_response = f"{response}<|im_end|>" return formatted_response else: @@ -236,7 +238,7 @@ def prepare_text_for_image_fusion(tokens, config, processor_output=None): from maxtext.multimodal.processor_llama4 import add_extra_tokens_for_images_llama4 # pylint: disable=import-outside-toplevel return add_extra_tokens_for_images_llama4(tokens, processor_output) # pyrefly: ignore[bad-argument-type] - elif vision_block in ["qwen3_omni", "qwen3_vl", "qwen3_5"]: + elif vision_block in ["qwen3_omni", "qwen3_vl", "qwen3_5", "cosmos3-nano-reasoner"]: from maxtext.multimodal.processor_qwen3_omni import add_extra_tokens_for_qwen3_omni # pylint: disable=import-outside-toplevel return add_extra_tokens_for_qwen3_omni(tokens, config, processor_output) @@ -260,7 +262,7 @@ def get_dummy_image_shape_for_init(model_name, batch_size=1, num_image_per_seque from maxtext.multimodal.processor_llama4 import get_dummy_image_shape_for_init_llama4 # pylint: disable=import-outside-toplevel image_shape = get_dummy_image_shape_for_init_llama4(batch_size, num_image_per_sequence) - elif vision_block in ["qwen3_omni", "qwen3_vl", "qwen3_5"]: + elif vision_block in ["qwen3_omni", "qwen3_vl", "qwen3_5", "cosmos3-nano-reasoner"]: from maxtext.multimodal.processor_qwen3_omni import get_dummy_image_shape_for_init_qwen3_omni # pylint: disable=import-outside-toplevel image_shape = get_dummy_image_shape_for_init_qwen3_omni(batch_size) @@ -308,7 +310,7 @@ def get_bidirectional_mask_vision(config, decoder_input_tokens, is_video: bool = from maxtext.multimodal.processor_llama4 import LLAMA4_PATCH_TOKEN # pylint: disable=import-outside-toplevel bidirectional_mask_vision = decoder_input_tokens == LLAMA4_PATCH_TOKEN - elif decoder_block in ["qwen3", "qwen3_moe", "qwen3_5"]: + elif decoder_block in ["qwen3", "qwen3_moe", "qwen3_5", "cosmos3-nano-reasoner"]: from maxtext.multimodal.processor_qwen3_omni import QwenTokens # pylint: disable=import-outside-toplevel tokens = QwenTokens(config) diff --git a/src/maxtext/utils/globals.py b/src/maxtext/utils/globals.py index 30f6e65124..0c73e6f020 100644 --- a/src/maxtext/utils/globals.py +++ b/src/maxtext/utils/globals.py @@ -91,6 +91,7 @@ "olmo3-7b": "allenai/Olmo-3-7B-Instruct", "olmo3-7b-pt": "allenai/Olmo-3-1025-7B", "olmo3-32b": "allenai/Olmo-3-32B-Think", + "cosmos3-nano-reasoner": "nvidia/Cosmos3-Nano", # "default" is not HF model, but adding to to avoid confusing warning about tokenizer_path "default": os.path.join(MAXTEXT_ASSETS_ROOT, "tokenizers/tokenizer.llama2"), } diff --git a/tests/assets/test_image_reasoning.jpg b/tests/assets/test_image_reasoning.jpg new file mode 100644 index 0000000000..113afdb1b5 Binary files /dev/null and b/tests/assets/test_image_reasoning.jpg differ diff --git a/tests/unit/param_mapping_test.py b/tests/unit/param_mapping_test.py index cea6485817..bc2583fa2d 100644 --- a/tests/unit/param_mapping_test.py +++ b/tests/unit/param_mapping_test.py @@ -111,6 +111,34 @@ def test_qwen3_next_mapping_scanned(self): mapping = param_mapping.QWEN3_NEXT_MAXTEXT_TO_HF_PARAM_MAPPING(config, maxtext_config, scan_layers=True) self.assertIn("params-decoder-layers-layer_0-input_layernorm-scale", mapping) + def test_cosmos3_text_mapping(self): + config = { + "text_config": {"num_hidden_layers": 2, "hidden_size": 256}, + } + maxtext_config = mock.Mock() + maxtext_config.use_multimodal = False + mapping = param_mapping.COSMOS3_MAXTEXT_TO_HF_PARAM_MAPPING(config, maxtext_config, scan_layers=False) + self.assertIn("params-token_embedder-embedding", mapping) + self.assertEqual(mapping["params-decoder-layers_0-self_attention-query-kernel"], "layers.0.self_attn.to_q.weight") + + def test_cosmos3_text_mapping_scanned(self): + config = { + "text_config": {"num_hidden_layers": 4, "hidden_size": 256}, + } + maxtext_config = mock.Mock() + maxtext_config.use_multimodal = False + mapping = param_mapping.COSMOS3_MAXTEXT_TO_HF_PARAM_MAPPING(config, maxtext_config, scan_layers=True) + self.assertIn("params-decoder-layers-self_attention-query-kernel", mapping) + self.assertEqual( + mapping["params-decoder-layers-self_attention-query-kernel"], + [ + "layers.0.self_attn.to_q.weight", + "layers.1.self_attn.to_q.weight", + "layers.2.self_attn.to_q.weight", + "layers.3.self_attn.to_q.weight", + ], + ) + def test_deepseek_mapping(self): config = { "num_hidden_layers": 4,