-
Notifications
You must be signed in to change notification settings - Fork 617
docs(nnx): add native LoRA Gemma4 and Qwen3 notebooks and update tutorial guides #4417
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
Merged
Changes from all commits
Commits
File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| 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. | ||
|
|
||
| ______________________________________________________________________ | ||
|
|
||
| ## 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. | | ||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.