Skip to content

Commit 83a529c

Browse files
committed
docs(nnx): add native LoRA Gemma4, Qwen3, and Llama3 tutorial notebooks
1 parent f17047b commit 83a529c

6 files changed

Lines changed: 922 additions & 1 deletion

File tree

‎docs/guides/run_python_notebook.md‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -194,6 +194,11 @@ jupyter lab --ip=0.0.0.0 --port=8888 --no-browser --allow-root
194194

195195
- **`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).
196196

197+
### Parameter-Efficient Fine-Tuning (PEFT/LoRA)
198+
199+
- **`qwen3_native_lora_demo.ipynb`** → Qwen3-0.6B PEFT training and evaluation with native LoRA and QLoRA. Includes both SFT training on [OpenAI's GSM8K dataset](https://huggingface.co/datasets/openai/gsm8k) and pre-training. Runs successfully on free-tier Google Colab TPUs.
200+
- **`gemma4_native_lora_demo.ipynb`** → Gemma4-e2b PEFT training and evaluation with native LoRA and QLoRA, demonstrating multi-query attention (MQA) support under Flax NNX. 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).
201+
197202
## Common Pitfalls & Debugging
198203

199204
| Issue | Solution |

‎docs/tutorials/post_training_index.md‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,7 @@ MaxText was co-designed with key Google led innovations to provide a unified pos
2727
- [SFT on Single-Host TPUs](./posttraining/sft.md)
2828
- [SFT on Multi-Host TPUs](./posttraining/sft_on_multi_host.md)
2929
- **LoRA (Low-Rank Adaptation)**
30+
- [Native LoRA/QLoRA on Single-Host TPUs](./posttraining/native_lora.md)
3031
- [LoRA on Single-Host TPUs](./posttraining/lora.md)
3132
- [LoRA on Multi-Host TPUs](./posttraining/lora_on_multi_host.md)
3233
- **DPO (Direct Preference Optimization) and ORPO (Odds-Ratio Policy Optimization)**
@@ -79,6 +80,7 @@ posttraining/rl_qwen3_30b.md
7980
posttraining/rl_gptoss_20b.md
8081
posttraining/knowledge_distillation.md
8182
posttraining/lora.md
83+
posttraining/native_lora.md
8284
posttraining/lora_on_multi_host.md
8385
posttraining/multimodal.md
8486
posttraining/full_finetuning.md
Lines changed: 166 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,166 @@
1+
<!--
2+
Copyright 2023–2026 Google LLC
3+
4+
Licensed under the Apache License, Version 2.0 (the "License");
5+
you may not use this file except in compliance with the License.
6+
You may obtain a copy of the License at
7+
8+
https://www.apache.org/licenses/LICENSE-2.0
9+
10+
Unless required by applicable law or agreed to in writing, software
11+
distributed under the License is distributed on an "AS IS" BASIS,
12+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13+
See the License for the specific language governing permissions and
14+
limitations under the License.
15+
-->
16+
17+
# Native NNX LoRA on single-host TPUs
18+
19+
**Native Low-Rank Adaptation (LoRA)** under pure Flax NNX in MaxText provides a highly optimized, state-of-the-art parameter-efficient fine-tuning (PEFT) framework.
20+
21+
Unlike traditional Linen/Tunix adapter wrappers, Native NNX LoRA operates by directly wrapping NNX modules. This allows for:
22+
23+
- **Zero adapter-wrapping overhead**: Cleaner model codebases and simplified parameter matching.
24+
- **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.
25+
- **Int8 Weight Quantization (QLoRA)**: Full support for memory-efficient 8-bit quantization during fine-tuning.
26+
27+
This tutorial provides step-by-step instructions for performing native LoRA/QLoRA fine-tuning and pre-training on single-host TPUs using pure Flax NNX.
28+
29+
______________________________________________________________________
30+
31+
## 🚀 Quick Experimentation with Notebooks
32+
33+
For interactive playground setups on Google Colab or local JupyterLab, we provide fully detailed demo notebooks:
34+
35+
- **Qwen3 Native LoRA Demo**: [qwen3_native_lora_demo.ipynb](https://github.com/AI-Hypercomputer/maxtext/blob/main/src/maxtext/examples/qwen3_native_lora_demo.ipynb)
36+
- **Gemma4 Native LoRA Demo**: [gemma4_native_lora_demo.ipynb](https://github.com/AI-Hypercomputer/maxtext/blob/main/src/maxtext/examples/gemma4_native_lora_demo.ipynb)
37+
38+
______________________________________________________________________
39+
40+
## Setup environment variables
41+
42+
Log in to Hugging Face. Provide your access token when prompted:
43+
44+
```bash
45+
hf auth login
46+
```
47+
48+
Set the following environment variables before running LoRA Fine-tuning.
49+
50+
```sh
51+
# -- Model configuration --
52+
export MODEL_NAME=<MODEL_NAME> # e.g., 'qwen3-0.6b' or 'gemma4-e2b'
53+
export TOKENIZER_PATH=<TOKENIZER_PATH> # e.g., 'Qwen/Qwen3-0.6B' or 'google/gemma-4-E2B-it'
54+
55+
# -- MaxText configuration --
56+
export BASE_OUTPUT_DIRECTORY=<GCS_BUCKET> # e.g., gs://my-bucket/my-output-directory or /path/to/my-output-directory
57+
export RUN_NAME=<RUN_NAME> # e.g., $(date +%Y-%m-%d-%H-%M-%S)
58+
export STEPS=<STEPS> # e.g., 1000
59+
export PER_DEVICE_BATCH_SIZE=<BATCH_SIZE_PER_DEVICE> # e.g., 1
60+
export LORA_RANK=<LORA_RANK> # e.g., 16
61+
export LORA_ALPHA=<LORA_ALPHA> # e.g., 32.0
62+
export LEARNING_RATE=<LEARNING_RATE> # e.g., 3e-6
63+
export MAX_TARGET_LENGTH=<MAX_TARGET_LENGTH> # e.g., 1024
64+
65+
# -- Dataset configuration --
66+
export DATASET_NAME=<DATASET_NAME> # e.g., openai/gsm8k
67+
export TRAIN_SPLIT=<TRAIN_SPLIT> # e.g., train
68+
export HF_DATA_DIR=<DATASET_PATH> # e.g., main
69+
export TRAIN_DATA_COLUMNS=<DATA_COLUMNS> # e.g., "['question','answer']"
70+
```
71+
72+
______________________________________________________________________
73+
74+
## Get your model checkpoint
75+
76+
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.
77+
78+
### Option 1: Using an existing MaxText checkpoint
79+
80+
If you already have a MaxText-compatible model checkpoint, simply set the following environment variable and move on to the next section.
81+
82+
```sh
83+
export MAXTEXT_CKPT_PATH=<CKPT_PATH> # e.g., gs://my-bucket/my-model-checkpoint/0/items or /path/to/my-model-checkpoint/0/items
84+
```
85+
86+
### Option 2: Converting a Hugging Face checkpoint
87+
88+
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.
89+
90+
```sh
91+
export MAXTEXT_CKPT_PATH=<CKPT_PATH> # e.g., gs://my-bucket/my-model-checkpoint/0/items or /path/to/my-model-checkpoint/0/items
92+
```
93+
94+
______________________________________________________________________
95+
96+
## Run Native LoRA Fine-Tuning
97+
98+
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`:
99+
100+
```sh
101+
python3 -m maxtext.trainers.post_train.sft.train_sft_native \
102+
src/maxtext/configs/post_train/sft.yml \
103+
run_name="${RUN_NAME?}" \
104+
base_output_directory="${BASE_OUTPUT_DIRECTORY?}" \
105+
model_name="${MODEL_NAME?}" \
106+
load_parameters_path="${MAXTEXT_CKPT_PATH?}" \
107+
tokenizer_path="${TOKENIZER_PATH?}" \
108+
hf_path="${DATASET_NAME?}" \
109+
train_split="${TRAIN_SPLIT?}" \
110+
hf_data_dir="${HF_DATA_DIR?}" \
111+
train_data_columns="${TRAIN_DATA_COLUMNS?}" \
112+
steps="${STEPS?}" \
113+
per_device_batch_size="${PER_DEVICE_BATCH_SIZE?}" \
114+
max_target_length="${MAX_TARGET_LENGTH?}" \
115+
learning_rate="${LEARNING_RATE?}" \
116+
weight_dtype=bfloat16 \
117+
dtype=bfloat16 \
118+
formatting_func_path="maxtext.input_pipeline.instruction_data_processing.math_qa_formatting" \
119+
formatting_func_kwargs="{'template_path': 'src/maxtext/examples/chat_templates/math_qa.json'}" \
120+
lora.enable_lora=True \
121+
lora.lora_rank="${LORA_RANK?}" \
122+
lora.lora_alpha="${LORA_ALPHA?}"
123+
```
124+
125+
______________________________________________________________________
126+
127+
## Run Native Pre-training with QLoRA (8-bit Quantization)
128+
129+
To run a standard native pre-training loop with memory-efficient 8-bit quantized weights, execute:
130+
131+
```sh
132+
python3 -m maxtext.trainers.pre_train.train \
133+
src/maxtext/configs/base.yml \
134+
run_name="native_qlora_pretrain_demo" \
135+
model_name="gemma4-e2b" \
136+
scan_layers=False \
137+
steps=10 \
138+
dataset_type="synthetic" \
139+
per_device_batch_size=1 \
140+
max_target_length=32 \
141+
enable_checkpointing=True \
142+
checkpoint_period=5 \
143+
base_output_directory="/tmp/native_qlora_pretrain_checkpoint" \
144+
attention="dot_product" \
145+
weight_dtype="bfloat16" \
146+
dtype="bfloat16" \
147+
lora.enable_lora=True \
148+
lora.lora_weight_qtype="int8" \
149+
lora.lora_tile_size=32 \
150+
lora.lora_rank=4 \
151+
lora.lora_alpha=8.0
152+
```
153+
154+
______________________________________________________________________
155+
156+
## ⚙️ LoRA/QLoRA Configuration Reference
157+
158+
All low-rank adaptation properties are prefixed under the `lora.` namespace inside the configuration. The key arguments are:
159+
160+
| Parameter | Type | Default | Description |
161+
| ------------------------ | ------- | ------- | ----------------------------------------------------------- |
162+
| `lora.enable_lora` | `bool` | `False` | Enables/Disables native LoRA wrapping. |
163+
| `lora.lora_rank` | `int` | `4` | The low-rank dimension ($r$) of the adapters. |
164+
| `lora.lora_alpha` | `float` | `8.0` | Scaling hyperparameter ($\alpha$) for the low-rank updates. |
165+
| `lora.lora_weight_qtype` | `str` | `""` | Set to `"int8"` to enable 8-bit quantized weights (QLoRA). |
166+
| `lora.lora_tile_size` | `int` | `32` | Tiling dimension for quantized linear layers. |

0 commit comments

Comments
 (0)