Skip to content

Repository files navigation

Astra: Activation-Space Tail-Eigenvector Low-Rank Adaptation of Large Language Models

Paper GitHub license

🔍Overview

In this work, we propose Astra (Activation-Space Tail-Eigenvector Low-Rank Adaptation), a novel PEFT method that leverages the tail eigenvectors of the model output activations-estimated from a small task-specific calibration set-to construct task-adaptive low-rank adapters. By constraining updates to the subspace spanned by these tail eigenvectors, Astra achieves faster convergence and improved downstream performance with a significantly reduced parameter budget. Extensive experiments across natural language understanding (NLU) and natural language generation (NLG) tasks demonstrate that Astra consistently outperforms existing PEFT baselines across 16 benchmarks and even surpasses full fine-tuning (FFT) in certain scenarios.

image image

🎯Quick Start

🤗 PEFT Integration

Astra is integrated directly in Hugging Face PEFT. The examples in this repository use the PEFT implementation; until the next PEFT release is published, the dependencies pin the PEFT commit that contains Astra.

import torch
from peft import LoraConfig, get_peft_model
from peft.tuners.lora import AstraConfig, preprocess_astra
from transformers import AutoModelForCausalLM, AutoTokenizer

from dataset.loader import get_calibration_dataloader

model = AutoModelForCausalLM.from_pretrained(
    "meta-llama/Llama-2-7b-hf", dtype=torch.bfloat16, device_map="auto"
)
tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-2-7b-hf")
calibration_dataloader = get_calibration_dataloader(
    "metamath",
    tokenizer,
    num_samples=256,
    batch_size=1,
    seq_len=512,
    padding="max_length",
    calib_on_inputs=True,
)

def run_model():
    for batch in calibration_dataloader:
        batch = {key: value.to(model.device) for key, value in batch.items()}
        with torch.no_grad():
            model(**batch)

lora_config = LoraConfig(
    init_lora_weights="astra",
    r=128,
    lora_alpha=128,
    target_modules=["q_proj", "v_proj", "k_proj", "o_proj", "gate_proj", "down_proj", "up_proj"],
    astra_config=AstraConfig(),
)
preprocess_astra(model, lora_config, run_model=run_model)
peft_model = get_peft_model(model, lora_config)

target_modules must be passed explicitly because preprocessing runs before get_peft_model, which is where PEFT normally infers model-specific defaults.

⚙️Install dependencies

# step 1: create a virtual environment
conda create -n astra python=3.10

# step 2: activate the virtual environment
conda activate astra

# step 3: install dependencies from requirements.txt
pip install -r requirements.txt

The current dependency pins PEFT at commit 73f9a1a9, which includes Astra. Replace this Git dependency with the next PEFT release once available.

💾 Save and convert an Astra adapter

Astra preprocessing creates a residual base model and saves the untrained adapter to astra_init. The initial adapter is required to convert a trained Astra adapter into a standard LoRA adapter:

# After preprocessing, before training:
peft_model.peft_config["default"].init_lora_weights = True
peft_model.save_pretrained(os.path.join(residual_model_path, "astra_init"))

# After training:
peft_model.save_pretrained(
    lora_output_dir,
    path_initial_model_for_weight_conversion=os.path.join(residual_model_path, "astra_init"),
)

The converted adapter can be loaded on top of the original base model and used with standard LoRA tooling.

📦 Prepare datasets

We use the processed datasets uploaded to Huggingface Hub by PiSSA. One can download the datasets using the following code:

from datasets import load_dataset
train_data = load_dataset("fxmeng/pissa-dataset", split="train") # from PiSSA

# MetamathQA dataset
math_types = {"GSM_Rephrased", "GSM_AnsAug", "GSM_SV", "GSM_FOBAR", "MATH_Rephrased", "MATH_AnsAug", "MATH_SV", "MATH_FOBAR"}
train_data = train_data.filter(lambda example: example["type"] in math_types)
train_data.to_json("dataset/metamath/train.json")

# CodeFeedback-Python dataset
code_types = {"python"}
train_data = train_data.filter(lambda example: example["type"] in code_types)
train_data.to_json("dataset/python/train.json")

# Commonsense reasoning
train_data = load_dataset("zwhe99/commonsense_170k", split="train")
train_data.to_json("dataset/commonsense/train.json")

🔁 Reproduce Results

To reproduce the results, please run the following bash scripts:

# metamath
bash scripts/metamath/run.sh

# code
bash scripts/code/run.sh

# commonsense
bash scripts/commonsense/run.sh

About

[ACL'26] The official implementation for "Astra: Activation-Space Tail-Eigenvector Low-Rank Adaptation of Large Language Models"

Topics

Resources

Stars

3 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages