From 41895cd6d5d8b2d38757fbc0c8abc0bc9d5f9f1d Mon Sep 17 00:00:00 2001 From: Steph Prince <40640337+stephprince@users.noreply.github.com> Date: Fri, 14 Mar 2025 13:04:23 -0700 Subject: [PATCH 1/7] add msa refinement model --- src/metfish/refinement_model/__init__.py | 0 .../refinement_model/refinement_model.py | 142 +++++++++++++++++ src/metfish/refinement_model/train.py | 144 ++++++++++++++++++ 3 files changed, 286 insertions(+) create mode 100644 src/metfish/refinement_model/__init__.py create mode 100644 src/metfish/refinement_model/refinement_model.py create mode 100644 src/metfish/refinement_model/train.py diff --git a/src/metfish/refinement_model/__init__.py b/src/metfish/refinement_model/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/metfish/refinement_model/refinement_model.py b/src/metfish/refinement_model/refinement_model.py new file mode 100644 index 0000000..eda3325 --- /dev/null +++ b/src/metfish/refinement_model/refinement_model.py @@ -0,0 +1,142 @@ + +import time +import torch +import pytorch_lightning as pl + +from openfold.utils.import_weights import import_jax_weights_ +from openfold.utils.tensor_utils import tensor_tree_map +from openfold.model.model import AlphaFold + +from metfish.msa_model.utils.loss import compute_saxs + + +class SAXSLoss(torch.nn.Module): + """SAXS """ + def __init__(self, config): + super(SAXSLoss, self).__init__() + self.config = config + + def forward(self, out, batch): + all_atom_pred_pos = out["final_atom_positions"] + all_atom_mask = batch["all_atom_mask"] + all_atom_true_pos = batch["all_atom_positions"] + step = self.config.saxs_loss.step + dmax = self.config.saxs_loss.dmax + + # calculate predicted and true saxs + pred_saxs = compute_saxs(all_atom_pos=all_atom_pred_pos, all_atom_mask=all_atom_mask, step=step, dmax=dmax) + true_saxs = compute_saxs(all_atom_pos=all_atom_true_pos, all_atom_mask=all_atom_mask, step=step, dmax=dmax) + + # get L1 loss + l1_loss = torch.nn.L1Loss(reduction="sum") + loss = l1_loss(pred_saxs, true_saxs) + + return loss + + +class MSARefinementModel(torch.nn.Module): + def __init__(self, config, training=True): + super(MSARefinementModel, self).__init__() + self.config = config + self.af_model = AlphaFold(config) + self.training = training + + for param in self.af_model.parameters(): + param.requires_grad = False + + # initialize parameters + self.w = None + self.b = None + + def initialize_parameters(self, msa): + self.w = torch.nn.Parameter(torch.ones_like(msa)) + self.b = torch.nn.Parameter(torch.zeros_like(msa)) + + def forward(self, batch): + device = batch['aatype'].device + self.w = self.w.to(device) + self.b = self.b.to(device) + + # refine msa with linear layer + # TODO - msa cluster profile here may mean the msa features after embedding... need to modify if so + msa_refined = self.w * batch['msa_feat'] + self.b + batch['msa_feat'] = msa_refined + + # run through alphafold + outputs = self.af_model(batch) + + return outputs + +# define the lightning module for training +class MSARefinementModelWrapper(pl.LightningModule): + def __init__(self, config, training=True, lr_mul=1.0, lr_add=0.05): + super().__init__() + self.save_hyperparameters() + self.config = config + self.model = MSARefinementModel(config) + self.training = training + if training: + self.loss = SAXSLoss(config.loss) + self.cached_weights = None + self.lr_mul = 1.0 + self.lr_add = 0.05 + + # activate manual optimization + self.automatic_optimization = False + + self.num_iterations = 100 + self.last_log_time = time.time() + + # TODO - initialize MSA refinement parameters as part of config file + self.model.initialize_parameters(torch.zeros(256)) + + def forward(self, batch): + return self.model(batch) + + def _log(self, loss, iter=None, train=True): + self.log( + f"loss_iter{iter}", + loss, + on_step=False, on_epoch=True, logger=True, sync_dist=True + ) + self.log('dur', time.time() - self.last_log_time, sync_dist=True) + self.last_log_time = time.time() + + def training_step(self, batch): + self.model.initialize_parameters(batch['msa_feat']) + opt = self.optimizers() + + for n in range(self.num_iterations): + print(f'Running iteration {n} / {self.num_iterations}') + + # clear gradients + opt.zero_grad() + + # forward pass + outputs = self.model(batch) + + # calculate loss + batch_no_recycling = tensor_tree_map(lambda t: t[..., -1], batch) # remove recycling dimension + loss = self.loss(outputs, batch_no_recycling) + + # backwards pass and update weights + self.manual_backward(loss, retain_graph=True) + opt.step() + + # log the loss + self._log(loss.detach(), iter=n) + + return loss + + def configure_optimizers(self, learning_rate: float = 1e-3, eps: float = 1e-5,): + optimizer = torch.optim.Adam([ + {'params': [self.model.w], 'lr': self.lr_mul}, + {'params': [self.model.b], 'lr': self.lr_add} + ], eps=eps) + + return optimizer + + def load_from_jax(self, jax_path): + import_jax_weights_( + self.model.af_model, jax_path, version='model_3' + ) diff --git a/src/metfish/refinement_model/train.py b/src/metfish/refinement_model/train.py new file mode 100644 index 0000000..1e654a3 --- /dev/null +++ b/src/metfish/refinement_model/train.py @@ -0,0 +1,144 @@ + +import argparse +import torch +import os +import pytorch_lightning as pl + +from pytorch_lightning.loggers import CSVLogger, WandbLogger +from torch.utils.data import DataLoader + +from metfish.msa_model.config import model_config +from metfish.msa_model.data.data_modules import MSASAXSDataset +from metfish.refinement_model.refinement_model import MSARefinementModelWrapper + +# gives a speedup on Ampere-class GPUs +torch.set_float32_matmul_precision("high") + +parser = argparse.ArgumentParser() +parser.add_argument( + "data_dir", type=str, + help="Directory containing training pdb, saxs, and msa data", +) +parser.add_argument( + "output_dir", type=str, + help='''Directory in which to output checkpoints, logs, etc. Ignored + if not on rank 0''', +) +parser.add_argument( + "--ckpt_path", type=str, + help='''Path to a model checkpoint from which to resume training.''', +) +parser.add_argument( + "--gpus_per_node", type=int, default=1, help='Number of gpus per node (will use all 4 per node on perlmutter).' +) +parser.add_argument( + "--num_nodes", type=int, default=1, help='Number of nodes to use for training.' +) +parser.add_argument( + "--batch_size", type=int, default=2, help='Batch size for each training step' +) +parser.add_argument( + "--seed", type=int, default=1, + help="Random seed" +) +parser.add_argument( + "--use_wandb", action="store_true", default=True, + help="Whether to log metrics to Weights & Biases" +) +parser.add_argument( + "--fast_dev_run", default=False, action='store_true', + help="Whether to run a fast dev run of a single batch for testing purposes" +) +parser.add_argument( + "--jax_param_path", type=str, default="/pscratch/sd/s/smprince/projects/alphaflow/params_model_1.npz", # these are the original AF weights, + help="""Path to an .npz JAX parameter file with which to initialize the model""" +) +parser.add_argument( + "--precision", type=str, default='bf16-mixed', + help='Sets precision, lower precision improves runtime performance.', +) +parser.add_argument( + "--max_epochs", type=int, default=100, +) +parser.add_argument( + "--log_every_n_steps", type=int, default=25, +) +parser.add_argument( + "--job_name", type=str, + help='''Name of job to be used for logging purposes.''', +) + +def main(data_dir="/global/cfs/cdirs/m3513/metfish/PDB70_verB_fixed_data/result", + output_dir="/pscratch/sd/s/smprince/projects/metfish/model_outputs", + ckpt_path=None, + gpus_per_node=1, + num_nodes=1, + batch_size=2, + seed=1, + use_wandb=True, + fast_dev_run=False, + jax_param_path="/pscratch/sd/s/smprince/projects/alphaflow/params_model_1.npz", + resume_from_ckpt=False, + precision='bf16-mixed', + max_epochs=100, + log_every_n_steps=25, + job_name='default', + ): + + # set up data paths and configuration + pdb_dir = f"{data_dir}/pdb" + saxs_dir = f"{data_dir}/saxs_r" + msa_dir = f"{data_dir}/msa" + csv_dir = f"{data_dir}/scripts" + training_csv = f'{csv_dir}/input_training.csv' # NOTE - this was msa_dir for training v_1 + + pl.seed_everything(seed, workers=True) + strategy = "ddp" if (gpus_per_node > 1) or num_nodes > 1 else "auto" + config = model_config('initial_training', train=True, low_prec=True) + data_config = config.data + data_config.common.use_templates = False + data_config.common.max_recycling_iters = 0 + + # set up training and test datasets and dataloaders + train_dataset = MSASAXSDataset(data_config, training_csv, msa_dir=msa_dir, saxs_dir=saxs_dir, pdb_dir=pdb_dir, saxs_ext='.pr.csv', pdb_prefix='') + train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=8) + + # initialize model + refinement_model = MSARefinementModelWrapper(config) + + # add logging + loggers = [CSVLogger(f"{output_dir}/lightning_logs", name="msasaxs")] + if use_wandb: + os.environ["WANDB__SERVICE_WAIT"] = "300" + os.environ["WANDB_MODE"] = "offline" + # wandb.init() + loggers.append(WandbLogger(name="msasaxs", save_dir=f"{output_dir}/lightning_logs")) + + # initialize trainer + trainer = pl.Trainer(accelerator="gpu", + strategy=strategy, + max_epochs=max_epochs, + limit_train_batches=1.0, + logger=loggers, + log_every_n_steps=log_every_n_steps, + default_root_dir=output_dir, + devices=gpus_per_node, + num_nodes=num_nodes, + fast_dev_run=fast_dev_run, + precision=precision, + ) + + # load existing weights + if jax_param_path and not resume_from_ckpt: + refinement_model.load_from_jax(jax_param_path) + print(f"Successfully loaded JAX parameters at {jax_param_path}...") + + # fit the model + trainer.fit(model=refinement_model, train_dataloaders=train_loader, ckpt_path=ckpt_path) + + print('done') + +if __name__ == "__main__": + main(data_dir="/global/cfs/cdirs/m3513/metfish/NMR_training/data_for_training", + output_dir="/pscratch/sd/s/smprince/projects/metfish/model_outputs/", + fast_dev_run=True,) \ No newline at end of file From 4bb3a6b29239cde5fbb70178f72cada4ad766c8b Mon Sep 17 00:00:00 2001 From: Steph Prince <40640337+stephprince@users.noreply.github.com> Date: Fri, 14 Mar 2025 13:23:40 -0700 Subject: [PATCH 2/7] clean up older code --- src/metfish/refinement_model/train.py | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/src/metfish/refinement_model/train.py b/src/metfish/refinement_model/train.py index 1e654a3..3a4e718 100644 --- a/src/metfish/refinement_model/train.py +++ b/src/metfish/refinement_model/train.py @@ -136,9 +136,7 @@ def main(data_dir="/global/cfs/cdirs/m3513/metfish/PDB70_verB_fixed_data/result" # fit the model trainer.fit(model=refinement_model, train_dataloaders=train_loader, ckpt_path=ckpt_path) - print('done') - if __name__ == "__main__": - main(data_dir="/global/cfs/cdirs/m3513/metfish/NMR_training/data_for_training", - output_dir="/pscratch/sd/s/smprince/projects/metfish/model_outputs/", - fast_dev_run=True,) \ No newline at end of file + args = parser.parse_args() + args_dict = vars(args) + main(**args_dict) From 838a07800e90875cc631f4877e455027edb76eb1 Mon Sep 17 00:00:00 2001 From: Steph Prince <40640337+stephprince@users.noreply.github.com> Date: Fri, 14 Mar 2025 13:57:29 -0700 Subject: [PATCH 3/7] add comments --- .../refinement_model/refinement_model.py | 23 +++++++++---------- src/metfish/refinement_model/train.py | 8 +++---- 2 files changed, 15 insertions(+), 16 deletions(-) diff --git a/src/metfish/refinement_model/refinement_model.py b/src/metfish/refinement_model/refinement_model.py index eda3325..59fc746 100644 --- a/src/metfish/refinement_model/refinement_model.py +++ b/src/metfish/refinement_model/refinement_model.py @@ -11,7 +11,7 @@ class SAXSLoss(torch.nn.Module): - """SAXS """ + """SAXS loss for MSA refinement""" def __init__(self, config): super(SAXSLoss, self).__init__() self.config = config @@ -41,12 +41,9 @@ def __init__(self, config, training=True): self.af_model = AlphaFold(config) self.training = training + # freeze AF parameters to allow gradient flow through AF model for param in self.af_model.parameters(): param.requires_grad = False - - # initialize parameters - self.w = None - self.b = None def initialize_parameters(self, msa): self.w = torch.nn.Parameter(torch.ones_like(msa)) @@ -69,7 +66,7 @@ def forward(self, batch): # define the lightning module for training class MSARefinementModelWrapper(pl.LightningModule): - def __init__(self, config, training=True, lr_mul=1.0, lr_add=0.05): + def __init__(self, config, training=True, num_iterations=100, lr_mul=1.0, lr_add=0.05): super().__init__() self.save_hyperparameters() self.config = config @@ -78,13 +75,15 @@ def __init__(self, config, training=True, lr_mul=1.0, lr_add=0.05): if training: self.loss = SAXSLoss(config.loss) self.cached_weights = None - self.lr_mul = 1.0 - self.lr_add = 0.05 + + # initial learning rates from Fadini et al. + self.lr_mul = lr_mul + self.lr_add = lr_add + self.num_iterations = num_iterations # activate manual optimization self.automatic_optimization = False - - self.num_iterations = 100 + self.last_log_time = time.time() # TODO - initialize MSA refinement parameters as part of config file @@ -120,7 +119,7 @@ def training_step(self, batch): loss = self.loss(outputs, batch_no_recycling) # backwards pass and update weights - self.manual_backward(loss, retain_graph=True) + self.manual_backward(loss, retain_graph=True) # retain graph to use same computation graph multiple times opt.step() # log the loss @@ -128,7 +127,7 @@ def training_step(self, batch): return loss - def configure_optimizers(self, learning_rate: float = 1e-3, eps: float = 1e-5,): + def configure_optimizers(self, eps: float = 1e-5,): optimizer = torch.optim.Adam([ {'params': [self.model.w], 'lr': self.lr_mul}, {'params': [self.model.b], 'lr': self.lr_add} diff --git a/src/metfish/refinement_model/train.py b/src/metfish/refinement_model/train.py index 3a4e718..018c938 100644 --- a/src/metfish/refinement_model/train.py +++ b/src/metfish/refinement_model/train.py @@ -4,6 +4,7 @@ import os import pytorch_lightning as pl +from pathlib import Path from pytorch_lightning.loggers import CSVLogger, WandbLogger from torch.utils.data import DataLoader @@ -107,18 +108,17 @@ def main(data_dir="/global/cfs/cdirs/m3513/metfish/PDB70_verB_fixed_data/result" refinement_model = MSARefinementModelWrapper(config) # add logging - loggers = [CSVLogger(f"{output_dir}/lightning_logs", name="msasaxs")] + Path(f"{output_dir}/lightning_logs/{job_name}").mkdir(parents=True, exist_ok=True) + loggers = [CSVLogger(f"{output_dir}/lightning_logs/{job_name}", name=job_name)] if use_wandb: os.environ["WANDB__SERVICE_WAIT"] = "300" os.environ["WANDB_MODE"] = "offline" - # wandb.init() - loggers.append(WandbLogger(name="msasaxs", save_dir=f"{output_dir}/lightning_logs")) + loggers.append(WandbLogger(name=job_name, save_dir=f"{output_dir}/lightning_logs/{job_name}")) # initialize trainer trainer = pl.Trainer(accelerator="gpu", strategy=strategy, max_epochs=max_epochs, - limit_train_batches=1.0, logger=loggers, log_every_n_steps=log_every_n_steps, default_root_dir=output_dir, From d1c24b125f28ba3a4d72316e5c8b2922fc6e1837 Mon Sep 17 00:00:00 2001 From: Steph Prince <40640337+stephprince@users.noreply.github.com> Date: Mon, 17 Mar 2025 16:30:14 -0700 Subject: [PATCH 4/7] update model initialization --- src/metfish/refinement_model/config.py | 627 ++++++++++++++++++ .../refinement_model/refinement_model.py | 22 +- src/metfish/refinement_model/train.py | 4 +- 3 files changed, 642 insertions(+), 11 deletions(-) create mode 100644 src/metfish/refinement_model/config.py diff --git a/src/metfish/refinement_model/config.py b/src/metfish/refinement_model/config.py new file mode 100644 index 0000000..d5897d3 --- /dev/null +++ b/src/metfish/refinement_model/config.py @@ -0,0 +1,627 @@ +import copy +import importlib +import ml_collections as mlc + + +def set_inf(c, inf): + for k, v in c.items(): + if isinstance(v, mlc.ConfigDict): + set_inf(v, inf) + elif k == "inf": + c[k] = inf + + +def enforce_config_constraints(config): + def string_to_setting(s): + path = s.split('.') + setting = config + for p in path: + setting = setting[p] + + return setting + + mutually_exclusive_bools = [ + ( + "model.template.average_templates", + "model.template.offload_templates" + ), + ( + "globals.use_lma", + "globals.use_flash", + ), + ] + + for s1, s2 in mutually_exclusive_bools: + s1_setting = string_to_setting(s1) + s2_setting = string_to_setting(s2) + if(s1_setting and s2_setting): + raise ValueError(f"Only one of {s1} and {s2} may be set at a time") + + fa_is_installed = importlib.util.find_spec("flash_attn") is not None + if(config.globals.use_flash and not fa_is_installed): + raise ValueError("use_flash requires that FlashAttention is installed") + + if( + config.globals.offload_inference and + not config.model.template.average_templates + ): + config.model.template.offload_templates = True + + +def model_config( + name, + train=False, + low_prec=False, + long_sequence_inference=False +): + c = copy.deepcopy(config) + # TRAINING PRESETS + if name == "initial_training": + # AF2 Suppl. Table 4, "initial training" setting + pass + elif name == "finetuning": + # AF2 Suppl. Table 4, "finetuning" setting + c.data.train.crop_size = 384 + c.data.train.max_extra_msa = 5120 + c.data.train.max_msa_clusters = 512 + c.loss.violation.weight = 1. + c.loss.experimentally_resolved.weight = 0.01 + elif name == "finetuning_ptm": + c.data.train.max_extra_msa = 5120 + c.data.train.crop_size = 384 + c.data.train.max_msa_clusters = 512 + c.loss.violation.weight = 1. + c.loss.experimentally_resolved.weight = 0.01 + c.model.heads.tm.enabled = True + c.loss.tm.weight = 0.1 + elif name == "finetuning_no_templ": + # AF2 Suppl. Table 4, "finetuning" setting + c.data.train.crop_size = 384 + c.data.train.max_extra_msa = 5120 + c.data.train.max_msa_clusters = 512 + c.model.template.enabled = False + c.loss.violation.weight = 1. + c.loss.experimentally_resolved.weight = 0.01 + elif name == "finetuning_no_templ_ptm": + # AF2 Suppl. Table 4, "finetuning" setting + c.data.train.crop_size = 384 + c.data.train.max_extra_msa = 5120 + c.data.train.max_msa_clusters = 512 + c.model.template.enabled = False + c.loss.violation.weight = 1. + c.loss.experimentally_resolved.weight = 0.01 + c.model.heads.tm.enabled = True + c.loss.tm.weight = 0.1 + # INFERENCE PRESETS + elif name == "model_1": + # AF2 Suppl. Table 5, Model 1.1.1 + c.data.train.max_extra_msa = 5120 + c.data.predict.max_extra_msa = 5120 + c.data.common.reduce_max_clusters_by_max_templates = True + c.data.common.use_templates = True + c.data.common.use_template_torsion_angles = True + c.model.template.enabled = True + elif name == "model_2": + # AF2 Suppl. Table 5, Model 1.1.2 + c.data.common.reduce_max_clusters_by_max_templates = True + c.data.common.use_templates = True + c.data.common.use_template_torsion_angles = True + c.model.template.enabled = True + elif name == "model_3": + # AF2 Suppl. Table 5, Model 1.2.1 + c.data.train.max_extra_msa = 5120 + c.data.predict.max_extra_msa = 5120 + c.model.template.enabled = False + elif name == "model_4": + # AF2 Suppl. Table 5, Model 1.2.2 + c.data.train.max_extra_msa = 5120 + c.data.predict.max_extra_msa = 5120 + c.model.template.enabled = False + elif name == "model_5": + # AF2 Suppl. Table 5, Model 1.2.3 + c.model.template.enabled = False + elif name == "model_1_ptm": + c.data.train.max_extra_msa = 5120 + c.data.predict.max_extra_msa = 5120 + c.data.common.reduce_max_clusters_by_max_templates = True + c.data.common.use_templates = True + c.data.common.use_template_torsion_angles = True + c.model.template.enabled = True + c.model.heads.tm.enabled = True + c.loss.tm.weight = 0.1 + elif name == "model_2_ptm": + c.data.common.reduce_max_clusters_by_max_templates = True + c.data.common.use_templates = True + c.data.common.use_template_torsion_angles = True + c.model.template.enabled = True + c.model.heads.tm.enabled = True + c.loss.tm.weight = 0.1 + elif name == "model_3_ptm": + c.data.train.max_extra_msa = 5120 + c.data.predict.max_extra_msa = 5120 + c.model.template.enabled = False + c.model.heads.tm.enabled = True + c.loss.tm.weight = 0.1 + elif name == "model_4_ptm": + c.data.train.max_extra_msa = 5120 + c.data.predict.max_extra_msa = 5120 + c.model.template.enabled = False + c.model.heads.tm.enabled = True + c.loss.tm.weight = 0.1 + elif name == "model_5_ptm": + c.model.template.enabled = False + c.model.heads.tm.enabled = True + c.loss.tm.weight = 0.1 + else: + raise ValueError("Invalid model name") + + if long_sequence_inference: + assert(not train) + c.globals.offload_inference = True + c.globals.use_lma = True + c.globals.use_flash = False + c.model.template.offload_inference = True + c.model.template.template_pair_stack.tune_chunk_size = False + c.model.extra_msa.extra_msa_stack.tune_chunk_size = False + c.model.evoformer_stack.tune_chunk_size = False + + if train: + c.globals.blocks_per_ckpt = 1 + c.globals.chunk_size = None + c.globals.use_lma = False + c.globals.offload_inference = False + c.model.template.average_templates = False + c.model.template.offload_templates = False + + if low_prec: + c.globals.eps = 1e-4 + # If we want exact numerical parity with the original, inf can't be + # a global constant + set_inf(c, 1e4) + + enforce_config_constraints(c) + + return c + + +c_z = mlc.FieldReference(128, field_type=int) +c_m = mlc.FieldReference(256, field_type=int) +c_t = mlc.FieldReference(64, field_type=int) +c_e = mlc.FieldReference(64, field_type=int) +c_s = mlc.FieldReference(384, field_type=int) +c_x = mlc.FieldReference(512, field_type=int) +blocks_per_ckpt = mlc.FieldReference(None, field_type=int) +chunk_size = mlc.FieldReference(4, field_type=int) +aux_distogram_bins = mlc.FieldReference(64, field_type=int) +tm_enabled = mlc.FieldReference(False, field_type=bool) +eps = mlc.FieldReference(1e-8, field_type=float) +templates_enabled = mlc.FieldReference(True, field_type=bool) +embed_template_torsion_angles = mlc.FieldReference(True, field_type=bool) +tune_chunk_size = mlc.FieldReference(True, field_type=bool) + +NUM_RES = "num residues placeholder" +NUM_MSA_SEQ = "msa placeholder" +NUM_EXTRA_SEQ = "extra msa placeholder" +NUM_TEMPLATES = "num templates placeholder" + +config = mlc.ConfigDict( + { + "data": { + "common": { + "feat": { + "aatype": [NUM_RES], + "all_atom_mask": [NUM_RES, None], + "all_atom_positions": [NUM_RES, None, None], + "alt_chi_angles": [NUM_RES, None], + "atom14_alt_gt_exists": [NUM_RES, None], + "atom14_alt_gt_positions": [NUM_RES, None, None], + "atom14_atom_exists": [NUM_RES, None], + "atom14_atom_is_ambiguous": [NUM_RES, None], + "atom14_gt_exists": [NUM_RES, None], + "atom14_gt_positions": [NUM_RES, None, None], + "atom37_atom_exists": [NUM_RES, None], + "backbone_rigid_mask": [NUM_RES], + "backbone_rigid_tensor": [NUM_RES, None, None], + "bert_mask": [NUM_MSA_SEQ, NUM_RES], + "chi_angles_sin_cos": [NUM_RES, None, None], + "chi_mask": [NUM_RES, None], + "extra_deletion_value": [NUM_EXTRA_SEQ, NUM_RES], + "extra_has_deletion": [NUM_EXTRA_SEQ, NUM_RES], + "extra_msa": [NUM_EXTRA_SEQ, NUM_RES], + "extra_msa_mask": [NUM_EXTRA_SEQ, NUM_RES], + "extra_msa_row_mask": [NUM_EXTRA_SEQ], + "is_distillation": [], + "msa_feat": [NUM_MSA_SEQ, NUM_RES, None], + "msa_mask": [NUM_MSA_SEQ, NUM_RES], + "msa_row_mask": [NUM_MSA_SEQ], + "no_recycling_iters": [], + "pseudo_beta": [NUM_RES, None], + "pseudo_beta_mask": [NUM_RES], + "residue_index": [NUM_RES], + "residx_atom14_to_atom37": [NUM_RES, None], + "residx_atom37_to_atom14": [NUM_RES, None], + "resolution": [], + "rigidgroups_alt_gt_frames": [NUM_RES, None, None, None], + "rigidgroups_group_exists": [NUM_RES, None], + "rigidgroups_group_is_ambiguous": [NUM_RES, None], + "rigidgroups_gt_exists": [NUM_RES, None], + "rigidgroups_gt_frames": [NUM_RES, None, None, None], + "saxs": [512], + "seq_length": [], + "seq_mask": [NUM_RES], + "target_feat": [NUM_RES, None], + "template_aatype": [NUM_TEMPLATES, NUM_RES], + "template_all_atom_mask": [NUM_TEMPLATES, NUM_RES, None], + "template_all_atom_positions": [ + NUM_TEMPLATES, NUM_RES, None, None, + ], + "template_alt_torsion_angles_sin_cos": [ + NUM_TEMPLATES, NUM_RES, None, None, + ], + "template_backbone_rigid_mask": [NUM_TEMPLATES, NUM_RES], + "template_backbone_rigid_tensor": [ + NUM_TEMPLATES, NUM_RES, None, None, + ], + "template_mask": [NUM_TEMPLATES], + "template_pseudo_beta": [NUM_TEMPLATES, NUM_RES, None], + "template_pseudo_beta_mask": [NUM_TEMPLATES, NUM_RES], + "template_sum_probs": [NUM_TEMPLATES, None], + "template_torsion_angles_mask": [ + NUM_TEMPLATES, NUM_RES, None, + ], + "template_torsion_angles_sin_cos": [ + NUM_TEMPLATES, NUM_RES, None, None, + ], + "true_msa": [NUM_MSA_SEQ, NUM_RES], + "use_clamped_fape": [], + }, + "masked_msa": { + "profile_prob": 0.1, + "same_prob": 0.1, + "uniform_prob": 0.1, + }, + "max_recycling_iters": 3, + "msa_cluster_features": True, + "reduce_msa_clusters_by_max_templates": False, + "resample_msa_in_recycling": True, + "template_features": [ + "template_all_atom_positions", + "template_sum_probs", + "template_aatype", + "template_all_atom_mask", + ], + "unsupervised_features": [ + "aatype", + "residue_index", + "msa", + "num_alignments", + "seq_length", + "between_segment_residues", + "deletion_matrix", + "no_recycling_iters", + "saxs" + ], + "use_templates": templates_enabled, + "use_template_torsion_angles": embed_template_torsion_angles, + }, + "supervised": { + "clamp_prob": 0.9, + "supervised_features": [ + "all_atom_mask", + "all_atom_positions", + "resolution", + "use_clamped_fape", + "is_distillation", + ], + }, + "predict": { + "fixed_size": True, + "subsample_templates": False, # We want top templates. + "masked_msa_replace_fraction": 0.15, + "max_msa_clusters": 512, + "max_extra_msa": 1024, + "max_template_hits": 4, + "max_templates": 4, + "crop": False, + "crop_size": None, + "supervised": False, + "uniform_recycling": False, + }, + "eval": { + "fixed_size": True, + "subsample_templates": False, # We want top templates. + "masked_msa_replace_fraction": 0.15, + "max_msa_clusters": 128, + "max_extra_msa": 1024, + "max_template_hits": 4, + "max_templates": 4, + "crop": False, + "crop_size": None, + "supervised": True, + "uniform_recycling": False, + }, + "train": { + "fixed_size": True, + "subsample_templates": True, + "masked_msa_replace_fraction": 0.15, + "max_msa_clusters": 128, + "max_extra_msa": 1024, + "max_template_hits": 4, + "max_templates": 4, + "shuffle_top_k_prefiltered": 20, + "crop": True, + "crop_size": 256, + "supervised": True, + "clamp_prob": 0.9, + "max_distillation_msa_clusters": 1000, + "uniform_recycling": True, + "distillation_prob": 0.75, + }, + "data_module": { + "use_small_bfd": False, + "data_loaders": { + "batch_size": 1, + "num_workers": 16, + "pin_memory": True, + }, + }, + }, + # Recurring FieldReferences that can be changed globally here + "globals": { + "blocks_per_ckpt": blocks_per_ckpt, + "chunk_size": chunk_size, + # Use Staats & Rabe's low-memory attention algorithm. Mutually + # exclusive with use_flash. + "use_lma": False, + # Use FlashAttention in selected modules. Mutually exclusive with + # use_lma. Doesn't work that well on long sequences (>1000 residues). + "use_flash": False, + "offload_inference": False, + "c_z": c_z, + "c_m": c_m, + "c_t": c_t, + "c_e": c_e, + "c_s": c_s, + "eps": eps, + }, + "model": { + "_mask_trans": False, + "input_embedder": { + "tf_dim": 22, + "msa_dim": 49, + "c_z": c_z, + "c_m": c_m, + "relpos_k": 32, + }, + "recycling_embedder": { + "c_z": c_z, + "c_m": c_m, + "min_bin": 3.25, + "max_bin": 20.75, + "no_bins": 15, + "inf": 1e8, + }, + "template": { + "distogram": { + "min_bin": 3.25, + "max_bin": 50.75, + "no_bins": 39, + }, + "template_angle_embedder": { + # DISCREPANCY: c_in is supposed to be 51. + "c_in": 57, + "c_out": c_m, + }, + "template_pair_embedder": { + "c_in": 88, + "c_out": c_t, + }, + "template_pair_stack": { + "c_t": c_t, + # DISCREPANCY: c_hidden_tri_att here is given in the supplement + # as 64. In the code, it's 16. + "c_hidden_tri_att": 16, + "c_hidden_tri_mul": 64, + "no_blocks": 2, + "no_heads": 4, + "pair_transition_n": 2, + "dropout_rate": 0.25, + "blocks_per_ckpt": blocks_per_ckpt, + "tune_chunk_size": tune_chunk_size, + "inf": 1e9, + }, + "template_pointwise_attention": { + "c_t": c_t, + "c_z": c_z, + # DISCREPANCY: c_hidden here is given in the supplement as 64. + # It's actually 16. + "c_hidden": 16, + "no_heads": 4, + "inf": 1e5, # 1e9, + }, + "inf": 1e5, # 1e9, + "eps": eps, # 1e-6, + "enabled": templates_enabled, + "embed_angles": embed_template_torsion_angles, + "use_unit_vector": False, + # Approximate template computation, saving memory. + # In our experiments, results are equivalent to or better than + # the stock implementation. Should be enabled for all new + # training runs. + "average_templates": False, + # Offload template embeddings to CPU memory. Vastly reduced + # memory consumption at the cost of a modest increase in + # runtime. Useful for inference on very long sequences. + # Mutually exclusive with average_templates. Automatically + # enabled if offload_inference is set. + "offload_templates": False, + }, + "extra_msa": { + "extra_msa_embedder": { + "c_in": 25, + "c_out": c_e, + }, + "extra_msa_stack": { + "c_m": c_e, + "c_z": c_z, + "c_hidden_msa_att": 8, + "c_hidden_opm": 32, + "c_hidden_mul": 128, + "c_hidden_pair_att": 32, + "no_heads_msa": 8, + "no_heads_pair": 4, + "no_blocks": 4, + "transition_n": 4, + "msa_dropout": 0.15, + "pair_dropout": 0.25, + "clear_cache_between_blocks": False, + "tune_chunk_size": tune_chunk_size, + "inf": 1e9, + "eps": eps, # 1e-10, + "ckpt": blocks_per_ckpt is not None, + }, + "enabled": True, + }, + "evoformer_stack": { + "c_m": c_m, + "c_z": c_z, + "c_hidden_msa_att": 32, + "c_hidden_opm": 32, + "c_hidden_mul": 128, + "c_hidden_pair_att": 32, + "c_s": c_s, + "no_heads_msa": 8, + "no_heads_pair": 4, + "no_blocks": 48, + "transition_n": 4, + "msa_dropout": 0.15, + "pair_dropout": 0.25, + "blocks_per_ckpt": blocks_per_ckpt, + "clear_cache_between_blocks": False, + "tune_chunk_size": tune_chunk_size, + "inf": 1e9, + "eps": eps, # 1e-10, + }, + "structure_module": { + "c_s": c_s, + "c_z": c_z, + "c_ipa": 16, + "c_resnet": 128, + "no_heads_ipa": 12, + "no_qk_points": 4, + "no_v_points": 8, + "dropout_rate": 0.1, + "no_blocks": 8, + "no_transition_layers": 1, + "no_resnet_blocks": 2, + "no_angles": 7, + "trans_scale_factor": 10, + "epsilon": eps, # 1e-12, + "inf": 1e5, + }, + "heads": { + "lddt": { + "no_bins": 50, + "c_in": c_s, + "c_hidden": 128, + }, + "distogram": { + "c_z": c_z, + "no_bins": aux_distogram_bins, + }, + "tm": { + "c_z": c_z, + "no_bins": aux_distogram_bins, + "enabled": tm_enabled, + }, + "masked_msa": { + "c_m": c_m, + "c_out": 23, + }, + "experimentally_resolved": { + "c_s": c_s, + "c_out": 37, + }, + }, + }, + "relax": { + "max_iterations": 0, # no max + "tolerance": 2.39, + "stiffness": 10.0, + "max_outer_iterations": 20, + "exclude_residues": [], + }, + "loss": { + "distogram": { + "min_bin": 2.3125, + "max_bin": 21.6875, + "no_bins": 64, + "eps": eps, # 1e-6, + "weight": 0.3, + }, + "experimentally_resolved": { + "eps": eps, # 1e-8, + "min_resolution": 0.1, + "max_resolution": 3.0, + "weight": 0.0, + }, + "fape": { + "backbone": { + "clamp_distance": 10.0, + "loss_unit_distance": 10.0, + "weight": 0.5, + }, + "sidechain": { + "clamp_distance": 10.0, + "length_scale": 10.0, + "weight": 0.5, + }, + "eps": 1e-4, + "weight": 1.0, + }, + "plddt_loss": { + "min_resolution": -1.0, # allows examples with unknown resolution + "max_resolution": 3.0, + "cutoff": 15.0, + "no_bins": 50, + "eps": eps, # 1e-10, + "weight": 0.01, + }, + "masked_msa": { + "eps": eps, # 1e-8, + "weight": 2.0, + }, + "supervised_chi": { + "chi_weight": 0.5, + "angle_norm_weight": 0.01, + "eps": eps, # 1e-6, + "weight": 1.0, + }, + "violation": { + "violation_tolerance_factor": 12.0, + "clash_overlap_tolerance": 1.5, + "eps": eps, # 1e-6, + "weight": 0.0, + }, + "tm": { + "max_bin": 31, + "no_bins": 64, + "min_resolution": 0.1, + "max_resolution": 3.0, + "eps": eps, # 1e-8, + "weight": 0., + "enabled": tm_enabled, + }, + "saxs_loss": { + "use_l1": False, + "dmax": 256, # pad to 512 for input data, use 512/step + "step": 0.5, + "eps": eps, # 1e-10, + "weight": 5.0 + }, + "eps": eps, + "saxs_loss_only": False, + }, + "ema": {"decay": 0.999}, + } +) \ No newline at end of file diff --git a/src/metfish/refinement_model/refinement_model.py b/src/metfish/refinement_model/refinement_model.py index 59fc746..3c181f3 100644 --- a/src/metfish/refinement_model/refinement_model.py +++ b/src/metfish/refinement_model/refinement_model.py @@ -45,6 +45,8 @@ def __init__(self, config, training=True): for param in self.af_model.parameters(): param.requires_grad = False + self.initialize_parameters(torch.ones((1))) + def initialize_parameters(self, msa): self.w = torch.nn.Parameter(torch.ones_like(msa)) self.b = torch.nn.Parameter(torch.zeros_like(msa)) @@ -56,8 +58,8 @@ def forward(self, batch): # refine msa with linear layer # TODO - msa cluster profile here may mean the msa features after embedding... need to modify if so - msa_refined = self.w * batch['msa_feat'] + self.b - batch['msa_feat'] = msa_refined + msa_feat_refined = self.w * batch['msa_feat'] + self.b + batch.update({'msa_feat': msa_feat_refined}) # run through alphafold outputs = self.af_model(batch) @@ -83,12 +85,8 @@ def __init__(self, config, training=True, num_iterations=100, lr_mul=1.0, lr_add # activate manual optimization self.automatic_optimization = False - self.last_log_time = time.time() - # TODO - initialize MSA refinement parameters as part of config file - self.model.initialize_parameters(torch.zeros(256)) - def forward(self, batch): return self.model(batch) @@ -102,7 +100,6 @@ def _log(self, loss, iter=None, train=True): self.last_log_time = time.time() def training_step(self, batch): - self.model.initialize_parameters(batch['msa_feat']) opt = self.optimizers() for n in range(self.num_iterations): @@ -111,15 +108,18 @@ def training_step(self, batch): # clear gradients opt.zero_grad() + # create new copy of a batch for each iteration + iteration_batch = {k: v for k, v in batch.items()} + # forward pass - outputs = self.model(batch) + outputs = self.model(iteration_batch) # calculate loss batch_no_recycling = tensor_tree_map(lambda t: t[..., -1], batch) # remove recycling dimension loss = self.loss(outputs, batch_no_recycling) # backwards pass and update weights - self.manual_backward(loss, retain_graph=True) # retain graph to use same computation graph multiple times + self.manual_backward(loss) opt.step() # log the loss @@ -127,6 +127,10 @@ def training_step(self, batch): return loss + def on_train_batch_start(self, batch, batch_idx): + self.model.initialize_parameters(batch['msa_feat']) + self.trainer.strategy.setup_optimizers(self.trainer) # reset optimizer based on batch msa feature size + def configure_optimizers(self, eps: float = 1e-5,): optimizer = torch.optim.Adam([ {'params': [self.model.w], 'lr': self.lr_mul}, diff --git a/src/metfish/refinement_model/train.py b/src/metfish/refinement_model/train.py index 018c938..8004d65 100644 --- a/src/metfish/refinement_model/train.py +++ b/src/metfish/refinement_model/train.py @@ -74,7 +74,7 @@ def main(data_dir="/global/cfs/cdirs/m3513/metfish/PDB70_verB_fixed_data/result" ckpt_path=None, gpus_per_node=1, num_nodes=1, - batch_size=2, + batch_size=1, seed=1, use_wandb=True, fast_dev_run=False, @@ -82,7 +82,7 @@ def main(data_dir="/global/cfs/cdirs/m3513/metfish/PDB70_verB_fixed_data/result" resume_from_ckpt=False, precision='bf16-mixed', max_epochs=100, - log_every_n_steps=25, + log_every_n_steps=1, job_name='default', ): From 3aee9149c3445afddfb50a018b96d68acd9311bb Mon Sep 17 00:00:00 2001 From: Steph Prince <40640337+stephprince@users.noreply.github.com> Date: Tue, 25 Mar 2025 14:53:19 -0700 Subject: [PATCH 5/7] update optimization model --- .../refinement_model/refinement_model.py | 94 +------------ src/metfish/refinement_model/train.py | 133 ++++++++++-------- src/metfish/utils.py | 22 ++- 3 files changed, 102 insertions(+), 147 deletions(-) diff --git a/src/metfish/refinement_model/refinement_model.py b/src/metfish/refinement_model/refinement_model.py index 3c181f3..a87a8ba 100644 --- a/src/metfish/refinement_model/refinement_model.py +++ b/src/metfish/refinement_model/refinement_model.py @@ -1,10 +1,5 @@ -import time import torch -import pytorch_lightning as pl - -from openfold.utils.import_weights import import_jax_weights_ -from openfold.utils.tensor_utils import tensor_tree_map from openfold.model.model import AlphaFold from metfish.msa_model.utils.loss import compute_saxs @@ -40,106 +35,23 @@ def __init__(self, config, training=True): self.config = config self.af_model = AlphaFold(config) self.training = training + if training: + self.loss = SAXSLoss(config.loss) # freeze AF parameters to allow gradient flow through AF model for param in self.af_model.parameters(): param.requires_grad = False - self.initialize_parameters(torch.ones((1))) - def initialize_parameters(self, msa): self.w = torch.nn.Parameter(torch.ones_like(msa)) self.b = torch.nn.Parameter(torch.zeros_like(msa)) def forward(self, batch): - device = batch['aatype'].device - self.w = self.w.to(device) - self.b = self.b.to(device) - # refine msa with linear layer - # TODO - msa cluster profile here may mean the msa features after embedding... need to modify if so msa_feat_refined = self.w * batch['msa_feat'] + self.b batch.update({'msa_feat': msa_feat_refined}) # run through alphafold outputs = self.af_model(batch) - return outputs - -# define the lightning module for training -class MSARefinementModelWrapper(pl.LightningModule): - def __init__(self, config, training=True, num_iterations=100, lr_mul=1.0, lr_add=0.05): - super().__init__() - self.save_hyperparameters() - self.config = config - self.model = MSARefinementModel(config) - self.training = training - if training: - self.loss = SAXSLoss(config.loss) - self.cached_weights = None - - # initial learning rates from Fadini et al. - self.lr_mul = lr_mul - self.lr_add = lr_add - self.num_iterations = num_iterations - - # activate manual optimization - self.automatic_optimization = False - self.last_log_time = time.time() - - def forward(self, batch): - return self.model(batch) - - def _log(self, loss, iter=None, train=True): - self.log( - f"loss_iter{iter}", - loss, - on_step=False, on_epoch=True, logger=True, sync_dist=True - ) - self.log('dur', time.time() - self.last_log_time, sync_dist=True) - self.last_log_time = time.time() - - def training_step(self, batch): - opt = self.optimizers() - - for n in range(self.num_iterations): - print(f'Running iteration {n} / {self.num_iterations}') - - # clear gradients - opt.zero_grad() - - # create new copy of a batch for each iteration - iteration_batch = {k: v for k, v in batch.items()} - - # forward pass - outputs = self.model(iteration_batch) - - # calculate loss - batch_no_recycling = tensor_tree_map(lambda t: t[..., -1], batch) # remove recycling dimension - loss = self.loss(outputs, batch_no_recycling) - - # backwards pass and update weights - self.manual_backward(loss) - opt.step() - - # log the loss - self._log(loss.detach(), iter=n) - - return loss - - def on_train_batch_start(self, batch, batch_idx): - self.model.initialize_parameters(batch['msa_feat']) - self.trainer.strategy.setup_optimizers(self.trainer) # reset optimizer based on batch msa feature size - - def configure_optimizers(self, eps: float = 1e-5,): - optimizer = torch.optim.Adam([ - {'params': [self.model.w], 'lr': self.lr_mul}, - {'params': [self.model.b], 'lr': self.lr_add} - ], eps=eps) - - return optimizer - - def load_from_jax(self, jax_path): - import_jax_weights_( - self.model.af_model, jax_path, version='model_3' - ) + return outputs \ No newline at end of file diff --git a/src/metfish/refinement_model/train.py b/src/metfish/refinement_model/train.py index 8004d65..6f485e2 100644 --- a/src/metfish/refinement_model/train.py +++ b/src/metfish/refinement_model/train.py @@ -2,15 +2,19 @@ import argparse import torch import os -import pytorch_lightning as pl +import lightning.pytorch as pl from pathlib import Path -from pytorch_lightning.loggers import CSVLogger, WandbLogger from torch.utils.data import DataLoader +from lightning.fabric import Fabric +from lightning.pytorch.loggers import CSVLogger, WandbLogger +from lightning.pytorch.callbacks import ModelCheckpoint +from openfold.utils.import_weights import import_jax_weights_ from metfish.msa_model.config import model_config from metfish.msa_model.data.data_modules import MSASAXSDataset -from metfish.refinement_model.refinement_model import MSARefinementModelWrapper +from metfish.refinement_model.refinement_model import MSARefinementModel +from metfish.refinement_model.model_wrapper import train # gives a speedup on Ampere-class GPUs torch.set_float32_matmul_precision("high") @@ -29,15 +33,6 @@ "--ckpt_path", type=str, help='''Path to a model checkpoint from which to resume training.''', ) -parser.add_argument( - "--gpus_per_node", type=int, default=1, help='Number of gpus per node (will use all 4 per node on perlmutter).' -) -parser.add_argument( - "--num_nodes", type=int, default=1, help='Number of nodes to use for training.' -) -parser.add_argument( - "--batch_size", type=int, default=2, help='Batch size for each training step' -) parser.add_argument( "--seed", type=int, default=1, help="Random seed" @@ -46,66 +41,65 @@ "--use_wandb", action="store_true", default=True, help="Whether to log metrics to Weights & Biases" ) -parser.add_argument( - "--fast_dev_run", default=False, action='store_true', - help="Whether to run a fast dev run of a single batch for testing purposes" -) parser.add_argument( "--jax_param_path", type=str, default="/pscratch/sd/s/smprince/projects/alphaflow/params_model_1.npz", # these are the original AF weights, help="""Path to an .npz JAX parameter file with which to initialize the model""" ) +parser.add_argument( + "--resume_from_ckpt", default=False, action='store_true', + help="Whether to use a model checkpoint from which to restore training state" +) parser.add_argument( "--precision", type=str, default='bf16-mixed', help='Sets precision, lower precision improves runtime performance.', ) parser.add_argument( - "--max_epochs", type=int, default=100, + "--job_name", type=str, + help='''Name of job to be used for logging purposes.''', ) parser.add_argument( - "--log_every_n_steps", type=int, default=25, + "--save_intermediate_pdb", default=False, action='store_true', + help="Whether to save intermediate pdb files for every step of the optimization" ) parser.add_argument( - "--job_name", type=str, - help='''Name of job to be used for logging purposes.''', + "--overwrite", default=False, action='store_true', + help="Whether to skip optimization if model checkpoints already exist" ) - -def main(data_dir="/global/cfs/cdirs/m3513/metfish/PDB70_verB_fixed_data/result", - output_dir="/pscratch/sd/s/smprince/projects/metfish/model_outputs", +def main(data_dir, + output_dir, ckpt_path=None, - gpus_per_node=1, - num_nodes=1, batch_size=1, seed=1, use_wandb=True, - fast_dev_run=False, jax_param_path="/pscratch/sd/s/smprince/projects/alphaflow/params_model_1.npz", resume_from_ckpt=False, precision='bf16-mixed', - max_epochs=100, - log_every_n_steps=1, - job_name='default', + job_name='optimization', + save_intermediate_pdb=False, + overwrite=False, ): # set up data paths and configuration - pdb_dir = f"{data_dir}/pdb" + pdb_dir = f"{data_dir}/pdbs" saxs_dir = f"{data_dir}/saxs_r" msa_dir = f"{data_dir}/msa" - csv_dir = f"{data_dir}/scripts" - training_csv = f'{csv_dir}/input_training.csv' # NOTE - this was msa_dir for training v_1 + training_csv = f'{data_dir}/input_no_training_data.csv' pl.seed_everything(seed, workers=True) - strategy = "ddp" if (gpus_per_node > 1) or num_nodes > 1 else "auto" config = model_config('initial_training', train=True, low_prec=True) data_config = config.data data_config.common.use_templates = False data_config.common.max_recycling_iters = 0 # set up training and test datasets and dataloaders - train_dataset = MSASAXSDataset(data_config, training_csv, msa_dir=msa_dir, saxs_dir=saxs_dir, pdb_dir=pdb_dir, saxs_ext='.pr.csv', pdb_prefix='') + train_dataset = MSASAXSDataset(data_config, training_csv, msa_dir=msa_dir, saxs_dir=saxs_dir, pdb_dir=pdb_dir, saxs_ext='_atom_only.csv', pdb_prefix='', pdb_ext='_atom_only.pdb') train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=8) - # initialize model - refinement_model = MSARefinementModelWrapper(config) + # initialize model and load existing weights if needed + refinement_model = MSARefinementModel(config) + if jax_param_path and not resume_from_ckpt: + import_jax_weights_(refinement_model.af_model, jax_param_path, version='model_3') + print(f"Successfully loaded JAX parameters at {jax_param_path}...") # add logging Path(f"{output_dir}/lightning_logs/{job_name}").mkdir(parents=True, exist_ok=True) @@ -115,26 +109,55 @@ def main(data_dir="/global/cfs/cdirs/m3513/metfish/PDB70_verB_fixed_data/result" os.environ["WANDB_MODE"] = "offline" loggers.append(WandbLogger(name=job_name, save_dir=f"{output_dir}/lightning_logs/{job_name}")) - # initialize trainer - trainer = pl.Trainer(accelerator="gpu", - strategy=strategy, - max_epochs=max_epochs, - logger=loggers, - log_every_n_steps=log_every_n_steps, - default_root_dir=output_dir, - devices=gpus_per_node, - num_nodes=num_nodes, - fast_dev_run=fast_dev_run, - precision=precision, - ) - - # load existing weights - if jax_param_path and not resume_from_ckpt: - refinement_model.load_from_jax(jax_param_path) - print(f"Successfully loaded JAX parameters at {jax_param_path}...") + # add checkpointing + Path(f"{output_dir}/checkpoints/{job_name}").mkdir(parents=True, exist_ok=True) + callbacks = [ModelCheckpoint(dirpath=f"{output_dir}/checkpoints/{job_name}")] + + # configure fabric trainer + fabric = Fabric(accelerator="gpu", + devices=1, + loggers=loggers, + callbacks=callbacks, + precision=precision) + + # run training for each value in dataset + data_loader = fabric.setup_dataloaders(train_loader) + for i, batch in enumerate(data_loader): + # skip if file already exists + seq_name = train_dataset.get_name(int(batch['batch_idx'])) + ckpt_path = f"{output_dir}/checkpoints/{job_name}/model_{seq_name}.ckpt" + if not overwrite and os.path.exists(ckpt_path.replace('.ckpt', '_phase2.ckpt')): + print(f"Skipping {seq_name} as model checkpoint already exists.") + continue + + intermediate_output_path = None + if save_intermediate_pdb: + intermediate_output_path = f"{output_dir}/intermediate_files/{job_name}/model_{seq_name}" + Path(intermediate_output_path).mkdir(parents=True, exist_ok=True) + + # initialize parameters and optimizer for each sequence + refinement_model.initialize_parameters(batch['msa_feat']) + optimizers_phase1 = torch.optim.Adam([ + {'params': [refinement_model.w], 'lr': 1.0}, + {'params': [refinement_model.b], 'lr': 0.05} + ], eps=1e-5) + optimizers_phase2 = torch.optim.Adam([ + {'params': [refinement_model.w], 'lr': 1e-3}, + {'params': [refinement_model.b], 'lr': 1e-3} + ], eps=1e-5) + + model, optimizer1, optimizer2 = fabric.setup(refinement_model, optimizers_phase1, optimizers_phase2) + + if resume_from_ckpt: + state = {"model": model, "optimizer1": optimizer1, "optimizer2": optimizer2, "iter": 0} + fabric.load(ckpt_path, state) - # fit the model - trainer.fit(model=refinement_model, train_dataloaders=train_loader, ckpt_path=ckpt_path) + # run training + print(f'Running optimization for {seq_name}') + train(fabric, model, optimizer1, optimizer2, batch, + ckpt_path=ckpt_path, + early_stopping=False, + intermediate_output_path=intermediate_output_path) if __name__ == "__main__": args = parser.parse_args() diff --git a/src/metfish/utils.py b/src/metfish/utils.py index 4bb7a7e..cd8e38f 100644 --- a/src/metfish/utils.py +++ b/src/metfish/utils.py @@ -15,6 +15,10 @@ from prody import parsePDB, ANM, GNM, extendModel, traverseMode, writePDB +from openfold.np import residue_constants +from openfold.np.protein import Protein +from metfish.msa_model.utils.tensor_utils import tensor_tree_map + n_elec_df = {el.symbol: el.number for el in elements} amino_acids = [a.upper() for a in SeqUtils.IUPACData.protein_letters_3to1.keys()] @@ -378,4 +382,20 @@ def write_conformers(out_dir, name, protein, pdb_ext='.pdb'): filenames.append(filename) conf_idx += 1 - return filenames \ No newline at end of file + return filenames + +def output_to_protein(output): + """Returns the pbd (file) string from the model given the model output.""" + output = tensor_tree_map(lambda x: x.cpu().numpy(), output) + final_atom_positions = output['final_atom_positions'] + final_atom_mask = output["atom37_atom_exists"] + pred = Protein( + aatype=output["aatype"], + atom_positions=final_atom_positions[0], + atom_mask=final_atom_mask, + residue_index=output["residue_index"] + 1, + b_factors=np.repeat(output["plddt"][...,None], residue_constants.atom_type_num, axis=-1)[0], + chain_index=output["chain_index"] if "chain_index" in output else None, + ) + + return pred \ No newline at end of file From 258b6a7f35d1490aa07352173e450bf887945e9e Mon Sep 17 00:00:00 2001 From: Steph Prince <40640337+stephprince@users.noreply.github.com> Date: Tue, 25 Mar 2025 14:55:48 -0700 Subject: [PATCH 6/7] add fabric training function --- src/metfish/refinement_model/model_wrapper.py | 118 ++++++++++++++++++ 1 file changed, 118 insertions(+) create mode 100644 src/metfish/refinement_model/model_wrapper.py diff --git a/src/metfish/refinement_model/model_wrapper.py b/src/metfish/refinement_model/model_wrapper.py new file mode 100644 index 0000000..40bac4f --- /dev/null +++ b/src/metfish/refinement_model/model_wrapper.py @@ -0,0 +1,118 @@ + +from tqdm import tqdm +from pathlib import Path +from openfold.utils.tensor_utils import tensor_tree_map +from openfold.np import protein + +from metfish.utils import output_to_protein + + +def train(fabric, model, optimizer1, optimizer2, batch, + ckpt_path=None, + num_runs_phase_1=3, + num_iterations_phase1=100, + num_iterations_phase2=500, + early_stopping=True, + intermediate_output_path=None,): + + # setup training + model.train() + best_loss = float('inf') + intermediate_pdb_path = Path(f'{intermediate_output_path}/{Path(ckpt_path).stem}') if intermediate_output_path is not None else None + ckpt_path_phase_1 = ckpt_path.replace('.ckpt', '_phase1.ckpt') + ckpt_path_phase_2 = ckpt_path.replace('.ckpt', '_phase2.ckpt') + + # phase 1 training + for r in range(num_runs_phase_1): + model.initialize_parameters(batch['msa_feat']) + + for i in tqdm(range(num_iterations_phase1)): + + # clear gradientsj + optimizer1.zero_grad() + + # create new copy of a batch for each iteration + iteration_batch = {k: v for k, v in batch.items()} + + # forward pass + outputs = model(iteration_batch) + + # calculate loss + batch_no_recycling = tensor_tree_map(lambda t: t[0, ..., -1], batch) # remove recycling dimension + loss = model.loss(outputs, batch_no_recycling) + + # backwards pass and update weights + fabric.backward(loss) + optimizer1.step() + + # log the loss + fabric.log("loss", loss) + + if intermediate_pdb_path is not None: + pdb_path_output = f'{intermediate_pdb_path}_phase1_run_{r}_iter_{i}.pdb' + save_intermediate_optimization_steps({**outputs, **batch_no_recycling}, pdb_path_output) + + # save checkpoint if best so far + if loss < best_loss and ckpt_path_phase_1 is not None: + best_loss = loss + state = {"model": model, "optimizer1": optimizer1, "optimizer2": optimizer2, "iter": i} + fabric.save(ckpt_path_phase_1, state) + + # load checkpoint with best outcome + fabric.load(ckpt_path_phase_1, state) + + # phase 2 training + no_improvement_count = 0 + for i in tqdm(range(num_iterations_phase2)): + # clear gradients + optimizer2.zero_grad() + + # create new copy of a batch for each iteration + iteration_batch = {k: v for k, v in batch.items()} + + # forward pass + outputs = model(iteration_batch) + + # calculate loss + batch_no_recycling = tensor_tree_map(lambda t: t[0, ..., -1], batch) + loss = model.loss(outputs, batch_no_recycling) + + # backwards pass and update weights + fabric.backward(loss) + optimizer2.step() + + # log the loss + fabric.log("loss", loss) + + # early stopping check + if early_stopping: + min_delta = 0.1 + patience = 50 + if loss < best_loss - min_delta: + best_loss = loss + no_improvement_count = 0 + else: + no_improvement_count += 1 + if no_improvement_count >= patience: + break + + # save intermediate output check + if intermediate_pdb_path is not None: + pdb_path_output = f'{intermediate_pdb_path}_phase2_iter_{i}.pdb' + save_intermediate_optimization_steps({**outputs, **batch_no_recycling}, pdb_path_output) + + state = {"model": model, "optimizer1": optimizer1, "optimizer2": optimizer2, "iter": i} + fabric.save(ckpt_path_phase_2, state) + + return loss + + +def save_intermediate_optimization_steps(outputs, path): + # copy and detach relevant tensors + out_to_prot_keys = ['final_atom_positions', 'plddt', 'atom37_atom_exists', 'aatype', 'residue_index', 'chain_index'] + output_info = {k: v.clone().detach() for k, v in outputs.items() if k in out_to_prot_keys} + unrelaxed_protein = output_to_protein(output_info) + + # save intermediate output + with open(path, 'w') as f: + f.write(protein.to_pdb(unrelaxed_protein)) \ No newline at end of file From 7caba1d24e8967a4aa47e095076379643326dab6 Mon Sep 17 00:00:00 2001 From: Steph Prince <40640337+stephprince@users.noreply.github.com> Date: Fri, 11 Apr 2025 15:36:22 -0700 Subject: [PATCH 7/7] update logging --- src/metfish/refinement_model/model_wrapper.py | 17 ++++++++--------- src/metfish/refinement_model/train.py | 4 ++-- 2 files changed, 10 insertions(+), 11 deletions(-) diff --git a/src/metfish/refinement_model/model_wrapper.py b/src/metfish/refinement_model/model_wrapper.py index 40bac4f..8850eef 100644 --- a/src/metfish/refinement_model/model_wrapper.py +++ b/src/metfish/refinement_model/model_wrapper.py @@ -20,12 +20,11 @@ def train(fabric, model, optimizer1, optimizer2, batch, best_loss = float('inf') intermediate_pdb_path = Path(f'{intermediate_output_path}/{Path(ckpt_path).stem}') if intermediate_output_path is not None else None ckpt_path_phase_1 = ckpt_path.replace('.ckpt', '_phase1.ckpt') - ckpt_path_phase_2 = ckpt_path.replace('.ckpt', '_phase2.ckpt') + ckpt_path_phase_2 = ckpt_path.replace('.ckpt', '_phase2.ckpt') # phase 1 training for r in range(num_runs_phase_1): model.initialize_parameters(batch['msa_feat']) - for i in tqdm(range(num_iterations_phase1)): # clear gradientsj @@ -45,13 +44,13 @@ def train(fabric, model, optimizer1, optimizer2, batch, fabric.backward(loss) optimizer1.step() - # log the loss - fabric.log("loss", loss) - if intermediate_pdb_path is not None: + fabric.log(f"loss/{intermediate_pdb_path.stem}_phase1", loss) pdb_path_output = f'{intermediate_pdb_path}_phase1_run_{r}_iter_{i}.pdb' save_intermediate_optimization_steps({**outputs, **batch_no_recycling}, pdb_path_output) - + else: + fabric.log("loss/phase1", loss) + # save checkpoint if best so far if loss < best_loss and ckpt_path_phase_1 is not None: best_loss = loss @@ -81,9 +80,6 @@ def train(fabric, model, optimizer1, optimizer2, batch, fabric.backward(loss) optimizer2.step() - # log the loss - fabric.log("loss", loss) - # early stopping check if early_stopping: min_delta = 0.1 @@ -98,8 +94,11 @@ def train(fabric, model, optimizer1, optimizer2, batch, # save intermediate output check if intermediate_pdb_path is not None: + fabric.log(f"loss/{intermediate_pdb_path.stem}_phase2", loss) pdb_path_output = f'{intermediate_pdb_path}_phase2_iter_{i}.pdb' save_intermediate_optimization_steps({**outputs, **batch_no_recycling}, pdb_path_output) + else: + fabric.log("loss/phase2", loss) state = {"model": model, "optimizer1": optimizer1, "optimizer2": optimizer2, "iter": i} fabric.save(ckpt_path_phase_2, state) diff --git a/src/metfish/refinement_model/train.py b/src/metfish/refinement_model/train.py index 6f485e2..1550123 100644 --- a/src/metfish/refinement_model/train.py +++ b/src/metfish/refinement_model/train.py @@ -65,8 +65,8 @@ "--overwrite", default=False, action='store_true', help="Whether to skip optimization if model checkpoints already exist" ) -def main(data_dir, - output_dir, +def main(data_dir="/global/cfs/cdirs/m3513/metfish/apo_holo_data", + output_dir="/pscratch/sd/s/smprince/projects/metfish/model_outputs", ckpt_path=None, batch_size=1, seed=1,