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
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)
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
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
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
49315664f728acdedeaedfded8ed922a649a33678f36679d6cfe98f1a49b6035
|
|
| MD5 |
68e738f39bb379e1f2624405ee4c246f
|
|
| BLAKE2b-256 |
15f1b212f8fb60a88921b019c787d077037c1282278763358fececcac3719e20
|
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
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
d0b672a7a1ccec31d94a896b12d58d1f666f5bcc89294d942838cb047bff7e70
|
|
| MD5 |
a00e194bd8127f75165290c92427b7db
|
|
| BLAKE2b-256 |
b6164c0c8dff06fadc748921ec239d3a8e66750bf67bc5eb80c04d7147b36929
|