| import argparse |
| from tqdm.auto import tqdm |
| import torch |
| import torch.utils.tensorboard |
| from torch_geometric.transforms import Compose |
|
|
| from onescience.datapipes.targetdiff import get_dataset |
| import onescience.utils.targetdiff.transforms_prop as utils_trans |
| import onescience.utils.targetdiff.misc as utils_misc |
| import numpy as np |
| from onescience.datapipes.targetdiff.protein_ligand import KMAP |
| from scripts.property_prediction.local_misc_prop import get_model, get_dataloader, get_eval_scores |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument('--ckpt_path', type=str) |
| parser.add_argument('--device', type=str, default='cuda') |
| parser.add_argument('--seed', type=int, default=2021) |
| args = parser.parse_args() |
| utils_misc.seed_all(args.seed) |
|
|
| |
| logger = utils_misc.get_logger('eval') |
| logger.info(args) |
|
|
| |
| logger.info(f'Loading model from {args.ckpt_path}') |
| ckpt_restore = torch.load(args.ckpt_path, map_location=torch.device('cpu')) |
| config = ckpt_restore['config'] |
| logger.info(f'ckpt_config: {config}') |
|
|
| |
| protein_featurizer = utils_trans.FeaturizeProteinAtom() |
| ligand_featurizer = utils_trans.FeaturizeLigandAtom() |
| transform = Compose([ |
| protein_featurizer, |
| ligand_featurizer, |
| ]) |
|
|
| |
| model = get_model(config, protein_featurizer.feature_dim, ligand_featurizer.feature_dim) |
| model.load_state_dict(ckpt_restore['model']) |
| model = model.to(args.device) |
| |
| |
| model.eval() |
|
|
| |
| |
| |
| logger.info('Loading dataset...') |
| dataset, subsets = get_dataset( |
| config=config.dataset, |
| transform=transform, |
| heavy_only=config.dataset.get('heavy_only', False) |
| ) |
| train_set, val_set, test_set = subsets['train'], subsets['val'], subsets['test'] |
| logger.info(f'Train set: {len(train_set)} Val set: {len(val_set)} Test set: {len(test_set)}') |
| train_loader, val_loader, test_loader = get_dataloader(train_set, val_set, test_set, config) |
|
|
| def validate(epoch, data_loader, prefix='Test'): |
| sum_loss, sum_n = 0, 0 |
| ytrue_arr, ypred_arr = [], [] |
| y_kind = [] |
| with torch.no_grad(): |
| model.eval() |
| for batch in tqdm(data_loader, desc=prefix): |
| batch = batch.to(args.device) |
| loss, pred = model.get_loss(batch, pos_noise_std=0., return_pred=True) |
| sum_loss += loss.item() * len(batch.y) |
| sum_n += len(batch.y) |
| ypred_arr.append(pred.view(-1)) |
| ytrue_arr.append(batch.y) |
| y_kind.append(batch.kind) |
| avg_loss = sum_loss / sum_n |
| logger.info('[%s] Epoch %03d | Loss %.6f' % ( |
| prefix, epoch, avg_loss, |
| )) |
| ypred_arr = torch.cat(ypred_arr).cpu().numpy().astype(np.float64) |
| ytrue_arr = torch.cat(ytrue_arr).cpu().numpy().astype(np.float64) |
| y_kind = torch.cat(y_kind).cpu().numpy() |
| rmse = get_eval_scores(ypred_arr, ytrue_arr, logger) |
| for k, v in KMAP.items(): |
| get_eval_scores(ypred_arr[y_kind == v], ytrue_arr[y_kind == v], logger, prefix=k) |
| return avg_loss |
|
|
| test_loss = validate(ckpt_restore['epoch'], test_loader) |
| print('Test loss: ', test_loss) |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|