Skip to main content

Efficiently select and merge expert LoRAs at Test-Time

Project description

Test-Time Model Merging (TTMM)

A library for efficiently selecting and merging expert LoRAs at Test-Time

Documentation

Please cite our work if you use this library in your research (bibtex below):

Installation

pip install ttmerge

Usage Example

import torch
from transformers import AutoTokenizer, AutoModelForCausalLM
from sentence_transformers import SentenceTransformer
from ttmerge import TestTimeMergingModel

# 1. Load base components
tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-3.2-1B")
base_model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-3.2-1B")
encoder = SentenceTransformer("all-mpnet-base-v2")

# 2. Load normalized corpus embeddings that represent expert domains
# Each row corresponds to an expert: row 0 is the normalized embedding vector of expert 0, row 1 is for expert 1, etc.
# Typically, you would load these pre-computed embeddings from a file
corpus_normalised_embeddings = torch.load("path/to/corpus_embeddings.pt")  # Shape: [n_experts, embedding_dim]

# Directory containing numbered subdirectories (0/, 1/, etc.) where expert LoRAs are stored.
# The individual expert LoRAs should be in the format that the PEFT library uses, 
# with each directory containing adapter_config.json and adapter_model.safetensors files.
adapter_location = "path/to/adapters" 

# 3. Initialize the TestTimeMergingModel
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
merging_model = TestTimeMergingModel(
    corpus=corpus_normalised_embeddings,
    tokenizer=tokenizer,
    encoder=encoder,
    base_model=base_model,
    device=device,
    adapter_location=adapter_location,
    verbose=True
)

# 4. Generate text using relevant expert models.
prompt = "Quantum computing is"
input_ids = tokenizer(prompt, return_tensors="pt").input_ids.to(device)

# The model automatically selects and merges the most relevant expert adapters
generated_ids = merging_model.generate(
    input_ids,
    max_length=512,
    temperature=0.7,
    do_sample=True
)

generated_text = tokenizer.decode(generated_ids[0], skip_special_tokens=True)
print(generated_text)

Development

CI checks

  • The code is auto-formatted using black ..
  • Static type checks can be run using pyright.

Citation

% Citation coming soon.

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

ttmerge-0.0.2.tar.gz (7.4 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

ttmerge-0.0.2-py3-none-any.whl (8.9 kB view details)

Uploaded Python 3

File details

Details for the file ttmerge-0.0.2.tar.gz.

File metadata

  • Download URL: ttmerge-0.0.2.tar.gz
  • Upload date:
  • Size: 7.4 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: poetry/2.1.2 CPython/3.11.5 Darwin/24.0.0

File hashes

Hashes for ttmerge-0.0.2.tar.gz
Algorithm Hash digest
SHA256 e794eedb9e231361268fc3170b526de269272cb28d7a20a959393620c9115727
MD5 6f3265530281ec25133b7139643477df
BLAKE2b-256 4ed7818903bb2882b985b1fdb3ab810cfa2ef8549933f26aacea7b1f67b6ab3d

See more details on using hashes here.

File details

Details for the file ttmerge-0.0.2-py3-none-any.whl.

File metadata

  • Download URL: ttmerge-0.0.2-py3-none-any.whl
  • Upload date:
  • Size: 8.9 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: poetry/2.1.2 CPython/3.11.5 Darwin/24.0.0

File hashes

Hashes for ttmerge-0.0.2-py3-none-any.whl
Algorithm Hash digest
SHA256 d2f1705e8e97389b2a072c3136fb101f44ab5f9f6d8f5f3a0d980e93733caae3
MD5 36486df0f2ad89a633b0bfe0d7db8bce
BLAKE2b-256 10d532e1fd18350fd3227bb48d2795c3851765c139abfb6a830777b573b0e548

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page