Team Ai
Modelpublic

OneScience-Group/ProteinMPNN

sourceHugging Facemitupdated 2mo agoView on Hugging Face
1likes24downloads
training.py271 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 23def main(args):24    import json, time, os, sys, glob25    import shutil26    import warnings27    import numpy as np28    import torch29    from torch import optim30    from torch.utils.data import DataLoader31    import queue32    import copy33    import torch.nn as nn34    import torch.nn.functional as F35    import random36    import os.path37    import subprocess38    from concurrent.futures import ProcessPoolExecutor    39    from proteinmpnn.utils import worker_init_fn, get_pdbs, loader_pdb, build_training_clusters, PDB_dataset, StructureDataset, StructureLoader40    from proteinmpnn.model_utils import featurize, loss_smoothed, loss_nll, get_std_opt, ProteinMPNN41 42    scaler = torch.cuda.amp.GradScaler()43     44    device = torch.device("cuda:0" if (torch.cuda.is_available()) else "cpu")45 46    base_folder = time.strftime(args.path_for_outputs, time.localtime())47 48    if base_folder[-1] != '/':49        base_folder += '/'50    if not os.path.exists(base_folder):51        os.makedirs(base_folder)52    subfolders = ['model_weights']53    for subfolder in subfolders:54        if not os.path.exists(base_folder + subfolder):55            os.makedirs(base_folder + subfolder)56 57    PATH = args.previous_checkpoint58 59    logfile = base_folder + 'log.txt'60    if not PATH:61        with open(logfile, 'w') as f:62            f.write('Epoch\tTrain\tValidation\n')63 64    data_path = args.path_for_training_data65    params = {66        "LIST"    : f"{data_path}/list.csv", 67        "VAL"     : f"{data_path}/valid_clusters.txt",68        "TEST"    : f"{data_path}/test_clusters.txt",69        "DIR"     : f"{data_path}",70        "DATCUT"  : "2030-Jan-01",71        "RESCUT"  : args.rescut, #resolution cutoff for PDBs72        "HOMO"    : 0.70 #min seq.id. to detect homo chains73    }74 75 76    LOAD_PARAM = {'batch_size': 1,77                  'shuffle': True,78                  'pin_memory':False,79                  'num_workers': 4}80 81   82    if args.debug:83        args.num_examples_per_epoch = 5084        args.max_protein_length = 100085        args.batch_size = 100086 87    train, valid, test = build_training_clusters(params, args.debug)88     89    train_set = PDB_dataset(list(train.keys()), loader_pdb, train, params)90    train_loader = torch.utils.data.DataLoader(train_set, worker_init_fn=worker_init_fn, **LOAD_PARAM)91    valid_set = PDB_dataset(list(valid.keys()), loader_pdb, valid, params)92    valid_loader = torch.utils.data.DataLoader(valid_set, worker_init_fn=worker_init_fn, **LOAD_PARAM)93 94 95    model = ProteinMPNN(node_features=args.hidden_dim, 96                        edge_features=args.hidden_dim, 97                        hidden_dim=args.hidden_dim, 98                        num_encoder_layers=args.num_encoder_layers, 99                        num_decoder_layers=args.num_encoder_layers, 100                        k_neighbors=args.num_neighbors, 101                        dropout=args.dropout, 102                        augment_eps=args.backbone_noise)103    model.to(device)104 105 106    if PATH:107        checkpoint = torch.load(PATH)108        total_step = checkpoint['step'] #write total_step from the checkpoint109        epoch = checkpoint['epoch'] #write epoch from the checkpoint110        model.load_state_dict(checkpoint['model_state_dict'])111    else:112        total_step = 0113        epoch = 0114 115    optimizer = get_std_opt(model.parameters(), args.hidden_dim, total_step)116 117 118    if PATH:119        optimizer.optimizer.load_state_dict(checkpoint['optimizer_state_dict'])120 121 122    with ProcessPoolExecutor(max_workers=12) as executor:123        q = queue.Queue(maxsize=3)124        p = queue.Queue(maxsize=3)125        for i in range(3):126            q.put_nowait(executor.submit(get_pdbs, train_loader, 1, args.max_protein_length, args.num_examples_per_epoch))127            p.put_nowait(executor.submit(get_pdbs, valid_loader, 1, args.max_protein_length, args.num_examples_per_epoch))128        pdb_dict_train = q.get().result()129        pdb_dict_valid = p.get().result()130       131        dataset_train = StructureDataset(pdb_dict_train, truncate=None, max_length=args.max_protein_length) 132        dataset_valid = StructureDataset(pdb_dict_valid, truncate=None, max_length=args.max_protein_length)133        134        loader_train = StructureLoader(dataset_train, batch_size=args.batch_size)135        loader_valid = StructureLoader(dataset_valid, batch_size=args.batch_size)136        137        reload_c = 0 138        for e in range(args.num_epochs):139            t0 = time.time()140            e = epoch + e141            model.train()142            train_sum, train_weights = 0., 0.143            train_acc = 0.144            if e % args.reload_data_every_n_epochs == 0:145                if reload_c != 0:146                    pdb_dict_train = q.get().result()147                    dataset_train = StructureDataset(pdb_dict_train, truncate=None, max_length=args.max_protein_length)148                    loader_train = StructureLoader(dataset_train, batch_size=args.batch_size)149                    pdb_dict_valid = p.get().result()150                    dataset_valid = StructureDataset(pdb_dict_valid, truncate=None, max_length=args.max_protein_length)151                    loader_valid = StructureLoader(dataset_valid, batch_size=args.batch_size)152                    q.put_nowait(executor.submit(get_pdbs, train_loader, 1, args.max_protein_length, args.num_examples_per_epoch))153                    p.put_nowait(executor.submit(get_pdbs, valid_loader, 1, args.max_protein_length, args.num_examples_per_epoch))154                reload_c += 1155            for _, batch in enumerate(loader_train):156                start_batch = time.time()157                X, S, mask, lengths, chain_M, residue_idx, mask_self, chain_encoding_all = featurize(batch, device)158                elapsed_featurize = time.time() - start_batch159                optimizer.zero_grad()160                mask_for_loss = mask*chain_M161                162                if args.mixed_precision:163                    with torch.cuda.amp.autocast():164                        log_probs = model(X, S, mask, chain_M, residue_idx, chain_encoding_all)165                        _, loss_av_smoothed = loss_smoothed(S, log_probs, mask_for_loss)166           167                    scaler.scale(loss_av_smoothed).backward()168                     169                    if args.gradient_norm > 0.0:170                        total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.gradient_norm)171 172                    scaler.step(optimizer)173                    scaler.update()174                else:175                    log_probs = model(X, S, mask, chain_M, residue_idx, chain_encoding_all)176                    _, loss_av_smoothed = loss_smoothed(S, log_probs, mask_for_loss)177                    loss_av_smoothed.backward()178 179                    if args.gradient_norm > 0.0:180                        total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.gradient_norm)181 182                    optimizer.step()183                184                loss, loss_av, true_false = loss_nll(S, log_probs, mask_for_loss)185            186                train_sum += torch.sum(loss * mask_for_loss).cpu().data.numpy()187                train_acc += torch.sum(true_false * mask_for_loss).cpu().data.numpy()188                train_weights += torch.sum(mask_for_loss).cpu().data.numpy()189 190                total_step += 1191 192            model.eval()193            with torch.no_grad():194                validation_sum, validation_weights = 0., 0.195                validation_acc = 0.196                for _, batch in enumerate(loader_valid):197                    X, S, mask, lengths, chain_M, residue_idx, mask_self, chain_encoding_all = featurize(batch, device)198                    log_probs = model(X, S, mask, chain_M, residue_idx, chain_encoding_all)199                    mask_for_loss = mask*chain_M200                    loss, loss_av, true_false = loss_nll(S, log_probs, mask_for_loss)201                    202                    validation_sum += torch.sum(loss * mask_for_loss).cpu().data.numpy()203                    validation_acc += torch.sum(true_false * mask_for_loss).cpu().data.numpy()204                    validation_weights += torch.sum(mask_for_loss).cpu().data.numpy()205            206            train_loss = train_sum / train_weights207            train_accuracy = train_acc / train_weights208            train_perplexity = np.exp(train_loss)209            validation_loss = validation_sum / validation_weights210            validation_accuracy = validation_acc / validation_weights211            validation_perplexity = np.exp(validation_loss)212            213            train_perplexity_ = np.format_float_positional(np.float32(train_perplexity), unique=False, precision=3)     214            validation_perplexity_ = np.format_float_positional(np.float32(validation_perplexity), unique=False, precision=3)215            train_accuracy_ = np.format_float_positional(np.float32(train_accuracy), unique=False, precision=3)216            validation_accuracy_ = np.format_float_positional(np.float32(validation_accuracy), unique=False, precision=3)217    218            t1 = time.time()219            dt = np.format_float_positional(np.float32(t1-t0), unique=False, precision=1) 220            with open(logfile, 'a') as f:221                f.write(f'epoch: {e+1}, step: {total_step}, time: {dt}, train: {train_perplexity_}, valid: {validation_perplexity_}, train_acc: {train_accuracy_}, valid_acc: {validation_accuracy_}\n')222            print(f'epoch: {e+1}, step: {total_step}, time: {dt}, train: {train_perplexity_}, valid: {validation_perplexity_}, train_acc: {train_accuracy_}, valid_acc: {validation_accuracy_}')223            224            checkpoint_filename_last = base_folder+'model_weights/epoch_last.pt'.format(e+1, total_step)225            torch.save({226                        'epoch': e+1,227                        'step': total_step,228                        'num_edges' : args.num_neighbors,229                        'noise_level': args.backbone_noise,230                        'model_state_dict': model.state_dict(),231                        'optimizer_state_dict': optimizer.optimizer.state_dict(),232                        }, checkpoint_filename_last)233 234            if (e+1) % args.save_model_every_n_epochs == 0:235                checkpoint_filename = base_folder+'model_weights/epoch{}_step{}.pt'.format(e+1, total_step)236                torch.save({237                        'epoch': e+1,238                        'step': total_step,239                        'num_edges' : args.num_neighbors,240                        'noise_level': args.backbone_noise, 241                        'model_state_dict': model.state_dict(),242                        'optimizer_state_dict': optimizer.optimizer.state_dict(),243                        }, checkpoint_filename)244 245 246if __name__ == "__main__":247    argparser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter)248 249    argparser.add_argument("--path_for_training_data", type=str, default="my_path/pdb_2021aug02", help="path for loading training data") 250    argparser.add_argument("--path_for_outputs", type=str, default="./exp_020", help="path for logs and model weights")251    argparser.add_argument("--previous_checkpoint", type=str, default="", help="path for previous model weights, e.g. file.pt")252    argparser.add_argument("--num_epochs", type=int, default=200, help="number of epochs to train for")253    argparser.add_argument("--save_model_every_n_epochs", type=int, default=10, help="save model weights every n epochs")254    argparser.add_argument("--reload_data_every_n_epochs", type=int, default=2, help="reload training data every n epochs")255    argparser.add_argument("--num_examples_per_epoch", type=int, default=1000000, help="number of training example to load for one epoch")256    argparser.add_argument("--batch_size", type=int, default=10000, help="number of tokens for one batch")257    argparser.add_argument("--max_protein_length", type=int, default=10000, help="maximum length of the protein complext")258    argparser.add_argument("--hidden_dim", type=int, default=128, help="hidden model dimension")259    argparser.add_argument("--num_encoder_layers", type=int, default=3, help="number of encoder layers") 260    argparser.add_argument("--num_decoder_layers", type=int, default=3, help="number of decoder layers")261    argparser.add_argument("--num_neighbors", type=int, default=48, help="number of neighbors for the sparse graph")   262    argparser.add_argument("--dropout", type=float, default=0.1, help="dropout level; 0.0 means no dropout")263    argparser.add_argument("--backbone_noise", type=float, default=0.2, help="amount of noise added to backbone during training")   264    argparser.add_argument("--rescut", type=float, default=3.5, help="PDB resolution cutoff")265    argparser.add_argument("--debug", type=bool, default=False, help="minimal data loading for debugging")266    argparser.add_argument("--gradient_norm", type=float, default=-1.0, help="clip gradient norm, set to negative to omit clipping")267    argparser.add_argument("--mixed_precision", type=bool, default=True, help="train with mixed precision")268 269    args = argparser.parse_args()    270    main(args)   271