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, andint4) 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:
Native LoRA & QLoRA Demo: native_lora_demo.ipynb
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 formaxtext[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) requirescan_layers=Falsedue to per-layer KV sharing. When running SFT with Gemma 4 models, ensurescan_layers=Falseis 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 |
|---|---|---|---|---|
|
|
|
|
Enables/Disables native LoRA wrapping. |
|
|
|
|
The low-rank dimension (\(r\)) of the adapters. |
|
|
|
|
Scaling hyperparameter (\(\alpha\)) for the low-rank updates. |
|
|
|
|
Set to quantization type (e.g., |
|
|
|
|
Tiling dimension for quantized linear layers. |