<!--
 Copyright 2023–2026 Google LLC

 Licensed under the Apache License, Version 2.0 (the "License");
 you may not use this file except in compliance with the License.
 You may obtain a copy of the License at

      https://www.apache.org/licenses/LICENSE-2.0

 Unless required by applicable law or agreed to in writing, software
 distributed under the License is distributed on an "AS IS" BASIS,
 WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
 See the License for the specific language governing permissions and
 limitations under the License.
 -->

# 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](https://github.com/google/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](../../run_maxtext/run_maxtext_via_cluster_toolkit.md) for `gcluster` setup.
- A Pathways-ready GKE cluster configured for Cluster Toolkit, including healthy Kueue and JobSet components (see [create a GKE cluster with Pathways](https://docs.cloud.google.com/ai-hypercomputer/docs/workloads/pathways-on-cloud/create-gke-cluster) and [Cluster Toolkit documentation](https://cloud.google.com/cluster-toolkit/docs/overview)).
- **Docker** installed and configured for sudoless use. Follow the steps to [configure sudoless Docker](https://docs.docker.com/engine/install/linux-postinstall/).

## 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](../build_maxtext.md).

## Environment configuration

Set up the following environment variables to configure your training run. Replace placeholders with your actual values.

```bash
# -- 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

1. **Command Line Argument**: Pass the `lora_module_path` argument directly in your training command.
2. **Task-Specific Config (`sft.yml`)**: Define the `lora_module_path` parameter in `src/maxtext/configs/post_train/sft.yml`.
3. **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.

```bash
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:

```bash
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](hf-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.

```bash
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`:

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

```bash
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-credentials` and `gcluster job config set` above.

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:

```sh
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.

```sh
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:

```bash
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-credentials` and `gcluster job config set` above.

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.

```sh
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 Face `adapter_model.safetensors` and `adapter_config.json` will be saved.
- `lora.lora_rank` / `lora.lora_alpha`: Must match the values used during the training phase to ensure the `adapter_config.json` is 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.
