Views
No views yet
esm2_t12_35M_UR50D for predicting post translational modification sites.1 "eval_loss": 0.4661065936088562,
2 "eval_accuracy": 0.9876599555715365,
3 "eval_auc": 0.8673592596422711,
4 "eval_precision": 0.14941997670219148,
5 "eval_recall": 0.7463955099754822
6 "eval_f1": 0.24899413187145658,
7 "eval_mcc": 0.3305508498121041,!pip install transformers -q
!pip install peft -q1from transformers import AutoModelForTokenClassification, AutoTokenizer
2from peft import PeftModel
3import torch
4
5# Path to the saved LoRA model
6model_path = "AmelieSchreiber/esm2_t12_35M_ptm_lora_2100K"
7# ESM2 base model
8base_model_path = "facebook/esm2_t12_35M_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 ptm site",
37 1: "ptm 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]))