File size: 6,055 Bytes
c26bf25 58816bc c26bf25 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 | 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}") |