|
| 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