JORA: JAX Tensor-Parallel LoRA Library for Retrieval Augmented Fine-Tuning

Fuente: arXiv
Saved in:
Bibliographic Details
Main Authors: Tahir, Anique, Cheng, Lu, Liu, Huan
Format: Preprint
Published: 2024
Subjects:
Online Access:
Tags: Add Tag
No Tags, Be the first to tag this record!
_version_ 1866914719483297792
author Tahir, Anique
Cheng, Lu
Liu, Huan
author_facet Tahir, Anique
Cheng, Lu
Liu, Huan
contents The scaling of Large Language Models (LLMs) for retrieval-based tasks, particularly in Retrieval Augmented Generation (RAG), faces significant memory constraints, especially when fine-tuning extensive prompt sequences. Current open-source libraries support full-model inference and fine-tuning across multiple GPUs but fall short of accommodating the efficient parameter distribution required for retrieved context. Addressing this gap, we introduce a novel framework for PEFT-compatible fine-tuning of Llama-2 models, leveraging distributed training. Our framework uniquely utilizes JAX's just-in-time (JIT) compilation and tensor-sharding for efficient resource management, thereby enabling accelerated fine-tuning with reduced memory requirements. This advancement significantly improves the scalability and feasibility of fine-tuning LLMs for complex RAG applications, even on systems with limited GPU resources. Our experiments show more than 12x improvement in runtime compared to Hugging Face/DeepSpeed implementation with four GPUs while consuming less than half the VRAM per GPU.
format Preprint
id arxiv_https___arxiv_org_abs_2403_11366
institution arXiv
publishDate 2024
record_format arxiv
spellingShingle JORA: JAX Tensor-Parallel LoRA Library for Retrieval Augmented Fine-Tuning
Tahir, Anique
Cheng, Lu
Liu, Huan
Machine Learning
Computation and Language
Distributed, Parallel, and Cluster Computing
The scaling of Large Language Models (LLMs) for retrieval-based tasks, particularly in Retrieval Augmented Generation (RAG), faces significant memory constraints, especially when fine-tuning extensive prompt sequences. Current open-source libraries support full-model inference and fine-tuning across multiple GPUs but fall short of accommodating the efficient parameter distribution required for retrieved context. Addressing this gap, we introduce a novel framework for PEFT-compatible fine-tuning of Llama-2 models, leveraging distributed training. Our framework uniquely utilizes JAX's just-in-time (JIT) compilation and tensor-sharding for efficient resource management, thereby enabling accelerated fine-tuning with reduced memory requirements. This advancement significantly improves the scalability and feasibility of fine-tuning LLMs for complex RAG applications, even on systems with limited GPU resources. Our experiments show more than 12x improvement in runtime compared to Hugging Face/DeepSpeed implementation with four GPUs while consuming less than half the VRAM per GPU.
title JORA: JAX Tensor-Parallel LoRA Library for Retrieval Augmented Fine-Tuning
topic Machine Learning
Computation and Language
Distributed, Parallel, and Cluster Computing
url https://arxiv.org/abs/2403.11366