diff --git a/docs/guides/run_python_notebook.md b/docs/guides/run_python_notebook.md index 1d15130f23..ea85908095 100644 --- a/docs/guides/run_python_notebook.md +++ b/docs/guides/run_python_notebook.md @@ -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 | diff --git a/docs/tutorials/post_training_index.md b/docs/tutorials/post_training_index.md index 03460338d3..02ef1c78b6 100644 --- a/docs/tutorials/post_training_index.md +++ b/docs/tutorials/post_training_index.md @@ -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)** @@ -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 diff --git a/docs/tutorials/posttraining/native_lora.md b/docs/tutorials/posttraining/native_lora.md new file mode 100644 index 0000000000..f0847ce548 --- /dev/null +++ b/docs/tutorials/posttraining/native_lora.md @@ -0,0 +1,175 @@ + + +# 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= # e.g., 'qwen3-0.6b' or 'gemma4-e2b' + +# -- MaxText configuration -- +export BASE_OUTPUT_DIRECTORY= # e.g., gs://my-bucket/my-output-directory or /path/to/my-output-directory +export RUN_NAME= # e.g., $(date +%Y-%m-%d-%H-%M-%S) +export STEPS= # e.g., 1000 +export PER_DEVICE_BATCH_SIZE= # e.g., 1 +export LORA_RANK= # e.g., 16 +export LORA_ALPHA= # e.g., 32.0 +export LEARNING_RATE= # e.g., 3e-6 +export MAX_TARGET_LENGTH= # e.g., 1024 + +# -- Dataset configuration -- +export DATASET_NAME= # e.g., openai/gsm8k +export TRAIN_SPLIT= # e.g., train +export HF_DATA_DIR= # e.g., main +export TRAIN_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= # 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= # 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. | diff --git a/src/maxtext/examples/lora_llama3_demo.ipynb b/src/maxtext/examples/lora_llama3_demo.ipynb index 4bafbde24e..99ce8481f3 100644 --- a/src/maxtext/examples/lora_llama3_demo.ipynb +++ b/src/maxtext/examples/lora_llama3_demo.ipynb @@ -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", diff --git a/src/maxtext/examples/native_lora_demo.ipynb b/src/maxtext/examples/native_lora_demo.ipynb new file mode 100644 index 0000000000..cce8529307 --- /dev/null +++ b/src/maxtext/examples/native_lora_demo.ipynb @@ -0,0 +1,310 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "[![Open In Colab](https://colab.research.google.com/assets/colab-badge.svg)](https://colab.research.google.com/github/AI-Hypercomputer/maxtext/blob/main/src/maxtext/examples/native_lora_demo.ipynb)\n", + "\n", + "# Parameter-Efficient Fine-Tuning (PEFT) Demo with Native LoRA and QLoRA" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Overview\n", + "\n", + "This tutorial demonstrates how to run native Parameter-Efficient Fine-Tuning (PEFT) using **LoRA** and **QLoRA** (base weight quantization such as `int8`, `nf4`, or `fp8`) in MaxText across supported models (such as **Gemma4** and **Qwen3**).\n", + "\n", + "We cover two workflows:\n", + "1. **Native LoRA / QLoRA Supervised Fine-Tuning (SFT)** on the GSM8K dataset, starting from converted Hugging Face base checkpoints.\n", + "2. **Native Pre-training with QLoRA** using memory-efficient quantized base weights." + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Prerequisites\n", + "\n", + "Before running this notebook, make sure your environment is set up for the method you are using. Follow the [Run MaxText Python Notebooks on TPUs](https://maxtext.readthedocs.io/en/latest/guides/run_python_notebook.html) guide and complete all steps for your chosen method (Google Colab, VS Code, or Local Jupyter Lab) before proceeding.\n", + "\n", + "If you run into issues, refer to the [Common Pitfalls & Debugging](https://maxtext.readthedocs.io/en/latest/guides/run_python_notebook.html#common-pitfalls-debugging) section of the guide." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "try:\n", + " import google.colab\n", + " print(\"Running the notebook on Google Colab\")\n", + " IN_COLAB = True\n", + "except ImportError:\n", + " print(\"Running the notebook on Visual Studio or JupyterLab\")\n", + " IN_COLAB = False" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Installation: MaxText and Post-Training Dependencies\n", + "\n", + "**Running the notebook on Visual Studio or JupyterLab:** Before proceeding, create a virtual environment and install the required post-training dependencies by following `Option 3: Installing [tpu-post-train]` in the [MaxText installation guide](https://maxtext.readthedocs.io/en/latest/install_maxtext.html#from-source). Once the environment is set up, ensure the notebook is running within it." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "if IN_COLAB:\n", + " # Clone the MaxText repository\n", + " !git clone https://github.com/AI-Hypercomputer/maxtext.git\n", + " %cd maxtext\n", + "\n", + " # Install uv, a fast Python package installer\n", + " !pip install uv\n", + " \n", + " # Install MaxText and post-training dependencies\n", + " import os\n", + " os.environ[\"UV_TORCH_BACKEND\"]=\"cpu\"\n", + " !uv pip install -e .[tpu-post-train] --resolution=lowest\n", + " !install_tpu_post_train_extra_deps" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "**Session restart Instructions for Colab:**\n", + "1. Navigate to the menu at the top of the screen.\n", + "2. Click on **Runtime**.\n", + "3. Select **Restart session** from the dropdown menu.\n", + "\n", + "You will be asked to confirm the action in a pop-up dialog. Click on **Yes**." + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Setup and Imports" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "import os\n", + "import sys\n", + "import subprocess\n", + "from etils import epath" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Model Configurations\n", + "\n", + "Select your target model family below (`gemma4-e2b` or `qwen3-0.6b`). All checkpoint directories, run names, and training configurations are dynamically derived from the model selection." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# Choose model: \"gemma4-e2b\" or \"qwen3-0.6b\"\n", + "MODEL_NAME = \"gemma4-e2b\" # @param [\"gemma4-e2b\", \"qwen3-0.6b\"]\n", + "SCAN_LAYERS = False if MODEL_NAME.startswith(\"gemma\") else True\n", + "\n", + "# Output directories and checkpoint paths derived dynamically from model name\n", + "BASE_OUTPUT_DIRECTORY = f\"/tmp/{MODEL_NAME}_output\"\n", + "MODEL_CHECKPOINT_PATH = f\"{BASE_OUTPUT_DIRECTORY}/{MODEL_NAME}_checkpoint\"\n", + "ORBAX_ITEMS_PATH = os.path.join(MODEL_CHECKPOINT_PATH, \"0/items\")\n", + "PRETRAIN_OUTPUT_DIRECTORY = f\"/tmp/{MODEL_NAME}_qlora_checkpoint\"\n", + "\n", + "print(f\"Selected model: {MODEL_NAME}\")\n", + "print(f\"Scan layers: {SCAN_LAYERS}\")\n", + "print(f\"Base output directory: {BASE_OUTPUT_DIRECTORY}\")\n", + "print(f\"Model checkpoint path: {ORBAX_ITEMS_PATH}\")" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Hugging Face Hub Login\n", + "\n", + "Model checkpoints are hosted on the Hugging Face Hub. Run the cell below to log in so that you can download the base model weights." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "if IN_COLAB:\n", + " from huggingface_hub import notebook_login\n", + " notebook_login()\n", + "else:\n", + " from huggingface_hub import login\n", + " login()" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Download & Convert Checkpoint from Hugging Face\n", + "\n", + "We convert the base Hugging Face checkpoint to the MaxText Orbax format using the standalone checkpoint conversion tool." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "if not epath.Path(ORBAX_ITEMS_PATH).exists():\n", + " print(f\"Converting checkpoint for {MODEL_NAME} from Hugging Face...\")\n", + " env = os.environ.copy()\n", + " env[\"JAX_PLATFORMS\"] = \"cpu\"\n", + " subprocess.run([\n", + " sys.executable, \"-m\", \"maxtext.checkpoint_conversion.to_maxtext\",\n", + " f\"model_name={MODEL_NAME}\",\n", + " f\"base_output_directory={MODEL_CHECKPOINT_PATH}\",\n", + " \"use_multimodal=false\",\n", + " f\"scan_layers={SCAN_LAYERS}\",\n", + " \"skip_jax_distributed_system=True\",\n", + " ], env=env, check=True)\n", + " print(f\"Checkpoint successfully converted to: {ORBAX_ITEMS_PATH}\")\n", + "else:\n", + " print(f\"Model checkpoint already exists at: {ORBAX_ITEMS_PATH}\")" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Native LoRA / QLoRA Supervised Fine-Tuning (SFT)\n", + "\n", + "We run native Supervised Fine-Tuning (SFT) on the real `openai/gsm8k` math dataset with LoRA / QLoRA enabled.\n", + "\n", + "We enable LoRA by passing `lora.enable_lora=True` (and optionally `lora.lora_weight_qtype=int8` or `nf4` for quantized weights), ensuring that only lightweight adapter parameters are trained while the base weights remain frozen." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "subprocess.run([\n", + " sys.executable, \"-m\", \"maxtext.trainers.post_train.sft.train_sft_native\",\n", + " f\"run_name=lora_{MODEL_NAME}_sft_demo\",\n", + " f\"load_parameters_path={ORBAX_ITEMS_PATH}\",\n", + " f\"model_name={MODEL_NAME}\",\n", + " f\"base_output_directory={BASE_OUTPUT_DIRECTORY}\",\n", + " f\"scan_layers={SCAN_LAYERS}\",\n", + " \"hf_path=openai/gsm8k\",\n", + " \"train_split=train\",\n", + " \"hf_data_dir=main\",\n", + " \"train_data_columns=['question', 'answer']\",\n", + " \"steps=5\",\n", + " \"per_device_batch_size=1\",\n", + " \"max_target_length=128\",\n", + " \"learning_rate=3e-6\",\n", + " \"weight_dtype=bfloat16\",\n", + " \"dtype=bfloat16\",\n", + " \"formatting_func_path=maxtext.input_pipeline.instruction_data_processing.math_qa_formatting\",\n", + " \"formatting_func_kwargs={'template_path': 'src/maxtext/examples/chat_templates/math_qa.json'}\",\n", + " \"lora.enable_lora=True\",\n", + " \"lora.lora_weight_qtype=int8\",\n", + " \"lora.lora_tile_size=32\",\n", + " \"lora.lora_rank=4\",\n", + " \"lora.lora_alpha=8.0\",\n", + "], check=True)" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Native Pre-training with QLoRA (Quantized Base Weights)\n", + "\n", + "Next, we demonstrate how to run a native pre-training loop with memory-efficient **QLoRA** enabled (`lora.enable_lora=True`, `lora.lora_weight_qtype=int8`)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "subprocess.run([\n", + " sys.executable, \"-m\", \"maxtext.trainers.pre_train.train\",\n", + " f\"run_name=qlora_{MODEL_NAME}_pretrain_demo\",\n", + " f\"model_name={MODEL_NAME}\",\n", + " f\"scan_layers={SCAN_LAYERS}\",\n", + " \"steps=1\",\n", + " \"dataset_type=synthetic\",\n", + " \"per_device_batch_size=1\",\n", + " \"max_target_length=32\",\n", + " \"enable_checkpointing=True\",\n", + " \"checkpoint_period=1\",\n", + " \"async_checkpointing=False\",\n", + " \"sharding_tolerance=1.0\",\n", + " f\"base_output_directory={PRETRAIN_OUTPUT_DIRECTORY}\",\n", + " \"attention=dot_product\",\n", + " \"weight_dtype=bfloat16\",\n", + " \"dtype=bfloat16\",\n", + " \"lora.enable_lora=True\",\n", + " \"lora.lora_weight_qtype=int8\",\n", + " \"lora.lora_tile_size=32\",\n", + " \"lora.lora_rank=4\",\n", + " \"lora.lora_alpha=8.0\",\n", + "], check=True)" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Conclusion\n", + "\n", + "This tutorial successfully demonstrated how to run native Parameter-Efficient Fine-Tuning (PEFT) with LoRA and QLoRA in MaxText." + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## \ud83d\udcda Learn More\n", + "\n", + "- **CLI Usage**: Refer to [Native LoRA and QLoRA Fine-tuning](https://github.com/AI-Hypercomputer/maxtext/blob/main/docs/tutorials/posttraining/native_lora.md) for detailed CLI orchestrations.\n", + "- **Configuration**: See `src/maxtext/configs/post_train/sft.yml` and `src/maxtext/configs/base.yml` for all available LoRA options (prefixed with `lora.`).\n", + "- **Documentation**: Check `src/maxtext/trainers/post_train/sft/train_sft_native.py` and `src/maxtext/trainers/pre_train/train.py` for native training implementations." + ] + } + ], + "metadata": { + "language_info": { + "name": "python" + } + }, + "nbformat": 4, + "nbformat_minor": 2 +} \ No newline at end of file