LoRA Fine-tuning on multi-host TPUs#
Low-Rank Adaptation (LoRA) is a Parameter-Efficient Fine-Tuning (PEFT) technique designed to optimize large language models while minimizing resource consumption.
Unlike traditional full-parameter fine-tuning, LoRA:
Freezes the pre-trained model weights, preserving the original knowledge.
Injects trainable rank decomposition matrices into the Transformer layers.
This tutorial provides step-by-step instructions for setting up the multi-host TPU environment and performing LoRA fine-tuning on a Hugging Face dataset using MaxText. In this tutorial we use a multi-host TPU such as v6e-256.
We use Tunix, a JAX-based library, to power these post-training tasks.
Let’s get started!
Prerequisites#
Before starting, ensure you have:
Access to a Google Cloud Project with TPU quotas.
A Hugging Face account with an access token for downloading models.
Permissions for Google Artifact Registry (Artifact Registry Writer role).
Cluster Toolkit installed and configured. Follow Running MaxText with Cluster Toolkit for
gclustersetup.A Pathways-ready GKE cluster configured for Cluster Toolkit, including healthy Kueue and JobSet components (see create a GKE cluster with Pathways and Cluster Toolkit documentation).
Docker installed and configured for sudoless use. Follow the steps to configure sudoless Docker.
Build and upload MaxText Docker image#
For instructions on building and uploading the MaxText Docker image with post-training dependencies, please refer to the official documentation.
Environment configuration#
Set up the following environment variables to configure your training run. Replace placeholders with your actual values.
# -- Model configuration --
# The MaxText model name. See `src/maxtext/configs/types.py` for `ModelName` for a
# full list of supported models.
export MODEL=<MODEL_NAME> # e.g., 'gemma4-26b'
# Your Hugging Face access token. Required to download gated models like Gemma.
# You can generate one at https://huggingface.co/settings/tokens.
export HF_TOKEN=<HF_TOKEN>
# -- MaxText configuration --
# Use a GCS bucket you own to store logs and checkpoints. Ideally in the same
# region as your TPUs to minimize latency and costs.
# You can list your buckets and their locations in the
# [Cloud Console](https://console.cloud.google.com/storage/browser) or via
# `gcloud storage buckets list --format="table(name, location)"`.
export BASE_OUTPUT_DIRECTORY=<GCS_BUCKET> # e.g., gs://my-bucket/maxtext-runs
# An arbitrary string to identify this specific run.
# We recommend to include the model, user, and timestamp.
# Note: Workload names cannot exceed 28 characters (or 22 characters when using Pathways due to Kubernetes 63-byte coordinator label limits, and 20 characters if appending '-convert' or '-to-hf' in the optional conversion steps) and must be valid DNS labels (lowercase alphanumeric and hyphens).
export RUN_NAME=<RUN_NAME>
# -- Workload configuration --
# Your GCP project ID. Find it on the [Cloud Console Dashboard](https://console.cloud.google.com/home/dashboard).
# If you've already set it in your local config, you can retrieve it via:
# gcloud config get-value project
export PROJECT_ID=<PROJECT_ID>
# The GCP location (region or zone) and name of your
# TPU-enabled GKE cluster. Both can be found on the
# [Cloud Console](https://console.cloud.google.com/kubernetes/list).
export LOCATION=<LOCATION> # e.g., 'europe-west4' (region) or 'us-central1-a' (zone)
export GKE_CLUSTER=<CLUSTER_NAME>
# Number of TPU slices for your workload.
export NUM_SLICES=<NUM_SLICES>
# Cluster Toolkit workload placement. Specify the compute type (machine type)
# and topology matching your TPU node pool.
# For example:
# - v5p-128: COMPUTE_TYPE='ct5p-hightpu-4t', TOPOLOGY='4x4x4'
# - v6e-256: COMPUTE_TYPE='ct6e-hightpu-4t', TOPOLOGY='16x16'
# To inspect the accelerator and topology labels on your GKE cluster nodes:
# kubectl get nodes -l cloud.google.com/gke-tpu-accelerator -o custom-columns=NAME:.metadata.name,ACCELERATOR:.metadata.labels.cloud\\.google\\.com/gke-tpu-accelerator,TOPOLOGY:.metadata.labels.cloud\\.google\\.com/gke-tpu-topology
export COMPUTE_TYPE=<COMPUTE_TYPE>
export TOPOLOGY=<TOPOLOGY>
# The Docker image you pushed in the previous step
export DOCKER_IMAGE=<IMAGE_NAME>
# -- Fine-Tuning configuration --
export STEPS=<STEPS> # e.g., 1000
export PER_DEVICE_BATCH_SIZE=<BATCH_SIZE_PER_DEVICE> # e.g., 1
export LORA_RANK=<LORA_RANK> # e.g., 16
export LORA_ALPHA=<LORA_ALPHA> # e.g., 32.0
export LEARNING_RATE=<LEARNING_RATE> # e.g., 3e-6
export MAX_TARGET_LENGTH=<MAX_TARGET_LENGTH> # e.g., 1024
# -- Dataset configuration --
export DATASET_NAME=<DATASET_NAME> # e.g., openai/gsm8k
export TRAIN_SPLIT=<TRAIN_SPLIT> # e.g., train
export HF_DATA_DIR=<DATASET_PATH> # e.g., main
export TRAIN_DATA_COLUMNS=<DATA_COLUMNS> # e.g., ['question','answer']
export CHAT_TEMPLATE_PATH=<TEMPLATE_PATH> # e.g., src/maxtext/examples/chat_templates/gemma_chat.json
# -- LoRA Conversion configuration (Optional) --
export HF_LORA_ADAPTER_PATH=<HF_LORA_ADAPTER_PATH> # e.g., 'username/adapter-name'
Customizing Trainable Layers (Optional)#
By default, MaxText determines which layers to apply LoRA to based on the model’s architecture by reading src/maxtext/configs/post_train/lora_module_path.yml.
If you need to fine-tune specific components (e.g., targeting only Attention layers to optimize memory usage), you can override these defaults through the following hierarchy:
Configuration Hierarchy#
Command Line Argument: Pass the
lora_module_pathargument directly in your training command.Task-Specific Config (
sft.yml): Define thelora_module_pathparameter insrc/maxtext/configs/post_train/sft.yml.Global Defaults: Automatic detection via the model-to-regex mapping defined in
lora_module_path.yml.
Get MaxText model checkpoint#
This section explains how to prepare your model checkpoint for use with MaxText. You have two options: using an existing MaxText checkpoint or converting a Hugging Face checkpoint.
Option 1: Using an existing MaxText checkpoint#
If you already have a MaxText-compatible model checkpoint, simply set the following environment variable and move on to the next section.
export MAXTEXT_CKPT_PATH=<CKPT_PATH> # e.g., gs://my-bucket/my-model-checkpoint/0/items
Note: Make sure that MAXTEXT_CKPT_PATH has the checkpoints created using the correct storage flags:
checkpoint_storage_use_zarr3=False
checkpoint_storage_use_ocdbt=False
Option 2: Converting a Hugging Face checkpoint#
Refer to the steps in Hugging Face to MaxText to convert a hugging face checkpoint to MaxText. Make sure you have correct checkpoint files converted and saved. Similar as Option 1, you can set the following environment and move on.
export MAXTEXT_CKPT_PATH=<CKPT_PATH> # gs://my-bucket/my-checkpoint-directory/0/items
Submit workload on GKE cluster#
This section provides the command to run LoRA Fine-Tuning on a GKE cluster.
Before submitting a job, configure access to the cluster with gcloud and gcluster:
gcloud container clusters get-credentials ${GKE_CLUSTER?} \
--location ${LOCATION?} \
--project ${PROJECT_ID?}
gcluster job config set project ${PROJECT_ID?}
gcluster job config set cluster ${GKE_CLUSTER?}
gcluster job config set location ${LOCATION?}
Run a Fresh LoRA Fine-Tuning on Hugging Face Dataset#
gcluster job submit \
--image=${DOCKER_IMAGE?} \
--name=${RUN_NAME?} \
--pathways \
--compute-type=${COMPUTE_TYPE?} \
--topology=${TOPOLOGY?} \
--num-slices=${NUM_SLICES?} \
--pathways-gcs-location=${BASE_OUTPUT_DIRECTORY?} \
--command="ENABLE_PATHWAYS_PERSISTENCE=1 \
python3 -m maxtext.trainers.post_train.sft.train_sft \
run_name=${RUN_NAME?} \
base_output_directory=${BASE_OUTPUT_DIRECTORY?} \
model_name=${MODEL?} \
load_parameters_path=${MAXTEXT_CKPT_PATH?} \
hf_access_token=${HF_TOKEN?} \
hf_path=${DATASET_NAME?} \
train_split=${TRAIN_SPLIT?} \
hf_data_dir=${HF_DATA_DIR?} \
train_data_columns=${TRAIN_DATA_COLUMNS?} \
steps=${STEPS?} \
per_device_batch_size=${PER_DEVICE_BATCH_SIZE?} \
max_target_length=${MAX_TARGET_LENGTH?} \
learning_rate=${LEARNING_RATE?} \
chat_template_path=${CHAT_TEMPLATE_PATH?} \
lora.enable_lora=True \
lora.lora_rank=${LORA_RANK?} \
lora.lora_alpha=${LORA_ALPHA?} \
checkpoint_storage_use_zarr3=False \
checkpoint_storage_use_ocdbt=False \
enable_single_controller=True"
Once the fine-tuning is completed, you can access your model checkpoints at ${BASE_OUTPUT_DIRECTORY}/${RUN_NAME}/checkpoints.
(Optional) Resume from a previous LoRA checkpoint#
If you want to resume training from a previous run or further fine-tune an existing LoRA adapter, you can specify the LoRA checkpoint path.
Step 1: Convert HF LoRA adapter to MaxText format#
For new deployments, run this conversion with Cluster Toolkit after configuring the cluster with
gcloud container clusters get-credentialsandgcluster job config setabove.
If your LoRA adapter is currently in Hugging Face format, you must convert it to MaxText format before it can be loaded. Use the integrated conversion utility:
gcluster job submit \
--image=${DOCKER_IMAGE?} \
--name=${RUN_NAME?}-convert \
--compute-type=${COMPUTE_TYPE?} \
--topology=${TOPOLOGY?} \
--num-slices=${NUM_SLICES?} \
--command="python3 -m maxtext.checkpoint_conversion.to_maxtext \
model_name=${MODEL?} \
hf_lora_adapter_path=${HF_LORA_ADAPTER_PATH?} \
base_output_directory=${BASE_OUTPUT_DIRECTORY?}/converted_adapter \
hf_access_token=${HF_TOKEN?} \
hardware=cpu \
skip_jax_distributed_system=True"
Step 2: Set the restore path#
Point LORA_RESTORE_PATH to the converted MaxText adapter directory (the directory containing the 0/items or Orbax files).
load_parameters_path: Points to the frozen base model weights (the original model).
lora_restore_path: Points to the previous LoRA adapter weights you wish to load.
export LORA_RESTORE_PATH=<LORA_RESTORE_PATH> # e.g., gs://my-bucket/run-1/checkpoints/0/items or /path/to/run-1/checkpoints/0/items
Step 3: Run LoRA Fine-Tuning with the Restore Path#
Once your environment variables and checkpoints are ready, you can start the LoRA fine-tuning process.
Execute the following command to begin training:
gcluster job submit \
--image=${DOCKER_IMAGE?} \
--name=${RUN_NAME?} \
--pathways \
--compute-type=${COMPUTE_TYPE?} \
--topology=${TOPOLOGY?} \
--num-slices=${NUM_SLICES?} \
--pathways-gcs-location=${BASE_OUTPUT_DIRECTORY?} \
--command="ENABLE_PATHWAYS_PERSISTENCE=1 \
python3 -m maxtext.trainers.post_train.sft.train_sft \
run_name=${RUN_NAME?} \
base_output_directory=${BASE_OUTPUT_DIRECTORY?} \
model_name=${MODEL?} \
load_parameters_path=${MAXTEXT_CKPT_PATH?} \
hf_access_token=${HF_TOKEN?} \
hf_path=${DATASET_NAME?} \
train_split=${TRAIN_SPLIT?} \
hf_data_dir=${HF_DATA_DIR?} \
train_data_columns=${TRAIN_DATA_COLUMNS?} \
steps=${STEPS?} \
per_device_batch_size=${PER_DEVICE_BATCH_SIZE?} \
max_target_length=${MAX_TARGET_LENGTH?} \
lora.lora_restore_path=${LORA_RESTORE_PATH?} \
learning_rate=${LEARNING_RATE?} \
chat_template_path=${CHAT_TEMPLATE_PATH?} \
lora.enable_lora=True \
lora.lora_rank=${LORA_RANK?} \
lora.lora_alpha=${LORA_ALPHA?} \
checkpoint_storage_use_zarr3=False \
checkpoint_storage_use_ocdbt=False \
enable_single_controller=True"
Your fine-tuned model checkpoints will be saved here: $BASE_OUTPUT_DIRECTORY/$RUN_NAME/checkpoints.
(Optional) Convert Fine-tuned LoRA to Hugging Face Format#
For new deployments, run this conversion with Cluster Toolkit after configuring the cluster with
gcloud container clusters get-credentialsandgcluster job config setabove.
After completing the fine-tuning process, your LoRA weights are stored in MaxText/Orbax format. To use these weights with the Hugging Face ecosystem (e.g., for inference or sharing), convert them back using the to_huggingface.py script.
gcluster job submit \
--image=${DOCKER_IMAGE?} \
--name="${RUN_NAME?}-to-hf" \
--compute-type=${COMPUTE_TYPE?} \
--topology=${TOPOLOGY?} \
--num-slices=1 \
--command="python3 -m maxtext.checkpoint_conversion.to_huggingface \
model_name=${MODEL?} \
lora.lora_restore_path=${BASE_OUTPUT_DIRECTORY?}/${RUN_NAME?}/checkpoints/${STEPS?}/model_params \
base_output_directory=${BASE_OUTPUT_DIRECTORY?}/hf_lora_adapter \
hf_access_token=${HF_TOKEN?}"
lora.lora_restore_path: Point this to the specific checkpoint directory (e.g.,.../checkpoints/1000/items) that you want to export.base_output_directory: The local or GCS directory where the Hugging Faceadapter_model.safetensorsandadapter_config.jsonwill be saved.lora.lora_rank/lora.lora_alpha: Must match the values used during the training phase to ensure theadapter_config.jsonis generated correctly.
A Note on Multi-Host Resharding#
When running LoRA fine-tuning in a multi-host environment (e.g., a TPU pod with 64 hosts managing 256 TPUs, such as Pathways), special care must be taken when resharding arrays.
In a single-host environment, the host has a global view of all devices, so a standard jax.device_put can easily distribute slices of data to all local TPUs. However, in a multi-host setup:
Addressability: A host only has a local view of its directly attached devices and cannot push data directly to TPUs managed by other hosts.
Memory Constraints: If every host tries to load the entire weight matrix into RAM just to extract its local piece, the host CPUs will run out of memory (OOM).
To solve this, MaxText uses jax.make_array_from_callback for a “safe reshard.” Instead of pushing data to the devices, this flips the paradigm. It creates a global jax.Array construct where each host locally executes a callback (lambda idx: val[idx]) to load only the specific slice of the data that its attached TPUs need. This completely bypasses cross-host device_put limitations and prevents OOMs since each host only indexes what it requires.