Skip to content

Commit 998e68f

Browse files
Merge pull request AI-Hypercomputer#4485 from AI-Hypercomputer:multi-dir-ckpt-load
PiperOrigin-RevId: 953422205
2 parents a91e533 + 542fb6d commit 998e68f

4 files changed

Lines changed: 910 additions & 0 deletions

File tree

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,39 @@
1+
# model config for omni-gemma3-qwen3
2+
# Combines LLM backbone from Qwen 3 4B with Vision features from Gemma 3 4B
3+
4+
# NOTE: model_name is set to "gemma3-4b" so encoders.py loads Gemma 3's Vision Tower without modification.
5+
# Meanwhile, decoders.py can build Qwen 3 LLM decoder directly from reading `decoder_block: "qwen3"` below.
6+
model_name: "gemma3-4b"
7+
use_multimodal: true
8+
9+
# Multimodal config for vision model (from gemma3-4b.yml)
10+
image_size_for_vit: 896
11+
num_channels_for_vit: 3
12+
patch_size_for_vit: 14
13+
conv_stride_for_vit: 14
14+
hidden_size_for_vit: 1152
15+
intermediate_size_for_vit: 4304
16+
num_hidden_layers_for_vit: 27
17+
num_attention_heads_for_vit: 16
18+
image_placeholder: <start_of_image>
19+
20+
# LLM backbone config (from qwen3-4b.yml)
21+
base_emb_dim: 2560
22+
base_num_query_heads: 32
23+
base_num_kv_heads: 8
24+
base_mlp_dim: 9728
25+
base_num_decoder_layers: 36
26+
head_dim: 128
27+
mlp_activations: ["silu", "linear"]
28+
vocab_size: 151936
29+
decoder_block: "qwen3"
30+
normalization_layer_epsilon: 1.0e-6
31+
rope_max_timescale: 1000000
32+
use_qk_norm: true
33+
logits_via_embedding: true
34+
normalize_embedding_logits: false
35+
enable_dropout: false
36+
tokenizer_type: "huggingface"
37+
38+
# Ensure no pretrained weights or checkpoints are loaded initially
39+
load_parameters_path: ""
Lines changed: 115 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,115 @@
1+
#!/bin/bash
2+
# Step 1 of 5 maxtext multimodal alignment proof of concept project.
3+
#
4+
# This file:
5+
# 1. downloads and converts a vision-language model (e.g. gemma3-4b) from Hugging Face to MaxText format
6+
# 2. downloads and converts an text-only model (e.g. qwen3-4b) from Hugging Face to MaxText format
7+
# 3. stitches the vision component and LLM checkpoints into a single omni checkpoint
8+
# 4. saves the stitched checkpoints to an output directory
9+
#
10+
# Example usage:
11+
# HF_TOKEN="your_token" BASE_OUTPUT_DIRECTORY="gs://your_bucket/omni_checkpoints" ./prepare_checkpoint.sh
12+
13+
set -e
14+
15+
# Hugging Face Token & Login
16+
HF_TOKEN="${HF_TOKEN:-}"
17+
18+
if [ -z "${BASE_OUTPUT_DIRECTORY}" ]; then
19+
echo "Error: BASE_OUTPUT_DIRECTORY is not set. Please set it as an environment variable."
20+
echo "Example: export BASE_OUTPUT_DIRECTORY=\"gs://your_bucket/omni_checkpoints\""
21+
exit 1
22+
fi
23+
BASE_OUTPUT_DIRECTORY="${BASE_OUTPUT_DIRECTORY%/}"
24+
25+
if [ -n "$HF_TOKEN" ]; then
26+
if command -v hf &> /dev/null; then
27+
echo "Logging into Hugging Face using hf..."
28+
hf auth login --token "$HF_TOKEN"
29+
elif command -v huggingface-cli &> /dev/null; then
30+
echo "Logging into Hugging Face using huggingface-cli..."
31+
huggingface-cli login --token "$HF_TOKEN"
32+
else
33+
echo "Neither hf nor huggingface-cli found. Skipping Hugging Face login."
34+
fi
35+
fi
36+
37+
# Configuration & Paths
38+
VISION_MAXTEXT_MODEL="gemma3-4b"
39+
VISION_HF_REPO="google/gemma-3-4b-it"
40+
41+
LLM_MAXTEXT_MODEL="qwen3-4b"
42+
LLM_HF_REPO="Qwen/Qwen3-4B"
43+
44+
# Automatically find maxtext package directory
45+
MAXTEXT_PKG_DIR=$(python3 -c "import os, maxtext; print(os.path.dirname(maxtext.__file__))")
46+
OMNI_CONFIG_PATH="${MAXTEXT_PKG_DIR}/experimental/omni_poc/omni-gemma3-qwen3.yml"
47+
48+
VISION_CKPT_DIR="${BASE_OUTPUT_DIRECTORY}/${VISION_MAXTEXT_MODEL}_converted"
49+
LLM_CKPT_DIR="${BASE_OUTPUT_DIRECTORY}/${LLM_MAXTEXT_MODEL}_converted"
50+
STITCHED_CKPT_DIR="${BASE_OUTPUT_DIRECTORY}/omni_stitched_${VISION_MAXTEXT_MODEL}_${LLM_MAXTEXT_MODEL}"
51+
52+
VISION_ITEMS_PATH="${VISION_CKPT_DIR}/0/items"
53+
LLM_ITEMS_PATH="${LLM_CKPT_DIR}/0/items"
54+
STITCHED_ITEMS_PATH="${STITCHED_CKPT_DIR}/0/items"
55+
56+
echo "Base Output Directory: ${BASE_OUTPUT_DIRECTORY}"
57+
echo "Vision Converted Path: ${VISION_ITEMS_PATH}"
58+
echo "LLM Converted Path: ${LLM_ITEMS_PATH}"
59+
echo "Stitched Target Path: ${STITCHED_ITEMS_PATH}"
60+
echo ""
61+
62+
export JAX_PLATFORMS=cpu
63+
64+
# Helper to check if local/GCS paths exist using python etils (same as python script)
65+
path_exists() {
66+
python3 -c "from etils import epath; import sys; sys.exit(0 if epath.Path(sys.argv[1]).exists() else 1)" "$1"
67+
}
68+
69+
# Step 1: Download & Convert Vision Model from Hugging Face -> MaxText
70+
echo "============================================================"
71+
if ! path_exists "$VISION_ITEMS_PATH"; then
72+
echo "Converting Vision Model (${VISION_MAXTEXT_MODEL}) from Hugging Face (${VISION_HF_REPO})..."
73+
python3 -m maxtext.checkpoint_conversion.to_maxtext \
74+
"${MAXTEXT_PKG_DIR}/configs/base.yml" \
75+
"model_name=${VISION_MAXTEXT_MODEL}" \
76+
"base_output_directory=${VISION_CKPT_DIR}" \
77+
"use_multimodal=True" \
78+
"scan_layers=True" \
79+
"skip_jax_distributed_system=True" \
80+
"--eager_load_method=transformers" \
81+
"--lazy_load_tensors=False" \
82+
"log_config=False"
83+
echo "Vision checkpoint conversion successful!"
84+
echo ""
85+
else
86+
echo "Step 1: Vision checkpoint already exists at ${VISION_ITEMS_PATH}"
87+
fi
88+
89+
# Step 2: Download & Convert Language Model from Hugging Face -> MaxText
90+
echo "============================================================"
91+
if ! path_exists "$LLM_ITEMS_PATH"; then
92+
echo "Converting Language Model (${LLM_MAXTEXT_MODEL}) from Hugging Face (${LLM_HF_REPO})..."
93+
python3 -m maxtext.checkpoint_conversion.to_maxtext \
94+
"${MAXTEXT_PKG_DIR}/configs/base.yml" \
95+
"model_name=${LLM_MAXTEXT_MODEL}" \
96+
"base_output_directory=${LLM_CKPT_DIR}" \
97+
"scan_layers=True" \
98+
"skip_jax_distributed_system=True" \
99+
"--eager_load_method=transformers" \
100+
"--lazy_load_tensors=False" \
101+
"log_config=False"
102+
echo "LLM checkpoint conversion successful!"
103+
echo ""
104+
else
105+
echo "Step 2: LLM checkpoint already exists at ${LLM_ITEMS_PATH}"
106+
fi
107+
108+
# Step 3: Checkpoint Stitching (Vision Tower + LLM Decoder + Fresh Projector)
109+
echo "============================================================"
110+
echo "Stitching Vision and LLM subtrees into unified Omni checkpoint..."
111+
python3 -m maxtext.experimental.omni_poc.utils.stitch_checkpoint \
112+
"$OMNI_CONFIG_PATH" \
113+
"vision_load_path=${VISION_ITEMS_PATH}" \
114+
"llm_load_path=${LLM_ITEMS_PATH}" \
115+
"stitched_output_path=${STITCHED_ITEMS_PATH}"

0 commit comments

Comments
 (0)