File size: 9,510 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 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 | import os
import shutil
import argparse
from tqdm.auto import tqdm
import torch
from torch.nn.utils import clip_grad_norm_
import torch.utils.tensorboard
import yaml
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
from onescience.utils.targetdiff.train import get_scheduler, get_optimizer
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
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 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')
parser.add_argument('--tag', type=str, default='')
args, config_overrides = parser.parse_known_args()
# Load configs
config = utils_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('.')]
utils_misc.seed_all(config.train.seed)
# Logging
log_dir = utils_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)
logger = utils_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 = utils_trans.FeaturizeProteinAtom()
ligand_featurizer = utils_trans.FeaturizeLigandAtom()
transform = Compose([
protein_featurizer,
ligand_featurizer,
])
# Datasets and loaders
logger.info('Loading dataset...')
dataset, subsets = get_dataset(
config=config.dataset,
transform=transform,
emb_path=config.dataset.emb_path if 'emb_path' in config.dataset else None,
heavy_only=config.dataset.heavy_only
)
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)
# Model
logger.info('Building model...')
model = get_model(config, protein_featurizer.feature_dim, ligand_featurizer.feature_dim)
model = model.to(args.device)
logger.info(f'# trainable parameters: {utils_misc.count_parameters(model) / 1e6:.4f} M')
# Optimizer and scheduler
optimizer = get_optimizer(config.train.optimizer, model)
scheduler = get_scheduler(config.train.scheduler, optimizer)
def train(epoch):
model.train()
optimizer.zero_grad()
it = 0
num_it = len(train_loader)
for batch in tqdm(train_loader, dynamic_ncols=True, desc=f'Epoch {epoch}', position=1):
it += 1
batch = batch.to(args.device)
# compute loss
loss = model.get_loss(batch, pos_noise_std=config.train.pos_noise_std)
loss.backward()
orig_grad_norm = clip_grad_norm_(model.parameters(), config.train.max_grad_norm)
optimizer.step()
optimizer.zero_grad()
if it % config.train.report_iter == 0:
logger.info('[Train] Epoch %03d Iter %04d | Loss %.6f | Lr %.4f * 1e-3' % (
epoch, it, loss.item(), optimizer.param_groups[0]['lr'] * 1000
))
writer.add_scalar('train/loss', loss, it + epoch * num_it)
writer.add_scalar('train/lr', optimizer.param_groups[0]['lr'], it + epoch * num_it)
writer.add_scalar('train/grad', orig_grad_norm, it + epoch * num_it)
writer.flush()
def validate(epoch, data_loader, scheduler, writer, prefix='Validate'):
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)
if scheduler:
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()
if writer:
writer.add_scalar('val/loss', avg_loss, epoch)
writer.add_scalar('val/rmse', rmse, epoch)
writer.flush()
return avg_loss
try:
best_val_loss = float('inf')
best_val_epoch = 0
patience = 0
for epoch in range(1, config.train.max_epochs + 1):
# with torch.autograd.detect_anomaly():
train(epoch)
if epoch % config.train.val_freq == 0 or epoch == config.train.max_epochs:
val_loss = validate(epoch, val_loader, scheduler, writer)
validate(epoch, test_loader, scheduler=None, writer=None, prefix='Test')
if val_loss < best_val_loss:
patience = 0
best_val_loss = val_loss
best_val_epoch = epoch
logger.info(f'Best val achieved at epoch {epoch}, val loss: {best_val_loss:.3f}')
logger.info(f'Eval on Test set:')
validate(epoch, test_loader, scheduler=None, writer=None, prefix='Test')
ckpt_path = os.path.join(ckpt_dir, '%d.pt' % epoch)
torch.save({
'config': config,
'model': model.state_dict(),
'optimizer': optimizer.state_dict(),
'scheduler': scheduler.state_dict(),
'epoch': epoch,
}, ckpt_path)
logger.info(f'Model {log_dir}/{epoch}.pt saved!')
else:
patience += 1
logger.info(f'Val loss does not improve, patience: {patience} '
f'(Best val loss: {best_val_loss:.3f} at epoch {best_val_epoch})')
except KeyboardInterrupt:
logger.info('Terminating...')
if __name__ == '__main__':
main()
|