Skip to content

Commit 542fb6d

Browse files
committed
Add omni_checkpoint_stitcher and omni-gemma3-qwen3 configuration
Add omni_checkpoint_stitcher and omni-gemma3-qwen3 configuration pylint add test file POC Phase 1: Implement checkpoint stitching and preparation pipeline reorg files better file description pyink pyink pyink improve print statement pyink pyink reorganize files rename files rm files small changes rm internal dir format
1 parent 7b4a493 commit 542fb6d

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)