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