Views
No views yet
esm2_t6_8M_UR50D,
trained on 166 protein sequences in the RNA binding sites dataset
using a 75/25 train/test split. It achieves an evaluation loss of 0.1791934072971344.1from transformers import AutoModelForTokenClassification, AutoTokenizer
2from peft import PeftModel
3import torch
4
5# Path to the saved LoRA model
6model_path = "AmelieSchreiber/esm2_t6_8M_UR50D_LoRA_RNA-binding"
7# ESM2 base model
8base_model_path = "facebook/esm2_t6_8M_UR50D"
9
10# Load the model
11base_model = AutoModelForTokenClassification.from_pretrained(base_model_path)
12loaded_model = PeftModel.from_pretrained(base_model, model_path)
13
14# Ensure the model is in evaluation mode
15loaded_model.eval()
16
17# Load the tokenizer
18loaded_tokenizer = AutoTokenizer.from_pretrained(base_model_path)
19
20# Protein sequence for inference
21protein_sequence = "MAVPETRPNHTIYINNLNEKIKKDELKKSLHAIFSRFGQILDILVSRSLKMRGQAFVIFKEVSSATNALRSMQGFPFYDKPMRIQYAKTDSDIIAKMKGT" # Replace with your actual sequence
22
23# Tokenize the sequence
24inputs = loaded_tokenizer(protein_sequence, return_tensors="pt", truncation=True, max_length=1024, padding='max_length')
25
26# Run the model
27with torch.no_grad():
28 logits = loaded_model(**inputs).logits
29
30# Get predictions
31tokens = loaded_tokenizer.convert_ids_to_tokens(inputs["input_ids"][0]) # Convert input ids back to tokens
32predictions = torch.argmax(logits, dim=2)
33
34# Define labels
35id2label = {
36 0: "No binding site",
37 1: "Binding site"
38}
39
40# Print the predicted labels for each token
41for token, prediction in zip(tokens, predictions[0].numpy()):
42 if token not in ['<pad>', '<cls>', '<eos>']:
43 print((token, id2label[prediction]))
44