Download scripts/train_diffusion.py from OneScience-Group/TargetDiff: direct link, hf CLI and curl.
- Browser
- Download file 12 kB
-
https://huggingface.co/OneScience-Group/TargetDiff/resolve/3f4bfb540c16acf7d89c469c4e1777fabc69090d/scripts/train_diffusion.py
- Command line
-
hf download hf://OneScience-Group/TargetDiff@3f4bfb540c16acf7d89c469c4e1777fabc69090d/scripts/train_diffusion.py
-
curl -L -o train_diffusion.py https://huggingface.co/OneScience-Group/TargetDiff/resolve/3f4bfb540c16acf7d89c469c4e1777fabc69090d/scripts/train_diffusion.py
12 kB
| import argparse | |
| import os | |
| import shutil | |
| import numpy as np | |
| import torch | |
| import torch.utils.tensorboard | |
| import yaml | |
| from sklearn.metrics import roc_auc_score | |
| from torch.nn.utils import clip_grad_norm_ | |
| from torch_geometric.loader import DataLoader | |
| from torch_geometric.transforms import Compose | |
| from tqdm.auto import tqdm | |
| import onescience.utils.targetdiff.misc as misc | |
| import onescience.utils.targetdiff.train as utils_train | |
| import onescience.utils.targetdiff.transforms as trans | |
| from onescience.datapipes.targetdiff import get_dataset | |
| from onescience.datapipes.targetdiff.pl_data import FOLLOW_BATCH | |
| from models.molopt_score_model import ScorePosNet3D | |
| REPO_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), '..', '..', '..', '..')) | |
| MODELS_SNAPSHOT_SRC = os.path.join(REPO_ROOT, 'src', 'onescience', 'models', 'targetdiff') | |
| def parse_override_value(raw_value, old_value): | |
| parsed_value = yaml.safe_load(raw_value) | |
| if old_value is None: | |
| return parsed_value | |
| if isinstance(old_value, bool): | |
| if isinstance(parsed_value, bool): | |
| return parsed_value | |
| return str(parsed_value).lower() in ('1', 'true', 'yes', 'y') | |
| if isinstance(old_value, tuple): | |
| if isinstance(parsed_value, str): | |
| return tuple(item.strip() for item in parsed_value.split(',')) | |
| return tuple(parsed_value) | |
| if isinstance(old_value, list): | |
| if isinstance(parsed_value, str): | |
| return [item.strip() for item in parsed_value.split(',')] | |
| return list(parsed_value) | |
| return type(old_value)(parsed_value) | |
| def apply_config_overrides(config, overrides): | |
| if len(overrides) % 2 != 0: | |
| raise ValueError('Config overrides must use "--key value" pairs.') | |
| for key_arg, raw_value in zip(overrides[::2], overrides[1::2]): | |
| if not key_arg.startswith('--'): | |
| raise ValueError(f'Config override key must start with "--": {key_arg}') | |
| key_path = key_arg[2:] | |
| parts = key_path.split('.') | |
| node = config | |
| for part in parts[:-1]: | |
| if part not in node: | |
| raise KeyError(f'Unknown config override: {key_path}') | |
| node = node[part] | |
| leaf = parts[-1] | |
| if leaf not in node: | |
| raise KeyError(f'Unknown config override: {key_path}') | |
| node[leaf] = parse_override_value(raw_value, node[leaf]) | |
| return config | |
| def get_auroc(y_true, y_pred, feat_mode): | |
| y_true = np.array(y_true) | |
| y_pred = np.array(y_pred) | |
| avg_auroc = 0. | |
| possible_classes = set(y_true) | |
| for c in possible_classes: | |
| auroc = roc_auc_score(y_true == c, y_pred[:, c]) | |
| avg_auroc += auroc * np.sum(y_true == c) | |
| mapping = { | |
| 'basic': trans.MAP_INDEX_TO_ATOM_TYPE_ONLY, | |
| 'add_aromatic': trans.MAP_INDEX_TO_ATOM_TYPE_AROMATIC, | |
| 'full': trans.MAP_INDEX_TO_ATOM_TYPE_FULL | |
| } | |
| print(f'atom: {mapping[feat_mode][c]} \t auc roc: {auroc:.4f}') | |
| return avg_auroc / len(y_true) | |
| if __name__ == '__main__': | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument('config', type=str) | |
| parser.add_argument('--device', type=str, default='cuda') | |
| parser.add_argument('--logdir', type=str, default='./logs_diffusion') | |
| parser.add_argument('--tag', type=str, default='') | |
| parser.add_argument('--train_report_iter', type=int, default=200) | |
| args, config_overrides = parser.parse_known_args() | |
| # Load configs | |
| config = misc.load_config(args.config) | |
| config = apply_config_overrides(config, config_overrides) | |
| config_name = os.path.basename(args.config)[:os.path.basename(args.config).rfind('.')] | |
| misc.seed_all(config.train.seed) | |
| # Logging | |
| log_dir = misc.get_new_log_dir(args.logdir, prefix=config_name, tag=args.tag) | |
| ckpt_dir = os.path.join(log_dir, 'checkpoints') | |
| os.makedirs(ckpt_dir, exist_ok=True) | |
| vis_dir = os.path.join(log_dir, 'vis') | |
| os.makedirs(vis_dir, exist_ok=True) | |
| logger = misc.get_logger('train', log_dir) | |
| writer = torch.utils.tensorboard.SummaryWriter(log_dir) | |
| logger.info(args) | |
| logger.info(config) | |
| shutil.copyfile(args.config, os.path.join(log_dir, os.path.basename(args.config))) | |
| local_models_dir = './models' | |
| if os.path.isdir(local_models_dir): | |
| shutil.copytree(local_models_dir, os.path.join(log_dir, 'models')) | |
| elif os.path.isdir(MODELS_SNAPSHOT_SRC): | |
| shutil.copytree(MODELS_SNAPSHOT_SRC, os.path.join(log_dir, 'models')) | |
| else: | |
| logger.warning('Skip model source snapshot: no ./models or migrated TargetDiff models directory found.') | |
| # Transforms | |
| protein_featurizer = trans.FeaturizeProteinAtom() | |
| ligand_featurizer = trans.FeaturizeLigandAtom(config.data.transform.ligand_atom_mode) | |
| transform_list = [ | |
| protein_featurizer, | |
| ligand_featurizer, | |
| trans.FeaturizeLigandBond(), | |
| ] | |
| if config.data.transform.random_rot: | |
| transform_list.append(trans.RandomRotation()) | |
| transform = Compose(transform_list) | |
| # Datasets and loaders | |
| logger.info('Loading dataset...') | |
| dataset, subsets = get_dataset( | |
| config=config.data, | |
| transform=transform | |
| ) | |
| train_set, val_set = subsets['train'], subsets['test'] | |
| logger.info(f'Training: {len(train_set)} Validation: {len(val_set)}') | |
| # follow_batch = ['protein_element', 'ligand_element'] | |
| collate_exclude_keys = ['ligand_nbh_list'] | |
| train_iterator = utils_train.inf_iterator(DataLoader( | |
| train_set, | |
| batch_size=config.train.batch_size, | |
| shuffle=True, | |
| num_workers=config.train.num_workers, | |
| follow_batch=FOLLOW_BATCH, | |
| exclude_keys=collate_exclude_keys | |
| )) | |
| val_loader = DataLoader(val_set, config.train.batch_size, shuffle=False, | |
| follow_batch=FOLLOW_BATCH, exclude_keys=collate_exclude_keys) | |
| # Model | |
| logger.info('Building model...') | |
| model = ScorePosNet3D( | |
| config.model, | |
| protein_atom_feature_dim=protein_featurizer.feature_dim, | |
| ligand_atom_feature_dim=ligand_featurizer.feature_dim | |
| ).to(args.device) | |
| # print(model) | |
| print(f'protein feature dim: {protein_featurizer.feature_dim} ligand feature dim: {ligand_featurizer.feature_dim}') | |
| logger.info(f'# trainable parameters: {misc.count_parameters(model) / 1e6:.4f} M') | |
| # Optimizer and scheduler | |
| optimizer = utils_train.get_optimizer(config.train.optimizer, model) | |
| scheduler = utils_train.get_scheduler(config.train.scheduler, optimizer) | |
| def train(it): | |
| model.train() | |
| optimizer.zero_grad() | |
| for _ in range(config.train.n_acc_batch): | |
| batch = next(train_iterator).to(args.device) | |
| protein_noise = torch.randn_like(batch.protein_pos) * config.train.pos_noise_std | |
| gt_protein_pos = batch.protein_pos + protein_noise | |
| results = model.get_diffusion_loss( | |
| protein_pos=gt_protein_pos, | |
| protein_v=batch.protein_atom_feature.float(), | |
| batch_protein=batch.protein_element_batch, | |
| ligand_pos=batch.ligand_pos, | |
| ligand_v=batch.ligand_atom_feature_full, | |
| batch_ligand=batch.ligand_element_batch | |
| ) | |
| loss, loss_pos, loss_v = results['loss'], results['loss_pos'], results['loss_v'] | |
| loss = loss / config.train.n_acc_batch | |
| loss.backward() | |
| orig_grad_norm = clip_grad_norm_(model.parameters(), config.train.max_grad_norm) | |
| optimizer.step() | |
| if it % args.train_report_iter == 0: | |
| logger.info( | |
| '[Train] Iter %d | Loss %.6f (pos %.6f | v %.6f) | Lr: %.6f | Grad Norm: %.6f' % ( | |
| it, loss, loss_pos, loss_v, optimizer.param_groups[0]['lr'], orig_grad_norm | |
| ) | |
| ) | |
| for k, v in results.items(): | |
| if torch.is_tensor(v) and v.squeeze().ndim == 0: | |
| writer.add_scalar(f'train/{k}', v, it) | |
| writer.add_scalar('train/lr', optimizer.param_groups[0]['lr'], it) | |
| writer.add_scalar('train/grad', orig_grad_norm, it) | |
| writer.flush() | |
| def validate(it): | |
| # fix time steps | |
| sum_loss, sum_loss_pos, sum_loss_v, sum_n = 0, 0, 0, 0 | |
| sum_loss_bond, sum_loss_non_bond = 0, 0 | |
| all_pred_v, all_true_v = [], [] | |
| all_pred_bond_type, all_gt_bond_type = [], [] | |
| with torch.no_grad(): | |
| model.eval() | |
| for batch in tqdm(val_loader, desc='Validate'): | |
| batch = batch.to(args.device) | |
| batch_size = batch.num_graphs | |
| t_loss, t_loss_pos, t_loss_v = [], [], [] | |
| for t in np.linspace(0, model.num_timesteps - 1, 10).astype(int): | |
| time_step = torch.tensor([t] * batch_size).to(args.device) | |
| results = model.get_diffusion_loss( | |
| protein_pos=batch.protein_pos, | |
| protein_v=batch.protein_atom_feature.float(), | |
| batch_protein=batch.protein_element_batch, | |
| ligand_pos=batch.ligand_pos, | |
| ligand_v=batch.ligand_atom_feature_full, | |
| batch_ligand=batch.ligand_element_batch, | |
| time_step=time_step | |
| ) | |
| loss, loss_pos, loss_v = results['loss'], results['loss_pos'], results['loss_v'] | |
| sum_loss += float(loss) * batch_size | |
| sum_loss_pos += float(loss_pos) * batch_size | |
| sum_loss_v += float(loss_v) * batch_size | |
| sum_n += batch_size | |
| all_pred_v.append(results['ligand_v_recon'].detach().cpu().numpy()) | |
| all_true_v.append(batch.ligand_atom_feature_full.detach().cpu().numpy()) | |
| avg_loss = sum_loss / sum_n | |
| avg_loss_pos = sum_loss_pos / sum_n | |
| avg_loss_v = sum_loss_v / sum_n | |
| atom_auroc = get_auroc(np.concatenate(all_true_v), np.concatenate(all_pred_v, axis=0), | |
| feat_mode=config.data.transform.ligand_atom_mode) | |
| if config.train.scheduler.type == 'plateau': | |
| scheduler.step(avg_loss) | |
| elif config.train.scheduler.type == 'warmup_plateau': | |
| scheduler.step_ReduceLROnPlateau(avg_loss) | |
| else: | |
| scheduler.step() | |
| logger.info( | |
| '[Validate] Iter %05d | Loss %.6f | Loss pos %.6f | Loss v %.6f e-3 | Avg atom auroc %.6f' % ( | |
| it, avg_loss, avg_loss_pos, avg_loss_v * 1000, atom_auroc | |
| ) | |
| ) | |
| writer.add_scalar('val/loss', avg_loss, it) | |
| writer.add_scalar('val/loss_pos', avg_loss_pos, it) | |
| writer.add_scalar('val/loss_v', avg_loss_v, it) | |
| writer.flush() | |
| return avg_loss | |
| try: | |
| best_loss, best_iter = None, None | |
| for it in range(1, config.train.max_iters + 1): | |
| # with torch.autograd.detect_anomaly(): | |
| train(it) | |
| if it % config.train.val_freq == 0 or it == config.train.max_iters: | |
| val_loss = validate(it) | |
| if best_loss is None or val_loss < best_loss: | |
| logger.info(f'[Validate] Best val loss achieved: {val_loss:.6f}') | |
| best_loss, best_iter = val_loss, it | |
| ckpt_path = os.path.join(ckpt_dir, '%d.pt' % it) | |
| torch.save({ | |
| 'config': config, | |
| 'model': model.state_dict(), | |
| 'optimizer': optimizer.state_dict(), | |
| 'scheduler': scheduler.state_dict(), | |
| 'iteration': it, | |
| }, ckpt_path) | |
| else: | |
| logger.info(f'[Validate] Val loss is not improved. ' | |
| f'Best val loss: {best_loss:.6f} at iter {best_iter}') | |
| except KeyboardInterrupt: | |
| logger.info('Terminating...') | |