Customize Gemma with Hugging Face Transformers

By Google for Developers

Share:

Key Concepts

  • Quantized LoRA (QLoRA): A parameter-efficient fine-tuning technique that quantizes pre-trained model weights to 4-bit and adds trainable LoRA adapters.
  • LoRA (Low-Rank Adaptation): A parameter-efficient fine-tuning method that introduces trainable low-rank matrices to adapt a pre-trained model to a specific task.
  • Parameter-Efficient Fine-Tuning (PEFT): Techniques that allow fine-tuning large language models with a small number of trainable parameters.
  • Text-to-SQL: A task where a model generates a SQL query based on a natural language instruction and a database schema.
  • Hugging Face Transformers: A library providing pre-trained models and tools for natural language processing.
  • BitsAndBytes: A library for quantizing models to reduce memory footprint.
  • PEFT (Parameter-Efficient Fine-Tuning): A library for parameter-efficient fine-tuning methods.
  • SFTTrainer: A class from the Hugging Face TRL library for supervised fine-tuning.
  • Gemma: A family of open-source language models by Google DeepMind.

Fine-Tuning Gemma with QLoRA for Text-to-SQL

This video demonstrates how to fine-tune the Gemma model for a text-to-SQL task using Hugging Face Transformers and QLoRA. The goal is to train the model to generate SQL queries from natural language instructions, given a database schema.

1. Introduction to QLoRA

  • QLoRA is introduced as a parameter-efficient fine-tuning method.
  • Pre-trained model weights are quantized to 4-bit and frozen.
  • LoRA linear adapters are added and trained.
  • This significantly reduces memory usage during training.
  • After training, the LoRA adapters can be merged into the pre-trained model or used separately for inference.

2. Setting up the Environment

  • Installing Required Libraries:
    • The first step involves installing necessary libraries using pip install.
    • Includes torch, tensorboard, transformers, datasets, bitsandbytes, and peft.
    • bitsandbytes is crucial for QLoRA, enabling 4-bit quantization.
    • peft provides tools for parameter-efficient fine-tuning.
  • Logging into Hugging Face Hub:
    • Access to Gemma requires accepting terms of use on the Hugging Face Hub.
    • Users need a Hugging Face account and must accept the Gemma license on the huggingface.co/jemma repository.

3. Preparing the Text-to-SQL Dataset

  • Dataset Creation:
    • A text-to-SQL dataset is used, where the model learns to generate SQL queries from natural language instructions.
    • The input consists of a user prompt containing the SQL schema and a user query.
    • The model's target output is the corresponding SQL query.
  • Loading and Sampling the Dataset:
    • The dataset is loaded from the Hugging Face Hub using the load_dataset method.
    • The dataset is downsampled to 12,000 examples for faster training.
  • Example Inspection:
    • An example from the dataset is printed to illustrate the input format.
    • The example includes the "assistance message" (SQL schema) and the "user message" (natural language query).
    • The model is expected to generate the correct SQL query based on this input.
    • Example: User asks "how many metal health parity violations were reported in each state and by violation type" and the model should generate "select state violation type..."

4. Fine-Tuning the Gemma Model

  • Loading the Model and Tokenizer:
    • The Gemma model is loaded using AutoModelForCausalLM from the transformers library.
    • The bitsandbytes configuration is used to quantize the model to 4-bit.
    • The tokenizer is loaded using the Gemma instruction template.
  • Configuring LoRA:
    • A LoRA configuration is defined to specify the size and number of adapter layers.
    • This configuration determines the trainable parameters added to the model.
  • Defining Training Hyperparameters:
    • Training hyperparameters are defined, including:
      • Number of epochs (e.g., 3).
      • Batch size.
      • Learning rate.
      • Logging frequency.
      • Use of TensorBoard for tracking training progress.
  • Setting up the Trainer:
    • The SFTTrainer from the Hugging Face TRL library is used to set up the training process.
    • The trainer takes the model, training arguments, dataset, LoRA configuration, and tokenizer as input.
    • The trainer applies the tokenizer to the training samples to create the Gemma instruction template.
  • Starting the Training:
    • The training is started using trainer.train().
    • The training process takes approximately 2 hours for 3 epochs and 150 steps.

5. Training Results and Model Saving

  • Training Progress:
    • The training loss decreases from approximately 2 to 0.6.
  • Model Saving:
    • The trained model, including the adapter configuration and model weights, is saved in the "Gemma-text-to-SQL" folder.
    • The saved model contains the adapter_config.json and adapter_model.bin files.
  • Inference:
    • The trained LoRA adapters can be merged back into the base model or used separately for inference.
    • The video refers to a separate guide or Jupyter notebook for running inference and generating SQL queries.

6. Conclusion

The video provides a practical guide to fine-tuning the Gemma model for a text-to-SQL task using QLoRA and Hugging Face Transformers. It covers the essential steps, from setting up the environment and preparing the data to configuring the training process and saving the trained model. The use of QLoRA enables efficient fine-tuning with reduced memory requirements, making it feasible to adapt large language models for specific tasks.

Chat with this Video

AI-Powered

Load the transcript when you're ready to chat so the initial page stays lighter.

Ready to summarize another video?

Summarize YouTube Video