SaProt / model /saprot /esm_mutation_model.py
anzhi2710gmailcom's picture
Upload folder using huggingface_hub
3b99abb verified
Raw
History Blame Contribute Delete
11.5 kB
import os
import torch
import copy
import json
import torchmetrics
import torch.distributed as dist
from scripts.utils.constants import aa_set, aa_list
from ..model_interface import register_model
from .base import SaprotBaseModel
@register_model
class EsmMutationModel(SaprotBaseModel):
def __init__(self,
use_bias_feature: bool = False,
MSA_log_path: str = None,
log_clinvar: bool = False,
log_dir: str = None,
**kwargs):
"""
Args:
use_bias_feature: Whether to use structure information as bias feature
MSA_log_path: If not None, the model will load MSA log from this path (following Tranception paper)
log_clinvar: If True, the model will log the predicted evolutionary indices for ClinVar variants
log_dir: If log_clinvar is True, the model will save the predicted evolutionary indices for ClinVar variants
**kwargs: other arguments for SaprotBaseModel
"""
self.use_bias_feature = use_bias_feature
self.MSA_log_path = MSA_log_path
self.MSA_info_dict = {}
if MSA_log_path:
with open(MSA_log_path, "r") as r:
for line in r:
data = json.loads(line)
data["MSA_log_prior"] = torch.tensor(data["MSA_log_prior"])
self.MSA_info_dict[data["DMS_id"]] = data
self.log_clinvar = log_clinvar
self.log_dir = log_dir
if log_clinvar:
self.mut_info_list = []
super().__init__(task="lm", **kwargs)
def initialize_metrics(self, stage):
return {f"{stage}_spearman": torchmetrics.SpearmanCorrCoef()}
def forward(self, wild_type, seqs, mut_info, structure_content, structure_type, plddt, struc_seq):
if self.use_bias_feature and getattr(self, "coords", None) is None:
structure_type = "cif" if structure_type == "mmcif" else structure_type
tmp_path = f"EsmMutationModel_{self.global_rank}.{structure_type}"
with open(tmp_path, "w") as f:
f.write(structure_content)
self.coords = parse_structure(tmp_path, ["A"])["A"]["coords"]
os.remove(tmp_path)
ins_seqs = []
ori_seqs = []
mut_data = []
# The running bottleneck is two forward passes of the model to deal with insertion
# Therefore we only forward pass the model twice for sequences with insertion
ins_dict = {}
for i, (seq, info) in enumerate(zip(seqs, mut_info)):
# We adopt the same strategy for esm2 model as in esm2 inverse folding paper
ori_seq = [aa for aa in wild_type]
ins_seq = copy.deepcopy(ori_seq)
tmp_data = []
ins_num = 0
# To indicate whether there is insertion in the sequence
flag = False
for single in info.split(":"):
# Mask the amino acid where the mutation happens
# -1 is added because the index starts from 1 and we need to convert it to 0
if single[0] in aa_set:
ori_aa, pos, mut_aa = single[0], int(single[1:-1]), single[-1]
ori_seq[pos - ins_num - 1] = self.tokenizer.mask_token
ins_seq[pos - 1] = self.tokenizer.mask_token
tmp_data.append((ori_aa, pos - ins_num, mut_aa, pos))
# For insertion
else:
ins_dict[i] = len(ins_dict)
flag = True
ins_num += 1
ins_pos = int(single[:-1])
ins_seq = ins_seq[:ins_pos - 1] + [self.tokenizer.mask_token] + ins_seq[ins_pos - 1:]
if flag:
ins_seqs.append(" ".join(ins_seq))
ori_seqs.append(" ".join(ori_seq))
mut_data.append(tmp_data)
device = self.device
if len(ins_seqs) > 0:
ins_inputs = self.tokenizer.batch_encode_plus(ins_seqs, return_tensors="pt", padding=True)
ins_inputs = {k: v.to(device) for k, v in ins_inputs.items()}
if self.use_bias_feature:
coords = [copy.deepcopy(self.coords) for _ in range(len(seqs))]
self.add_bias_feature(ins_inputs, coords)
ins_outputs = self.model(**ins_inputs)
ins_probs = ins_outputs['logits'].softmax(dim=-1)
ori_inputs = self.tokenizer.batch_encode_plus(ori_seqs, return_tensors="pt", padding=True)
ori_inputs = {k: v.to(device) for k, v in ori_inputs.items()}
if self.use_bias_feature:
coords = [copy.deepcopy(self.coords) for _ in range(len(seqs))]
self.add_bias_feature(ori_inputs, coords)
ori_outputs = self.model(**ori_inputs)
ori_probs = ori_outputs['logits'].softmax(dim=-1)
if self.MSA_log_path is not None:
aa2id = {"A": 5, "C": 6, "D": 7, "E": 8, "F": 9, "G": 10, "H": 11, "I": 12, "K": 13, "L": 14, "M": 15,
"N": 16, "P": 17, "Q": 18, "R": 19, "S": 20, "T": 21, "V": 22, "W": 23, "Y": 24}
DMS_id = os.path.basename(self.trainer.datamodule.test_lmdb)
MSA_info = self.MSA_info_dict[DMS_id]
MSA_log_prior = MSA_info["MSA_log_prior"].to(device)
st, ed = MSA_info["MSA_start"], MSA_info["MSA_end"]
preds = []
for i, data_list in enumerate(mut_data):
pred = 0
for data in data_list:
ori_aa, ori_pos, mut_aa, ins_pos = data
ori_prob = ori_probs[i, ori_pos, self.tokenizer.convert_tokens_to_ids(ori_aa)]
if i in ins_dict:
mut_prob = ins_probs[ins_dict[i], ins_pos, self.tokenizer.convert_tokens_to_ids(mut_aa)]
else:
mut_prob = ori_probs[i, ins_pos, self.tokenizer.convert_tokens_to_ids(mut_aa)]
# Add MSA info if available
if self.MSA_log_path is not None and st <= ori_pos -1 < ed:
ori_msa_prob = MSA_log_prior[ori_pos - 1 - st, aa2id[ori_aa]]
mut_msa_prob = MSA_log_prior[ori_pos - 1 - st, aa2id[mut_aa]]
pred += 0.4 * torch.log(mut_prob / ori_prob) + 0.6 * (mut_msa_prob - ori_msa_prob)
else:
# compute zero-shot score
pred += torch.log(mut_prob / ori_prob)
preds.append(pred)
if self.log_clinvar:
self.mut_info_list.append((mut_info, -torch.tensor(preds)))
return torch.tensor(preds).to(ori_probs)
def loss_func(self, stage, outputs, labels):
fitness = labels['labels']
self.test_spearman(outputs, fitness)
def on_test_epoch_end(self):
spearman = self.test_spearman.compute()
self.reset_metrics("test")
self.log("spearman", spearman)
if self.use_bias_feature:
self.coords = None
if self.log_clinvar:
# Get dataset name
name = os.path.basename(self.trainer.datamodule.test_lmdb)
device_rank = dist.get_rank()
log_path = f"{self.log_dir}/{name}_{device_rank}.csv"
with open(log_path, "w") as w:
w.write("protein_name,mutations,evol_indices\n")
for mut_info, preds in self.mut_info_list:
for mut, pred in zip(mut_info, preds):
w.write(f"{name},{mut},{pred}\n")
self.mut_info_list = []
def predict_mut(self, seq: str, mut_info: str) -> float:
"""
Predict the mutational effect of a given mutation
Args:
seq: The wild type sequence
mut_info: The mutation information in the format of "A123B", where A is the original amino acid, 123 is the
position and B is the mutated amino acid. If multiple mutations are provided, they should be
separated by colon, e.g. "A123B:C124D".
Returns:
The predicted mutational effect
"""
tokens = self.tokenizer.tokenize(seq)
for single in mut_info.split(":"):
pos = int(single[1:-1])
tokens[pos - 1] = self.tokenizer.mask_token
mask_seq = " ".join(tokens)
inputs = self.tokenizer(mask_seq, return_tensors="pt")
inputs = {k: v.to(self.device) for k, v in inputs.items()}
with torch.no_grad():
outputs = self.model(**inputs)
logits = outputs.logits
probs = logits.softmax(dim=-1)
score = 0
for single in mut_info.split(":"):
ori_aa, pos, mut_aa = single[0], int(single[1:-1]), single[-1]
ori_prob = probs[0, pos, self.tokenizer.convert_tokens_to_ids(ori_aa)]
mut_prob = probs[0, pos, self.tokenizer.convert_tokens_to_ids(mut_aa)]
score += torch.log(mut_prob / ori_prob)
return score
def predict_pos_mut(self, seq: str, pos: int) -> dict:
"""
Predict the mutational effect of mutations at a given position
Args:
seq: The wild type sequence
pos: The position of the mutation
Returns:
The predicted mutational effect
"""
tokens = self.tokenizer.tokenize(seq)
ori_aa = tokens[pos - 1][0]
tokens[pos - 1] = self.tokenizer.mask_token
mask_seq = " ".join(tokens)
inputs = self.tokenizer(mask_seq, return_tensors="pt")
inputs = {k: v.to(self.device) for k, v in inputs.items()}
with torch.no_grad():
outputs = self.model(**inputs)
logits = outputs.logits
probs = logits.softmax(dim=-1)[0, pos]
scores = {}
ori_prob = probs[self.tokenizer.convert_tokens_to_ids(ori_aa)]
for mut_aa in aa_list:
mut_prob = probs[self.tokenizer.convert_tokens_to_ids(mut_aa)]
score = torch.log(mut_prob / ori_prob).item()
scores[f"{ori_aa}{pos}{mut_aa}"] = score
return scores
def predict_pos_prob(self, seq: str, pos: int) -> dict:
"""
Predict the probability of all amino acids at a given position
Args:
seq: The wild type sequence
pos: The position of the mutation
Returns:
The predicted probability of all amino acids
"""
tokens = self.tokenizer.tokenize(seq)
tokens[pos - 1] = self.tokenizer.mask_token
mask_seq = " ".join(tokens)
inputs = self.tokenizer(mask_seq, return_tensors="pt")
inputs = {k: v.to(self.device) for k, v in inputs.items()}
with torch.no_grad():
outputs = self.model(**inputs)
logits = outputs.logits
probs = logits.softmax(dim=-1)[0, pos]
scores = {}
for aa in aa_list:
prob = probs[self.tokenizer.convert_tokens_to_ids(aa)]
scores[aa] = prob.item()
return scores