Team Ai
Modelpublic

OneScience-Group/ProteinMPNN

sourceHugging Facemitupdated 2mo agoView on Hugging Face
1likes20downloads
inference.py495 linesDownload Raw Back to scripts
1import argparse2import os3import sys4 5_PROJECT_ROOT = os.path.abspath(os.path.dirname(__file__))6while _PROJECT_ROOT and not os.path.isdir(os.path.join(_PROJECT_ROOT, "model")):7    _PARENT = os.path.dirname(_PROJECT_ROOT)8    if _PARENT == _PROJECT_ROOT:9        break10    _PROJECT_ROOT = _PARENT11_MODEL_ROOT = os.path.join(_PROJECT_ROOT, "model")12_ONESCIENCE_ROOT = os.environ.get("ONESCIENCE_ROOT")13for _path in (_MODEL_ROOT, _PROJECT_ROOT):14    if os.path.exists(_path) and _path not in sys.path:15        sys.path.insert(0, _path)16if _ONESCIENCE_ROOT:17    _ONESCIENCE_SRC = os.path.join(_ONESCIENCE_ROOT, "src")18    for _path in (_ONESCIENCE_SRC, _ONESCIENCE_ROOT):19        if os.path.exists(_path) and _path not in sys.path:20            sys.path.insert(0, _path)21import os.path22 23 24def _resolve_model_folder_path(args):25    if args.path_to_model_weights:26        return os.path.abspath(os.path.normpath(args.path_to_model_weights))27    if args.ca_only and args.use_soluble_model:28        raise ValueError("CA-SolubleMPNN is not available yet")29    if args.ca_only:30        variant = "ca_model_weights"31    elif args.use_soluble_model:32        variant = "soluble_model_weights"33    else:34        variant = "vanilla_model_weights"35    return os.path.join(_PROJECT_ROOT, "weight", variant)36 37 38def main(args):39 40    import json, time, os, sys, glob41    import shutil42    import warnings43    import numpy as np44    import torch45    from torch import optim46    from torch.utils.data import DataLoader47    from torch.utils.data.dataset import random_split, Subset48    import copy49    import torch.nn as nn50    import torch.nn.functional as F51    import random52    import os.path53    import subprocess54    55    from proteinmpnn.protein_mpnn_utils import loss_nll, loss_smoothed, gather_edges, gather_nodes, gather_nodes_t, cat_neighbors_nodes, _scores, _S_to_seq, tied_featurize, parse_PDB, parse_fasta56    from proteinmpnn.protein_mpnn_utils import StructureDataset, StructureDatasetPDB, ProteinMPNN57 58    if args.seed:59        seed=args.seed60    else:61        seed=int(np.random.randint(0, high=999, size=1, dtype=int)[0])62 63    torch.manual_seed(seed)64    random.seed(seed)65    np.random.seed(seed)   66    67    hidden_dim = 12868    num_layers = 3 69  70 71    try:72        model_folder_path = _resolve_model_folder_path(args)73    except ValueError as exc:74        print(f"WARNING: {exc}")75        sys.exit(1)76    if not args.path_to_model_weights:77        if args.ca_only:78            print("Using CA-ProteinMPNN!")79        elif args.use_soluble_model:80            print("Using ProteinMPNN trained on soluble proteins only!")81 82    checkpoint_path = os.path.join(model_folder_path, f'{args.model_name}.pt')83    folder_for_outputs = args.out_folder84    85    NUM_BATCHES = args.num_seq_per_target//args.batch_size86    BATCH_COPIES = args.batch_size87    temperatures = [float(item) for item in args.sampling_temp.split()]88    omit_AAs_list = args.omit_AAs89    alphabet = 'ACDEFGHIKLMNPQRSTVWYX'90    alphabet_dict = dict(zip(alphabet, range(21)))    91    print_all = args.suppress_print == 0 92    omit_AAs_np = np.array([AA in omit_AAs_list for AA in alphabet]).astype(np.float32)93    device = torch.device("cuda:0" if (torch.cuda.is_available()) else "cpu")94    if os.path.isfile(args.chain_id_jsonl):95        with open(args.chain_id_jsonl, 'r') as json_file:96            json_list = list(json_file)97        for json_str in json_list:98            chain_id_dict = json.loads(json_str)99    else:100        chain_id_dict = None101        if print_all:102            print(40*'-')103            print('chain_id_jsonl is NOT loaded')104        105    if os.path.isfile(args.fixed_positions_jsonl):106        with open(args.fixed_positions_jsonl, 'r') as json_file:107            json_list = list(json_file)108        for json_str in json_list:109            fixed_positions_dict = json.loads(json_str)110    else:111        if print_all:112            print(40*'-')113            print('fixed_positions_jsonl is NOT loaded')114        fixed_positions_dict = None115    116    117    if os.path.isfile(args.pssm_jsonl):118        with open(args.pssm_jsonl, 'r') as json_file:119            json_list = list(json_file)120        pssm_dict = {}121        for json_str in json_list:122            pssm_dict.update(json.loads(json_str))123    else:124        if print_all:125            print(40*'-')126            print('pssm_jsonl is NOT loaded')127        pssm_dict = None128    129    130    if os.path.isfile(args.omit_AA_jsonl):131        with open(args.omit_AA_jsonl, 'r') as json_file:132            json_list = list(json_file)133        for json_str in json_list:134            omit_AA_dict = json.loads(json_str)135    else:136        if print_all:137            print(40*'-')138            print('omit_AA_jsonl is NOT loaded')139        omit_AA_dict = None140    141    142    if os.path.isfile(args.bias_AA_jsonl):143        with open(args.bias_AA_jsonl, 'r') as json_file:144            json_list = list(json_file)145        for json_str in json_list:146            bias_AA_dict = json.loads(json_str)147    else:148        if print_all:149            print(40*'-')150            print('bias_AA_jsonl is NOT loaded')151        bias_AA_dict = None152    153    154    if os.path.isfile(args.tied_positions_jsonl):155        with open(args.tied_positions_jsonl, 'r') as json_file:156            json_list = list(json_file)157        for json_str in json_list:158            tied_positions_dict = json.loads(json_str)159    else:160        if print_all:161            print(40*'-')162            print('tied_positions_jsonl is NOT loaded')163        tied_positions_dict = None164 165    166    if os.path.isfile(args.bias_by_res_jsonl):167        with open(args.bias_by_res_jsonl, 'r') as json_file:168            json_list = list(json_file)169    170        for json_str in json_list:171            bias_by_res_dict = json.loads(json_str)172        if print_all:173            print('bias by residue dictionary is loaded')174    else:175        if print_all:176            print(40*'-')177            print('bias by residue dictionary is not loaded, or not provided')178        bias_by_res_dict = None179   180 181    if print_all: 182        print(40*'-')183    bias_AAs_np = np.zeros(len(alphabet))184    if bias_AA_dict:185            for n, AA in enumerate(alphabet):186                    if AA in list(bias_AA_dict.keys()):187                            bias_AAs_np[n] = bias_AA_dict[AA]188    189    if args.pdb_path:190        pdb_dict_list = parse_PDB(args.pdb_path, ca_only=args.ca_only)191        dataset_valid = StructureDatasetPDB(pdb_dict_list, truncate=None, max_length=args.max_length)192        all_chain_list = [item[-1:] for item in list(pdb_dict_list[0]) if item[:9]=='seq_chain'] #['A','B', 'C',...]193        if args.pdb_path_chains:194            designed_chain_list = [str(item) for item in args.pdb_path_chains.split()]195        else:196            designed_chain_list = all_chain_list197        fixed_chain_list = [letter for letter in all_chain_list if letter not in designed_chain_list]198        chain_id_dict = {}199        chain_id_dict[pdb_dict_list[0]['name']]= (designed_chain_list, fixed_chain_list)200    else:201        dataset_valid = StructureDataset(args.jsonl_path, truncate=None, max_length=args.max_length, verbose=print_all)202 203    checkpoint = torch.load(checkpoint_path, map_location=device) 204    noise_level_print = checkpoint['noise_level']205    model = ProteinMPNN(ca_only=args.ca_only, num_letters=21, node_features=hidden_dim, edge_features=hidden_dim, hidden_dim=hidden_dim, num_encoder_layers=num_layers, num_decoder_layers=num_layers, augment_eps=args.backbone_noise, k_neighbors=checkpoint['num_edges'])206    model.to(device)207    model.load_state_dict(checkpoint['model_state_dict'])208    model.eval()209 210    if print_all:211        print(40*'-')212        print('Number of edges:', checkpoint['num_edges'])213        print(f'Training noise level: {noise_level_print}A')214 215    # Build paths for experiment216    base_folder = folder_for_outputs217    if base_folder[-1] != '/':218        base_folder = base_folder + '/'219    if not os.path.exists(base_folder):220        os.makedirs(base_folder)221    222    if not os.path.exists(base_folder + 'seqs'):223        os.makedirs(base_folder + 'seqs')224    225    if args.save_score:226        if not os.path.exists(base_folder + 'scores'):227            os.makedirs(base_folder + 'scores')228 229    if args.score_only:230        if not os.path.exists(base_folder + 'score_only'):231            os.makedirs(base_folder + 'score_only')232   233 234    if args.conditional_probs_only:235        if not os.path.exists(base_folder + 'conditional_probs_only'):236            os.makedirs(base_folder + 'conditional_probs_only')237 238    if args.unconditional_probs_only:239        if not os.path.exists(base_folder + 'unconditional_probs_only'):240            os.makedirs(base_folder + 'unconditional_probs_only')241 242    if args.save_probs:243        if not os.path.exists(base_folder + 'probs'):244            os.makedirs(base_folder + 'probs') 245    246    # Timing247    start_time = time.time()248    total_residues = 0249    protein_list = []250    total_step = 0251    # Validation epoch252    with torch.no_grad():253        test_sum, test_weights = 0., 0.254        for ix, protein in enumerate(dataset_valid):255            score_list = []256            global_score_list = []257            all_probs_list = []258            all_log_probs_list = []259            S_sample_list = []260            batch_clones = [copy.deepcopy(protein) for i in range(BATCH_COPIES)]261            X, S, mask, lengths, chain_M, chain_encoding_all, chain_list_list, visible_list_list, masked_list_list, masked_chain_length_list_list, chain_M_pos, omit_AA_mask, residue_idx, dihedral_mask, tied_pos_list_of_lists_list, pssm_coef, pssm_bias, pssm_log_odds_all, bias_by_res_all, tied_beta = tied_featurize(batch_clones, device, chain_id_dict, fixed_positions_dict, omit_AA_dict, tied_positions_dict, pssm_dict, bias_by_res_dict, ca_only=args.ca_only)262            pssm_log_odds_mask = (pssm_log_odds_all > args.pssm_threshold).float() #1.0 for true, 0.0 for false263            name_ = batch_clones[0]['name']264            if args.score_only:265                loop_c = 0 266                if args.path_to_fasta:267                    fasta_names, fasta_seqs = parse_fasta(args.path_to_fasta, omit=["/"])268                    loop_c = len(fasta_seqs)269                for fc in range(1+loop_c):270                    if fc == 0:271                        structure_sequence_score_file = base_folder + '/score_only/' + batch_clones[0]['name'] + f'_pdb'272                    else:273                        structure_sequence_score_file = base_folder + '/score_only/' + batch_clones[0]['name'] + f'_fasta_{fc}'274                    native_score_list = []275                    global_native_score_list = []276                    if fc > 0:277                        input_seq_length = len(fasta_seqs[fc-1])278                        S_input = torch.tensor([alphabet_dict[AA] for AA in fasta_seqs[fc-1]], device=device)[None,:].repeat(X.shape[0], 1)279                        S[:,:input_seq_length] = S_input #assumes that S and S_input are alphabetically sorted for masked_chains280                    for j in range(NUM_BATCHES):281                        randn_1 = torch.randn(chain_M.shape, device=X.device)282                        log_probs = model(X, S, mask, chain_M*chain_M_pos, residue_idx, chain_encoding_all, randn_1)283                        mask_for_loss = mask*chain_M*chain_M_pos284                        scores = _scores(S, log_probs, mask_for_loss)285                        native_score = scores.cpu().data.numpy()286                        native_score_list.append(native_score)287                        global_scores = _scores(S, log_probs, mask)288                        global_native_score = global_scores.cpu().data.numpy()289                        global_native_score_list.append(global_native_score)290                    native_score = np.concatenate(native_score_list, 0)291                    global_native_score = np.concatenate(global_native_score_list, 0)292                    ns_mean = native_score.mean()293                    ns_mean_print = np.format_float_positional(np.float32(ns_mean), unique=False, precision=4)294                    ns_std = native_score.std()295                    ns_std_print = np.format_float_positional(np.float32(ns_std), unique=False, precision=4)296 297                    global_ns_mean = global_native_score.mean()298                    global_ns_mean_print = np.format_float_positional(np.float32(global_ns_mean), unique=False, precision=4)299                    global_ns_std = global_native_score.std()300                    global_ns_std_print = np.format_float_positional(np.float32(global_ns_std), unique=False, precision=4)301 302                    ns_sample_size = native_score.shape[0]303                    seq_str = _S_to_seq(S[0,], chain_M[0,])304                    np.savez(structure_sequence_score_file, score=native_score, global_score=global_native_score, S=S[0,].cpu().numpy(), seq_str=seq_str)305                    if print_all:306                        if fc == 0:307                            print(f'Score for {name_} from PDB, mean: {ns_mean_print}, std: {ns_std_print}, sample size: {ns_sample_size},  global score, mean: {global_ns_mean_print}, std: {global_ns_std_print}, sample size: {ns_sample_size}')308                        else:309                            print(f'Score for {name_}_{fc} from FASTA, mean: {ns_mean_print}, std: {ns_std_print}, sample size: {ns_sample_size},  global score, mean: {global_ns_mean_print}, std: {global_ns_std_print}, sample size: {ns_sample_size}')310            elif args.conditional_probs_only:311                if print_all:312                    print(f'Calculating conditional probabilities for {name_}')313                conditional_probs_only_file = base_folder + '/conditional_probs_only/' + batch_clones[0]['name']314                log_conditional_probs_list = []315                for j in range(NUM_BATCHES):316                    randn_1 = torch.randn(chain_M.shape, device=X.device)317                    log_conditional_probs = model.conditional_probs(X, S, mask, chain_M*chain_M_pos, residue_idx, chain_encoding_all, randn_1, args.conditional_probs_only_backbone)318                    log_conditional_probs_list.append(log_conditional_probs.cpu().numpy())319                concat_log_p = np.concatenate(log_conditional_probs_list, 0) #[B, L, 21]320                mask_out = (chain_M*chain_M_pos*mask)[0,].cpu().numpy()321                np.savez(conditional_probs_only_file, log_p=concat_log_p, S=S[0,].cpu().numpy(), mask=mask[0,].cpu().numpy(), design_mask=mask_out)322            elif args.unconditional_probs_only:323                if print_all:324                    print(f'Calculating sequence unconditional probabilities for {name_}')325                unconditional_probs_only_file = base_folder + '/unconditional_probs_only/' + batch_clones[0]['name']326                log_unconditional_probs_list = []327                for j in range(NUM_BATCHES):328                    log_unconditional_probs = model.unconditional_probs(X, mask, residue_idx, chain_encoding_all)329                    log_unconditional_probs_list.append(log_unconditional_probs.cpu().numpy())330                concat_log_p = np.concatenate(log_unconditional_probs_list, 0) #[B, L, 21]331                mask_out = (chain_M*chain_M_pos*mask)[0,].cpu().numpy()332                np.savez(unconditional_probs_only_file, log_p=concat_log_p, S=S[0,].cpu().numpy(), mask=mask[0,].cpu().numpy(), design_mask=mask_out)333            else:334                randn_1 = torch.randn(chain_M.shape, device=X.device)335                log_probs = model(X, S, mask, chain_M*chain_M_pos, residue_idx, chain_encoding_all, randn_1)336                mask_for_loss = mask*chain_M*chain_M_pos337                scores = _scores(S, log_probs, mask_for_loss) #score only the redesigned part338                native_score = scores.cpu().data.numpy()339                global_scores = _scores(S, log_probs, mask) #score the whole structure-sequence340                global_native_score = global_scores.cpu().data.numpy()341                # Generate some sequences342                ali_file = base_folder + '/seqs/' + batch_clones[0]['name'] + '.fa'343                score_file = base_folder + '/scores/' + batch_clones[0]['name'] + '.npz'344                probs_file = base_folder + '/probs/' + batch_clones[0]['name'] + '.npz'345                if print_all:346                    print(f'Generating sequences for: {name_}')347                t0 = time.time()348                with open(ali_file, 'w') as f:349                    for temp in temperatures:350                        for j in range(NUM_BATCHES):351                            randn_2 = torch.randn(chain_M.shape, device=X.device)352                            if tied_positions_dict == None:353                                sample_dict = model.sample(X, randn_2, S, chain_M, chain_encoding_all, residue_idx, mask=mask, temperature=temp, omit_AAs_np=omit_AAs_np, bias_AAs_np=bias_AAs_np, chain_M_pos=chain_M_pos, omit_AA_mask=omit_AA_mask, pssm_coef=pssm_coef, pssm_bias=pssm_bias, pssm_multi=args.pssm_multi, pssm_log_odds_flag=bool(args.pssm_log_odds_flag), pssm_log_odds_mask=pssm_log_odds_mask, pssm_bias_flag=bool(args.pssm_bias_flag), bias_by_res=bias_by_res_all)354                                S_sample = sample_dict["S"] 355                            else:356                                sample_dict = model.tied_sample(X, randn_2, S, chain_M, chain_encoding_all, residue_idx, mask=mask, temperature=temp, omit_AAs_np=omit_AAs_np, bias_AAs_np=bias_AAs_np, chain_M_pos=chain_M_pos, omit_AA_mask=omit_AA_mask, pssm_coef=pssm_coef, pssm_bias=pssm_bias, pssm_multi=args.pssm_multi, pssm_log_odds_flag=bool(args.pssm_log_odds_flag), pssm_log_odds_mask=pssm_log_odds_mask, pssm_bias_flag=bool(args.pssm_bias_flag), tied_pos=tied_pos_list_of_lists_list[0], tied_beta=tied_beta, bias_by_res=bias_by_res_all)357                            # Compute scores358                                S_sample = sample_dict["S"]359                            log_probs = model(X, S_sample, mask, chain_M*chain_M_pos, residue_idx, chain_encoding_all, randn_2, use_input_decoding_order=True, decoding_order=sample_dict["decoding_order"])360                            mask_for_loss = mask*chain_M*chain_M_pos361                            scores = _scores(S_sample, log_probs, mask_for_loss)362                            scores = scores.cpu().data.numpy()363                            364                            global_scores = _scores(S_sample, log_probs, mask) #score the whole structure-sequence365                            global_scores = global_scores.cpu().data.numpy()366                            367                            all_probs_list.append(sample_dict["probs"].cpu().data.numpy())368                            all_log_probs_list.append(log_probs.cpu().data.numpy())369                            S_sample_list.append(S_sample.cpu().data.numpy())370                            for b_ix in range(BATCH_COPIES):371                                masked_chain_length_list = masked_chain_length_list_list[b_ix]372                                masked_list = masked_list_list[b_ix]373                                seq_recovery_rate = torch.sum(torch.sum(torch.nn.functional.one_hot(S[b_ix], 21)*torch.nn.functional.one_hot(S_sample[b_ix], 21),axis=-1)*mask_for_loss[b_ix])/torch.sum(mask_for_loss[b_ix])374                                seq = _S_to_seq(S_sample[b_ix], chain_M[b_ix])375                                score = scores[b_ix]376                                score_list.append(score)377                                global_score = global_scores[b_ix]378                                global_score_list.append(global_score)379                                native_seq = _S_to_seq(S[b_ix], chain_M[b_ix])380                                if b_ix == 0 and j==0 and temp==temperatures[0]:381                                    start = 0382                                    end = 0383                                    list_of_AAs = []384                                    for mask_l in masked_chain_length_list:385                                        end += mask_l386                                        list_of_AAs.append(native_seq[start:end])387                                        start = end388                                    native_seq = "".join(list(np.array(list_of_AAs)[np.argsort(masked_list)]))389                                    l0 = 0390                                    for mc_length in list(np.array(masked_chain_length_list)[np.argsort(masked_list)])[:-1]:391                                        l0 += mc_length392                                        native_seq = native_seq[:l0] + '/' + native_seq[l0:]393                                        l0 += 1394                                    sorted_masked_chain_letters = np.argsort(masked_list_list[0])395                                    print_masked_chains = [masked_list_list[0][i] for i in sorted_masked_chain_letters]396                                    sorted_visible_chain_letters = np.argsort(visible_list_list[0])397                                    print_visible_chains = [visible_list_list[0][i] for i in sorted_visible_chain_letters]398                                    native_score_print = np.format_float_positional(np.float32(native_score.mean()), unique=False, precision=4)399                                    global_native_score_print = np.format_float_positional(np.float32(global_native_score.mean()), unique=False, precision=4)400                                    script_dir = os.path.dirname(os.path.realpath(__file__))401                                    try:402                                        commit_str = subprocess.check_output(f'git --git-dir {script_dir}/.git rev-parse HEAD', shell=True, stderr=subprocess.DEVNULL).decode().strip()403                                    except subprocess.CalledProcessError:404                                        commit_str = 'unknown'405                                    if args.ca_only:406                                        print_model_name = 'CA_model_name'407                                    else:408                                        print_model_name = 'model_name'409                                    f.write('>{}, score={}, global_score={}, fixed_chains={}, designed_chains={}, {}={}, git_hash={}, seed={}\n{}\n'.format(name_, native_score_print, global_native_score_print, print_visible_chains, print_masked_chains, print_model_name, args.model_name, commit_str, seed, native_seq)) #write the native sequence410                                start = 0411                                end = 0412                                list_of_AAs = []413                                for mask_l in masked_chain_length_list:414                                    end += mask_l415                                    list_of_AAs.append(seq[start:end])416                                    start = end417    418                                seq = "".join(list(np.array(list_of_AAs)[np.argsort(masked_list)]))419                                l0 = 0420                                for mc_length in list(np.array(masked_chain_length_list)[np.argsort(masked_list)])[:-1]:421                                    l0 += mc_length422                                    seq = seq[:l0] + '/' + seq[l0:]423                                    l0 += 1424                                score_print = np.format_float_positional(np.float32(score), unique=False, precision=4)425                                global_score_print = np.format_float_positional(np.float32(global_score), unique=False, precision=4)426                                seq_rec_print = np.format_float_positional(np.float32(seq_recovery_rate.detach().cpu().numpy()), unique=False, precision=4)427                                sample_number = j*BATCH_COPIES+b_ix+1428                                f.write('>T={}, sample={}, score={}, global_score={}, seq_recovery={}\n{}\n'.format(temp,sample_number,score_print,global_score_print,seq_rec_print,seq)) #write generated sequence429                if args.save_score:430                    np.savez(score_file, score=np.array(score_list, np.float32), global_score=np.array(global_score_list, np.float32))431                if args.save_probs:432                    all_probs_concat = np.concatenate(all_probs_list)433                    all_log_probs_concat = np.concatenate(all_log_probs_list)434                    S_sample_concat = np.concatenate(S_sample_list)435                    np.savez(probs_file, probs=np.array(all_probs_concat, np.float32), log_probs=np.array(all_log_probs_concat, np.float32), S=np.array(S_sample_concat, np.int32), mask=mask_for_loss.cpu().data.numpy(), chain_order=chain_list_list)436                t1 = time.time()437                dt = round(float(t1-t0), 4)438                num_seqs = len(temperatures)*NUM_BATCHES*BATCH_COPIES439                total_length = X.shape[1]440                if print_all:441                    print(f'{num_seqs} sequences of length {total_length} generated in {dt} seconds')442   443if __name__ == "__main__":444    argparser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter)445 446    argparser.add_argument("--suppress_print", type=int, default=0, help="0 for False, 1 for True")447 448  449    argparser.add_argument("--ca_only", action="store_true", default=False, help="Parse CA-only structures and use CA-only models (default: false)")   450    argparser.add_argument("--path_to_model_weights", type=str, default="", help="Path to model weights folder;") 451    argparser.add_argument("--model_name", type=str, default="v_48_020", help="ProteinMPNN model name: v_48_002, v_48_010, v_48_020, v_48_030; v_48_010=version with 48 edges 0.10A noise")452    argparser.add_argument("--use_soluble_model", action="store_true", default=False, help="Flag to load ProteinMPNN weights trained on soluble proteins only.")453 454 455    argparser.add_argument("--seed", type=int, default=0, help="If set to 0 then a random seed will be picked;")456 457    argparser.add_argument("--save_score", type=int, default=0, help="0 for False, 1 for True; save score=-log_prob to npy files")458    argparser.add_argument("--save_probs", type=int, default=0, help="0 for False, 1 for True; save MPNN predicted probabilites per position")459 460    argparser.add_argument("--score_only", type=int, default=0, help="0 for False, 1 for True; score input backbone-sequence pairs")461    argparser.add_argument("--path_to_fasta", type=str, default="", help="score provided input sequence in a fasta format; e.g. GGGGGG/PPPPS/WWW for chains A, B, C sorted alphabetically and separated by /")462 463 464    argparser.add_argument("--conditional_probs_only", type=int, default=0, help="0 for False, 1 for True; output conditional probabilities p(s_i given the rest of the sequence and backbone)")    465    argparser.add_argument("--conditional_probs_only_backbone", type=int, default=0, help="0 for False, 1 for True; if true output conditional probabilities p(s_i given backbone)") 466    argparser.add_argument("--unconditional_probs_only", type=int, default=0, help="0 for False, 1 for True; output unconditional probabilities p(s_i given backbone) in one forward pass")   467 468    argparser.add_argument("--backbone_noise", type=float, default=0.00, help="Standard deviation of Gaussian noise to add to backbone atoms")469    argparser.add_argument("--num_seq_per_target", type=int, default=1, help="Number of sequences to generate per target")470    argparser.add_argument("--batch_size", type=int, default=1, help="Batch size; can set higher for titan, quadro GPUs, reduce this if running out of GPU memory")471    argparser.add_argument("--max_length", type=int, default=200000, help="Max sequence length")472    argparser.add_argument("--sampling_temp", type=str, default="0.1", help="A string of temperatures, 0.2 0.25 0.5. Sampling temperature for amino acids. Suggested values 0.1, 0.15, 0.2, 0.25, 0.3. Higher values will lead to more diversity.")473    474    argparser.add_argument("--out_folder", type=str, help="Path to a folder to output sequences, e.g. /home/out/")475    argparser.add_argument("--pdb_path", type=str, default='', help="Path to a single PDB to be designed")476    argparser.add_argument("--pdb_path_chains", type=str, default='', help="Define which chains need to be designed for a single PDB ")477    argparser.add_argument("--jsonl_path", type=str, help="Path to a folder with parsed pdb into jsonl")478    argparser.add_argument("--chain_id_jsonl",type=str, default='', help="Path to a dictionary specifying which chains need to be designed and which ones are fixed, if not specied all chains will be designed.")479    argparser.add_argument("--fixed_positions_jsonl", type=str, default='', help="Path to a dictionary with fixed positions")480    argparser.add_argument("--omit_AAs", type=list, default='X', help="Specify which amino acids should be omitted in the generated sequence, e.g. 'AC' would omit alanine and cystine.")481    argparser.add_argument("--bias_AA_jsonl", type=str, default='', help="Path to a dictionary which specifies AA composion bias if neededi, e.g. {A: -1.1, F: 0.7} would make A less likely and F more likely.")482   483    argparser.add_argument("--bias_by_res_jsonl", default='', help="Path to dictionary with per position bias.") 484    argparser.add_argument("--omit_AA_jsonl", type=str, default='', help="Path to a dictionary which specifies which amino acids need to be omited from design at specific chain indices")485    argparser.add_argument("--pssm_jsonl", type=str, default='', help="Path to a dictionary with pssm")486    argparser.add_argument("--pssm_multi", type=float, default=0.0, help="A value between [0.0, 1.0], 0.0 means do not use pssm, 1.0 ignore MPNN predictions")487    argparser.add_argument("--pssm_threshold", type=float, default=0.0, help="A value between -inf + inf to restric per position AAs")488    argparser.add_argument("--pssm_log_odds_flag", type=int, default=0, help="0 for False, 1 for True")489    argparser.add_argument("--pssm_bias_flag", type=int, default=0, help="0 for False, 1 for True")490    491    argparser.add_argument("--tied_positions_jsonl", type=str, default='', help="Path to a dictionary with tied positions")492    493    args = argparser.parse_args()    494    main(args)   495