OneScience-Group/ProteinMPNN
124
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 