File size: 5,061 Bytes
fae1173 | 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 149 150 151 152 153 | import argparse
import os
import sys
from pathlib import Path
ROOT_DIR = Path(__file__).resolve().parents[1]
MODEL_DIR = ROOT_DIR / "model"
if str(MODEL_DIR) not in sys.path:
sys.path.insert(0, str(MODEL_DIR))
def load_esm_model(checkpoint_path):
import esm
import warnings
with warnings.catch_warnings():
warnings.simplefilter('ignore', UserWarning)
return esm.pretrained.load_model_and_alphabet(checkpoint_path)
def get_native_seq(pdbfile, chain):
import util
structure = util.load_structure(pdbfile, chain)
_ , native_seq = util.extract_coords_from_structure(structure)
return native_seq
def write_dms_lib(args):
'''Writes a deep mutational scanning library, including the native/wildtype (wt) of the
indicated target chain in the structure to an output Fasta file'''
from dms_utils import deep_mutational_scan
sequence = get_native_seq(args.pdbfile, args.chain)
Path(args.seqpath).parent.mkdir(parents=True, exist_ok=True)
with open(args.seqpath, 'w') as f:
f.write('>wt\n')
f.write(sequence+'\n')
for pos, wt, mt in deep_mutational_scan(sequence):
assert(sequence[pos] == wt)
mut_seq = sequence[:pos] + mt + sequence[(pos + 1):]
f.write('>' + str(wt) + str(pos+1+args.offset) + str(mt) + '\n')
f.write(mut_seq + '\n')
def get_top_n(args):
import pandas as pd
recs, rec_inds = [], []
scores_df = pd.read_csv(args.outpath).sort_values(by = 'log_likelihood', ascending = False)
for seqid in scores_df['seqid']:
res_ind = seqid[1:-1]
if (rec_inds.count(res_ind) < args.maxrep):
if args.upperbound == None or (int(res_ind) < int(args.upperbound)):
recs.append(seqid)
rec_inds.append(res_ind)
if len(recs) == args.n:
break
print(f'\n Chain {args.chain}')
print(*recs, sep='\n')
def get_model_checkpoint_path(filename):
# Expanding the user's home directory
return os.path.expanduser(f"~/.cache/torch/hub/checkpoints/{filename}")
def main():
parser = argparse.ArgumentParser(
description='Score sequences based on a given structure.'
)
parser.add_argument(
'pdbfile', type=str,
help='input filepath, either .pdb or .cif',
)
parser.add_argument(
'--seqpath', type=str,
help='filepath where fasta of dms library should be saveda',
)
parser.add_argument(
'--outpath', type=str,
help='output filepath for scores of variant sequences',
)
parser.add_argument(
'--chain', type=str,
help='chain id for the chain of interest', default='A',
)
parser.set_defaults(multichain_backbone=True)
parser.add_argument(
'--multichain-backbone', action='store_true',
help='use the backbones of all chains in the input for conditioning'
)
parser.add_argument(
'--singlechain-backbone', dest='multichain_backbone',
action='store_false',
help='use the backbone of only target chain in the input for conditioning'
)
parser.add_argument(
'--order', type=str, default=None,
help='for multichain, option to specify order of chains'
)
parser.add_argument(
'--n', type=int,
help='number of desired predictions to be output',
default=10,
)
parser.add_argument(
'--maxrep', type=int,
help='maximum representation of a single site in the top recommendations \
(eg: maxrep = 1 is a unique set where no wildtype residue is mutated more than once)',
default=1,
)
parser.add_argument(
'--offset', type=int,
help='integer offset for labeling of residue indices encoded in the structure',
default=0,
)
parser.add_argument(
'--upperbound', type=int,
help='only residue positions less than the user-defined upperbound are considered to be recommended for screening \
(but all positions are still conditioned for scoring)',
default=None,
)
parser.add_argument(
"--nogpu", action="store_true",
help="Do not use GPU even if available"
)
args = parser.parse_args()
if args.seqpath is None:
args.seqpath = f'output/{args.pdbfile[:-4]}-chain{args.chain}_dms.fasta'
if args.outpath is None:
args.outpath = f'output/{args.pdbfile[:-4]}-chain{args.chain}_scores.csv'
#write dms library for target chain
write_dms_lib(args)
model_checkpoint_path = get_model_checkpoint_path('esm_if1_20220410.pt')
model, alphabet = load_esm_model(model_checkpoint_path)
model = model.eval()
import score_log_likelihoods
if args.multichain_backbone:
score_log_likelihoods.score_multichain_backbone(model, alphabet, args)
else:
score_log_likelihoods.score_singlechain_backbone(model, alphabet, args)
get_top_n(args)
if __name__ == '__main__':
main()
|