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 expert 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
expert_embeddings = torch.load("path/to/expert_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(
    expert_embeddings=expert_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.3.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.3-py3-none-any.whl (8.9 kB view details)

Uploaded Python 3

File details

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

File metadata

  • Download URL: ttmerge-0.0.3.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.3.tar.gz
Algorithm Hash digest
SHA256 49315664f728acdedeaedfded8ed922a649a33678f36679d6cfe98f1a49b6035
MD5 68e738f39bb379e1f2624405ee4c246f
BLAKE2b-256 15f1b212f8fb60a88921b019c787d077037c1282278763358fececcac3719e20

See more details on using hashes here.

File details

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

File metadata

  • Download URL: ttmerge-0.0.3-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.3-py3-none-any.whl
Algorithm Hash digest
SHA256 d0b672a7a1ccec31d94a896b12d58d1f666f5bcc89294d942838cb047bff7e70
MD5 a00e194bd8127f75165290c92427b7db
BLAKE2b-256 b6164c0c8dff06fadc748921ec239d3a8e66750bf67bc5eb80c04d7147b36929

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