Native LoRA on single-host TPUs#

Native Low-Rank Adaptation (LoRA) in MaxText provides a highly optimized, state-of-the-art parameter-efficient fine-tuning (PEFT) framework.

Unlike traditional adapter wrappers, Native LoRA operates by directly wrapping core model layers. This allows for:

  • Zero adapter-wrapping overhead: Cleaner model codebases and simplified parameter matching.

  • Native Checkpoint Save and Restore: Full out-of-the-box compatibility with Orbax checkpointers, allowing frozen base weights and active adapter parameters to be saved/loaded seamlessly.

  • Weight Quantization (QLoRA): Full support for memory-efficient base weight quantization (including int8, nf4, fp8, and int4) during fine-tuning.

  • Native Trainer Scaling & ZeRO-1: Full access to MaxText’s native trainer infrastructure, including ZeRO-1 optimizer sharding (shard_optimizer_over_data), multidimensional mesh parallelism (FSDP, TP, CP), and Goodput telemetry.

This tutorial provides step-by-step instructions for performing native LoRA/QLoRA fine-tuning and pre-training on single-host TPUs.


🚀 Quick Experimentation with Notebooks#

For interactive playground setups on Google Colab or local JupyterLab, we provide a fully detailed demo notebook:


Install MaxText and Post-Training dependencies#

For instructions on installing MaxText with post-training dependencies on your VM, please refer to the official documentation and use the maxtext[tpu-post-train] installation path to include all necessary post-training dependencies.

Note: If you have previously installed MaxText with a different option (e.g., maxtext[tpu]), we strongly recommend using a fresh virtual environment for maxtext[tpu-post-train] to avoid potential library version conflicts.


Setup environment variables#

Log in to Hugging Face. Provide your access token when prompted:

hf auth login

Set the following environment variables before running LoRA Fine-tuning.

# -- 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., 'qwen3-0.6b' or 'gemma4-e2b'

# -- MaxText configuration --
export BASE_OUTPUT_DIRECTORY=<GCS_BUCKET> # e.g., gs://my-bucket/my-output-directory or /path/to/my-output-directory
export RUN_NAME=<RUN_NAME> # e.g., $(date +%Y-%m-%d-%H-%M-%S)
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']"

Get your 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 or /path/to/my-model-checkpoint/0/items

Option 2: Converting a Hugging Face checkpoint#

Refer to the steps in Hugging Face to MaxText to convert a Hugging Face checkpoint to MaxText. Similar to Option 1, you can set the following environment variable and move on.

export MAXTEXT_CKPT_PATH=<CKPT_PATH> # e.g., gs://my-bucket/my-model-checkpoint/0/items or /path/to/my-model-checkpoint/0/items

Run Native LoRA Fine-Tuning#

Execute the following command to begin LoRA fine-tuning on a Hugging Face dataset (e.g. GSM8K) using the native SFT entrypoint train_sft_native.py:

python3 -m maxtext.trainers.post_train.sft.train_sft_native \
    run_name=${RUN_NAME?} \
    base_output_directory=${BASE_OUTPUT_DIRECTORY?} \
    model_name=${MODEL?} \
    load_parameters_path=${MAXTEXT_CKPT_PATH?} \
    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?} \
    weight_dtype=bfloat16 \
    dtype=bfloat16 \
    formatting_func_path="maxtext.input_pipeline.instruction_data_processing.math_qa_formatting" \
    formatting_func_kwargs="{'template_path': 'src/maxtext/examples/chat_templates/math_qa.json'}" \
    lora.enable_lora=True \
    lora.lora_rank=${LORA_RANK?} \
    lora.lora_alpha=${LORA_ALPHA?}

Note for Gemma 4 architectures: Gemma 4 models (gemma4-e2b, gemma4-e4b) require scan_layers=False due to per-layer KV sharing. When running SFT with Gemma 4 models, ensure scan_layers=False is included.


Run Native Pre-training with QLoRA (8-bit Quantization)#

To run a standard native pre-training loop with memory-efficient 8-bit quantized weights, execute:

python3 -m maxtext.trainers.pre_train.train \
    run_name="native_qlora_pretrain_demo" \
    model_name=${MODEL?} \
    scan_layers=False \
    steps=10 \
    dataset_type="synthetic" \
    per_device_batch_size=1 \
    max_target_length=32 \
    enable_checkpointing=True \
    checkpoint_period=5 \
    base_output_directory="/tmp/native_qlora_pretrain_checkpoint" \
    attention="dot_product" \
    weight_dtype="bfloat16" \
    dtype="bfloat16" \
    lora.enable_lora=True \
    lora.lora_weight_qtype="int8" \
    lora.lora_tile_size=32 \
    lora.lora_rank=4 \
    lora.lora_alpha=8.0

⚙️ LoRA/QLoRA Configuration Reference#

All low-rank adaptation properties are prefixed under the lora. namespace inside the configuration. The key arguments are:

Parameter

Type

Schema Default

Recommended / Demo

Description

lora.enable_lora

bool

False

True

Enables/Disables native LoRA wrapping.

lora.lora_rank

int

0

4 (or 16)

The low-rank dimension (\(r\)) of the adapters.

lora.lora_alpha

float

0.0

8.0 (or 32.0)

Scaling hyperparameter (\(\alpha\)) for the low-rank updates.

lora.lora_weight_qtype

str

""

"int8"

Set to quantization type (e.g., "int8", "nf4", "fp8") to enable quantized weights (QLoRA).

lora.lora_tile_size

int

null

32

Tiling dimension for quantized linear layers.