Bio-CLIP / bio_clip_recommender.py
hohdavid's picture
modified tokenizer
58816bc verified
Raw
History Blame Contribute Delete
6.06 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
import pickle
import numpy as np
import re
from transformers import T5Tokenizer, T5EncoderModel
class BioCLIP(nn.Module):
def __init__(self, seq_dim=1024, text_dim=768, shared_dim=512):
super().__init__()
self.seq_projector = nn.Sequential(
nn.Linear(seq_dim, shared_dim),
nn.GELU(),
nn.Dropout(0.1),
nn.Linear(shared_dim, shared_dim)
)
self.text_projector = nn.Sequential(
nn.Linear(text_dim, shared_dim),
nn.GELU(),
nn.Dropout(0.1),
nn.Linear(shared_dim, shared_dim)
)
self.temperature = nn.Parameter(torch.ones([]) * np.log(1 / 0.07))
def forward(self, seq_emb, text_emb):
z_seq = F.normalize(self.seq_projector(seq_emb), p=2, dim=-1)
z_text = F.normalize(self.text_projector(text_emb), p=2, dim=-1)
return z_seq, z_text
class ProteinEmbedder:
def __init__(self, device="cpu"):
print("Loading ProtT5 Language Model (this takes a moment)...")
self.device = device
self.tokenizer = T5Tokenizer.from_pretrained("Rostlab/prot_t5_xl_uniref50", do_lower_case=False)
self.model = T5EncoderModel.from_pretrained("Rostlab/prot_t5_xl_uniref50").to(self.device)
self.model.eval()
print("ProtT5 Online!")
def embed_raw_sequence(self, sequence: str):
seq = re.sub(r"[UZOB]", "X", sequence.upper())
seq_spaced = " ".join(list(seq))
with torch.no_grad():
ids = self.tokenizer([seq_spaced], add_special_tokens=True, padding=True, return_tensors="pt")
input_ids = ids['input_ids'].to(self.device)
attention_mask = ids['attention_mask'].to(self.device)
embedding = self.model(input_ids=input_ids, attention_mask=attention_mask)
embedding = embedding.last_hidden_state
seq_len = (attention_mask[0] == 1).sum()
protein_emb = embedding[0, :seq_len-1].mean(dim=0)
return protein_emb.cpu().numpy()
class ProtocolRecommender:
def __init__(self, device="cpu"):
self.device = device
self.steps = ['lysis', 'elution', 'desalting']
self.models = {}
self.databases = {}
print("Loading Bio-CLIP Expert Models...")
for step in self.steps:
model = BioCLIP().to(self.device)
weights_path = f"bioclip_weights_{step}.pth"
model.load_state_dict(torch.load(weights_path, map_location=self.device, weights_only=True))
model.eval()
self.models[step] = model
db_path = f"aligned_spaces_{step}.pkl"
with open(db_path, "rb") as f:
db = pickle.load(f)
self.databases[step] = {
"text_vectors": torch.tensor(db["aligned_text"]).float().to(self.device),
"raw_texts": db["raw_texts"]
}
print("Ready! All systems online.\n")
def search(self, protein_sequence_1024d, top_k=3):
results = {}
if not isinstance(protein_sequence_1024d, torch.Tensor):
seq_tensor = torch.tensor(protein_sequence_1024d).float().unsqueeze(0).to(self.device)
else:
seq_tensor = protein_sequence_1024d.to(self.device)
if seq_tensor.dim() == 1:
seq_tensor = seq_tensor.unsqueeze(0)
with torch.no_grad():
for step in self.steps:
model = self.models[step]
db = self.databases[step]
z_query = model.seq_projector(seq_tensor)
z_query = F.normalize(z_query, p=2, dim=-1)
similarities = (z_query @ db["text_vectors"].T).squeeze()
top_scores, top_indices = torch.topk(similarities, k=top_k)
step_results = []
for score, idx in zip(top_scores.cpu().numpy(), top_indices.cpu().numpy()):
step_results.append({
"confidence": score,
"text": db["raw_texts"][idx]
})
results[step] = step_results
return results
if __name__ == "__main__":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
try:
embedder = ProteinEmbedder(device=device)
engine = ProtocolRecommender(device=device)
print("\n" + "="*60)
raw_amino_acids = "MSKGEELFTGVVPILVELDGDVNGHKFSVSGEGEGDATYGKLTLKFICTTGKLPVPWPTLVTTFSYGVQCFSRYPDHMKQHDFFKSAMPEGYVQERTIFFKDDGNYKTRAEVKFEGDTLVNRIELKGIDFKEDGNILGHKLEYNYNSHNVYIMADKQKNGIKVNFKIRHNIEDGSVQLADHYQQNTPIGDGPVLLPDNHYLSTQSALSKDPNEKRDHMVLLEFVTAAGITHGMDELYK"
print(f"User inputted sequence: {raw_amino_acids[:30]}... (Length: {len(raw_amino_acids)})")
print("="*60)
print("1. Passing sequence through ProtT5...")
new_protein_vector = embedder.embed_raw_sequence(raw_amino_acids)
print("2. Querying Bio-CLIP databases...\n")
recommendations = engine.search(new_protein_vector, top_k=3)
print("="*60)
print("BIO-CLIP RECOMMENDED PURIFICATION PIPELINE")
print("="*60)
for step in ['lysis', 'elution', 'desalting']:
print(f"\n--- {step.upper()} EXPERT ---")
for i, rec in enumerate(recommendations[step]):
confidence_pct = max(0, min(100, (rec['confidence'] + 0.1) * 100))
print(f"Option {i+1} [Confidence: {confidence_pct:.1f}%]")
print(f"Protocol: {rec['text']}\n")
except FileNotFoundError as e:
print(f"\nERROR: Could not find model files. Make sure you run this in the same folder as your .pth and .pkl files!")
print(f"Details: {e}")