File size: 4,882 Bytes
3ac1d94 | 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 | import argparse
import os
import shutil
import torch
from torch_geometric.transforms import Compose
import onescience.utils.targetdiff.misc as misc
import onescience.utils.targetdiff.transforms as trans
from onescience.datapipes.targetdiff.pl_data import ProteinLigandData, torchify_dict
from models.molopt_score_model import ScorePosNet3D
from scripts.sample_diffusion import sample_diffusion_ligand
from onescience.utils.targetdiff.data import PDBProtein
from onescience.utils.targetdiff import reconstruct
from rdkit import Chem
def pdb_to_pocket_data(pdb_path):
pocket_dict = PDBProtein(pdb_path).to_dict_atom()
data = ProteinLigandData.from_protein_ligand_dicts(
protein_dict=torchify_dict(pocket_dict),
ligand_dict={
'element': torch.empty([0, ], dtype=torch.long),
'pos': torch.empty([0, 3], dtype=torch.float),
'atom_feature': torch.empty([0, 8], dtype=torch.float),
'bond_index': torch.empty([2, 0], dtype=torch.long),
'bond_type': torch.empty([0, ], dtype=torch.long),
}
)
return data
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument('config', type=str)
parser.add_argument('--pdb_path', type=str)
parser.add_argument('--device', type=str, default='cuda:0')
parser.add_argument('--batch_size', type=int, default=100)
parser.add_argument('--result_path', type=str, default='./outputs_pdb')
parser.add_argument('--num_samples', type=int)
args = parser.parse_args()
logger = misc.get_logger('evaluate')
# Load config
config = misc.load_config(args.config)
logger.info(config)
misc.seed_all(config.sample.seed)
# Load checkpoint
ckpt = torch.load(config.model.checkpoint, map_location=args.device, weights_only=False)
logger.info(f"Training Config: {ckpt['config']}")
# Transforms
protein_featurizer = trans.FeaturizeProteinAtom()
ligand_atom_mode = ckpt['config'].data.transform.ligand_atom_mode
ligand_featurizer = trans.FeaturizeLigandAtom(ligand_atom_mode)
transform = Compose([
protein_featurizer,
])
# Load model
model = ScorePosNet3D(
ckpt['config'].model,
protein_atom_feature_dim=protein_featurizer.feature_dim,
ligand_atom_feature_dim=ligand_featurizer.feature_dim
).to(args.device)
model.load_state_dict(ckpt['model'], strict=False if 'train_config' in config.model else True)
logger.info(f'Successfully load the model! {config.model.checkpoint}')
# Load pocket
data = pdb_to_pocket_data(args.pdb_path)
data = transform(data)
if args.num_samples:
config.sample.num_samples = args.num_samples
all_pred_pos, all_pred_v, pred_pos_traj, pred_v_traj, pred_v0_traj, pred_vt_traj, time_list = sample_diffusion_ligand(
model, data, config.sample.num_samples,
batch_size=args.batch_size, device=args.device,
num_steps=config.sample.num_steps,
pos_only=config.sample.pos_only,
center_pos_mode=config.sample.center_pos_mode,
sample_num_atoms=config.sample.sample_num_atoms
)
result = {
'data': data,
'pred_ligand_pos': all_pred_pos,
'pred_ligand_v': all_pred_v,
'pred_ligand_pos_traj': pred_pos_traj,
'pred_ligand_v_traj': pred_v_traj
}
logger.info('Sample done!')
# reconstruction
gen_mols = []
n_recon_success, n_complete = 0, 0
for sample_idx, (pred_pos, pred_v) in enumerate(zip(all_pred_pos, all_pred_v)):
pred_atom_type = trans.get_atomic_number_from_index(pred_v, mode='add_aromatic')
try:
pred_aromatic = trans.is_aromatic_from_index(pred_v, mode='add_aromatic')
mol = reconstruct.reconstruct_from_generated(pred_pos, pred_atom_type, pred_aromatic)
smiles = Chem.MolToSmiles(mol)
except reconstruct.MolReconsError:
gen_mols.append(None)
continue
n_recon_success += 1
if '.' in smiles:
gen_mols.append(None)
continue
n_complete += 1
gen_mols.append(mol)
result['mols'] = gen_mols
logger.info('Reconstruction done!')
logger.info(f'n recon: {n_recon_success} n complete: {n_complete}')
result_path = args.result_path
os.makedirs(result_path, exist_ok=True)
shutil.copyfile(args.config, os.path.join(result_path, 'sample.yml'))
torch.save(result, os.path.join(result_path, f'sample.pt'))
mols_save_path = os.path.join(result_path, f'sdf')
os.makedirs(mols_save_path, exist_ok=True)
for idx, mol in enumerate(gen_mols):
if mol is not None:
sdf_writer = Chem.SDWriter(os.path.join(mols_save_path, f'{idx:03d}.sdf'))
sdf_writer.write(mol)
sdf_writer.close()
logger.info(f'Results are saved in {result_path}')
|