Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions docs/guides/run_python_notebook.md
Original file line number Diff line number Diff line change
Expand Up @@ -194,6 +194,10 @@ jupyter lab --ip=0.0.0.0 --port=8888 --no-browser --allow-root

- **`rl_llama3_demo.ipynb`** → GRPO/GSPO training on [OpenAI's GSM8K dataset](https://huggingface.co/datasets/openai/gsm8k). We recommend running this on a v5p-8 TPU VM using [Method 2](#method-2-visual-studio-code-with-tpu-recommended) or [Method 3](#method-3-local-jupyter-lab-with-tpu-recommended).

### Parameter-Efficient Fine-Tuning (PEFT/LoRA)

- **`native_lora_demo.ipynb`** → Interactive Parameter-Efficient Fine-Tuning (PEFT) and pre-training demo with native LoRA and QLoRA for supported models (such as Qwen3-0.6B and Gemma4-e2b). Includes SFT on [OpenAI's GSM8K dataset](https://huggingface.co/datasets/openai/gsm8k) and pre-training with QLoRA. Runs successfully on free-tier Google Colab TPUs or TPU VMs.

## Common Pitfalls & Debugging

| Issue | Solution |
Expand Down
2 changes: 2 additions & 0 deletions docs/tutorials/post_training_index.md
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ MaxText was co-designed with key Google led innovations to provide a unified pos
- [SFT on Single-Host TPUs](./posttraining/sft.md)
- [SFT on Multi-Host TPUs](./posttraining/sft_on_multi_host.md)
- **LoRA (Low-Rank Adaptation)**
- [Native LoRA/QLoRA on Single-Host TPUs](./posttraining/native_lora.md)
- [LoRA on Single-Host TPUs](./posttraining/lora.md)
- [LoRA on Multi-Host TPUs](./posttraining/lora_on_multi_host.md)
- **DPO (Direct Preference Optimization) and ORPO (Odds-Ratio Policy Optimization)**
Expand Down Expand Up @@ -79,6 +80,7 @@ posttraining/rl_qwen3_30b.md
posttraining/rl_gptoss_20b.md
posttraining/knowledge_distillation.md
posttraining/lora.md
posttraining/native_lora.md
posttraining/lora_on_multi_host.md
posttraining/multimodal.md
posttraining/full_finetuning.md
Expand Down
175 changes: 175 additions & 0 deletions docs/tutorials/posttraining/native_lora.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,175 @@
<!--
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.
-->

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

- **Native LoRA & QLoRA Demo**: [native_lora_demo.ipynb](https://github.com/AI-Hypercomputer/maxtext/blob/main/src/maxtext/examples/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](../../install_maxtext.md) 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.

______________________________________________________________________

Comment thread
SurbhiJainUSC marked this conversation as resolved.
## Setup environment variables

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

```bash
hf auth login
```

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

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

```sh
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](../../guides/checkpointing_solutions/convert_checkpoint.md#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.

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

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

```sh
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. |
2 changes: 1 addition & 1 deletion src/maxtext/examples/lora_llama3_demo.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -186,7 +186,7 @@
"metadata": {},
"outputs": [],
"source": [
"MODEL_NAME = \"llama3.1-8b\"\n",
"MODEL_NAME = \"llama3.1-8b-Instruct\"\n",
"TOKENIZER_PATH = \"meta-llama/Llama-3.1-8B-Instruct\"\n",
"tokenizer = transformers.AutoTokenizer.from_pretrained(TOKENIZER_PATH)\n",
"# This is the directory where the fine-tuned model checkpoint will be saved\n",
Expand Down
Loading
Loading