Download model/PXDesignBench/ColabDesign/colabdesign/tr/model.py from OneScience-Group/PXDesign: direct link, hf CLI and curl.
- Browser
- Download file 12.6 kB
-
https://huggingface.co/OneScience-Group/PXDesign/resolve/main/model/PXDesignBench/ColabDesign/colabdesign/tr/model.py
- Command line
-
hf download hf://OneScience-Group/PXDesign/model/PXDesignBench/ColabDesign/colabdesign/tr/model.py
-
curl -L -o model.py https://huggingface.co/OneScience-Group/PXDesign/resolve/main/model/PXDesignBench/ColabDesign/colabdesign/tr/model.py
12.6 kB
| import random, os | |
| import numpy as np | |
| import jax | |
| import jax.numpy as jnp | |
| import matplotlib.pyplot as plt | |
| from colabdesign.shared.utils import copy_dict, update_dict, Key, dict_to_str | |
| from colabdesign.shared.prep import prep_pos | |
| from colabdesign.shared.protein import _np_get_6D_binned | |
| from colabdesign.shared.model import design_model, soft_seq | |
| from .trrosetta import TrRosetta, get_model_params | |
| # borrow some stuff from AfDesign | |
| from colabdesign.af.prep import prep_pdb | |
| from colabdesign.af.alphafold.common import protein | |
| class mk_tr_model(design_model): | |
| def __init__(self, protocol="fixbb", num_models=1, | |
| sample_models=True, data_dir="params/tr", | |
| optimizer="sgd", learning_rate=0.1, | |
| loss_callback=None): | |
| assert protocol in ["fixbb","hallucination","partial"] | |
| self.protocol = protocol | |
| self._data_dir = "." if os.path.isfile(os.path.join("models",f"model_xaa.npy")) else data_dir | |
| self._loss_callback = loss_callback | |
| self._num = 1 | |
| # set default options | |
| self.opt = {"temp":1.0, "soft":1.0, "hard":1.0, "dropout":False, | |
| "num_models":num_models,"sample_models":sample_models, | |
| "weights":{}, "lr":1.0, "alpha":1.0, | |
| "learning_rate":learning_rate, "use_pssm":False, | |
| "norm_seq_grad":True} | |
| self._args = {"optimizer":optimizer} | |
| self._params = {} | |
| self._inputs = {} | |
| # setup model | |
| self._model = self._get_model() | |
| self._model_params = [] | |
| for k in list("abcde"): | |
| p = os.path.join(self._data_dir,os.path.join("models",f"model_xa{k}.npy")) | |
| self._model_params.append(get_model_params(p)) | |
| if protocol in ["hallucination","partial"]: | |
| self._bkg_model = TrRosetta(bkg_model=True) | |
| def _get_model(self): | |
| runner = TrRosetta() | |
| def _get_loss(inputs, outputs): | |
| opt = inputs["opt"] | |
| aux = {"outputs":outputs, "losses":{}} | |
| log_p = jax.tree_util.tree_map(jax.nn.log_softmax, outputs) | |
| # bkg loss | |
| if self.protocol in ["hallucination","partial"]: | |
| p = jax.tree_util.tree_map(jax.nn.softmax, outputs) | |
| log_q = jax.tree_util.tree_map(jax.nn.log_softmax, inputs["6D_bkg"]) | |
| aux["losses"]["bkg"] = {} | |
| for k in ["dist","omega","theta","phi"]: | |
| aux["losses"]["bkg"][k] = -(p[k]*(log_p[k]-log_q[k])).sum(-1).mean() | |
| # cce loss | |
| if self.protocol in ["fixbb","partial"]: | |
| if "pos" in opt: | |
| pos = opt["pos"] | |
| log_p = jax.tree_util.tree_map(lambda x:x[:,pos][pos,:], log_p) | |
| q = inputs["6D"] | |
| aux["losses"]["cce"] = {} | |
| for k in ["dist","omega","theta","phi"]: | |
| aux["losses"]["cce"][k] = -(q[k]*log_p[k]).sum(-1).mean() | |
| if self._loss_callback is not None: | |
| aux["losses"].update(self._loss_callback(outputs)) | |
| # weighted loss | |
| w = opt["weights"] | |
| tree_multi = lambda x,y: jax.tree_util.tree_map(lambda a,b:a*b, x,y) | |
| losses = {k:(tree_multi(v,w[k]) if k in w else v) for k,v in aux["losses"].items()} | |
| loss = sum(jax.tree_util.tree_leaves(losses)) | |
| return loss, aux | |
| def _model(params, model_params, inputs, key): | |
| inputs["params"] = params | |
| opt = inputs["opt"] | |
| seq = soft_seq(params["seq"], inputs["bias"], opt) | |
| if "fix_pos" in opt: | |
| if "pos" in self.opt: | |
| seq_ref = jax.nn.one_hot(inputs["batch"]["aatype_sub"],20) | |
| p = opt["pos"][opt["fix_pos"]] | |
| fix_seq = lambda x:x.at[...,p,:].set(seq_ref) | |
| else: | |
| seq_ref = jax.nn.one_hot(inputs["batch"]["aatype"],20) | |
| p = opt["fix_pos"] | |
| fix_seq = lambda x:x.at[...,p,:].set(seq_ref[...,p,:]) | |
| seq = jax.tree_util.tree_map(fix_seq, seq) | |
| inputs.update({"seq":seq["pseudo"][0], | |
| "prf":jnp.where(opt["use_pssm"],seq["pssm"],seq["pseudo"])[0]}) | |
| rate = jnp.where(opt["dropout"],0.15,0.0) | |
| outputs = runner(inputs, model_params, key, rate) | |
| loss, aux = _get_loss(inputs, outputs) | |
| aux.update({"seq":seq,"opt":opt}) | |
| return loss, aux | |
| return {"grad_fn":jax.jit(jax.value_and_grad(_model, has_aux=True, argnums=0)), | |
| "fn":jax.jit(_model)} | |
| def prep_inputs(self, pdb_filename=None, chain=None, length=None, | |
| pos=None, fix_pos=None, atoms_to_exclude=None, ignore_missing=True, | |
| **kwargs): | |
| ''' | |
| prep inputs for TrDesign | |
| ''' | |
| if self.protocol in ["fixbb", "partial"]: | |
| # parse PDB file and return features compatible with TrRosetta | |
| pdb = prep_pdb(pdb_filename, chain, ignore_missing=ignore_missing) | |
| self._inputs["batch"] = pdb["batch"] | |
| if fix_pos is not None: | |
| self.opt["fix_pos"] = prep_pos(fix_pos, **pdb["idx"])["pos"] | |
| if self.protocol == "partial" and pos is not None: | |
| self._pos_info = prep_pos(pos, **pdb["idx"]) | |
| p = self._pos_info["pos"] | |
| aatype = self._inputs["batch"]["aatype"] | |
| self._inputs["batch"] = jax.tree_util.tree_map(lambda x:x[p], self._inputs["batch"]) | |
| self.opt["pos"] = p | |
| if "fix_pos" in self.opt: | |
| sub_i,sub_p = [],[] | |
| p = p.tolist() | |
| for i in self.opt["fix_pos"].tolist(): | |
| if i in p: | |
| sub_i.append(i) | |
| sub_p.append(p.index(i)) | |
| self.opt["fix_pos"] = np.array(sub_p) | |
| self._inputs["batch"]["aatype_sub"] = aatype[sub_i] | |
| self._inputs["6D"] = _np_get_6D_binned(self._inputs["batch"]["all_atom_positions"], | |
| self._inputs["batch"]["all_atom_mask"]) | |
| self._len = len(self._inputs["batch"]["aatype"]) | |
| self.opt["weights"]["cce"] = {"dist":1/6,"omega":1/6,"theta":2/6,"phi":2/6} | |
| if atoms_to_exclude is not None: | |
| if "N" in atoms_to_exclude: | |
| # theta = [N]-CA-CB-CB | |
| self.opt["weights"]["cce"] = dict(dist=1/4,omega=1/4,phi=1/2,theta=0) | |
| if "CA" in atoms_to_exclude: | |
| # theta = N-[CA]-CB-CB | |
| # omega = [CA]-CB-CB-[CA] | |
| # phi = [CA]-CB-CB | |
| self.opt["weights"]["cce"] = dict(dist=1,omega=0,phi=0,theta=0) | |
| if self.protocol in ["hallucination", "partial"]: | |
| # compute background distribution | |
| if length is not None: self._len = length | |
| self._inputs["6D_bkg"] = [] | |
| key = jax.random.PRNGKey(0) | |
| for n in range(1,6): | |
| p = os.path.join(self._data_dir,os.path.join("bkgr_models",f"bkgr0{n}.npy")) | |
| self._inputs["6D_bkg"].append(self._bkg_model(get_model_params(p), key, self._len)) | |
| self._inputs["6D_bkg"] = jax.tree_util.tree_map(lambda *x:np.stack(x).mean(0), *self._inputs["6D_bkg"]) | |
| # reweight the background | |
| self.opt["weights"]["bkg"] = dict(dist=1/6,omega=1/6,phi=2/6,theta=2/6) | |
| self._opt = copy_dict(self.opt) | |
| self.restart(**kwargs) | |
| def set_opt(self, *args, **kwargs): | |
| ''' | |
| set [opt]ions | |
| ------------------- | |
| note: model.restart() resets the [opt]ions to their defaults | |
| use model.set_opt(..., set_defaults=True) | |
| or model.restart(..., reset_opt=False) to avoid this | |
| ------------------- | |
| model.set_opt(num_models=1) | |
| model.set_opt(con=dict(num=1)) or set_opt({"con":{"num":1}}) | |
| model.set_opt(lr=1, set_defaults=True) | |
| ''' | |
| if kwargs.pop("set_defaults", False): | |
| update_dict(self._opt, *args, **kwargs) | |
| update_dict(self.opt, *args, **kwargs) | |
| def restart(self, seed=None, opt=None, weights=None, | |
| seq=None, reset_opt=True, **kwargs): | |
| if reset_opt: | |
| self.opt = copy_dict(self._opt) | |
| self.set_opt(opt) | |
| self.set_weights(weights) | |
| self.set_seed(seed) | |
| # set sequence | |
| self.set_seq(seq, **kwargs) | |
| # setup optimizer | |
| self._k = 0 | |
| self.set_optimizer() | |
| # clear previous best | |
| self._tmp = {"best":{}} | |
| def run(self, backprop=True): | |
| '''run model to get outputs, losses and gradients''' | |
| # decide which model params to use | |
| ns = np.arange(5) | |
| m = min(self.opt["num_models"],len(ns)) | |
| if self.opt["sample_models"] and m != len(ns): | |
| model_num = np.random.choice(ns,(m,),replace=False) | |
| else: | |
| model_num = ns[:m] | |
| model_num = np.array(model_num).tolist() | |
| # run in serial | |
| aux_all = [] | |
| for n in model_num: | |
| model_params = self._model_params[n] | |
| self._inputs["opt"] = self.opt | |
| flags = [self._params, model_params, self._inputs, self.key()] | |
| if backprop: | |
| (loss,aux),grad = self._model["grad_fn"](*flags) | |
| else: | |
| loss,aux = self._model["fn"](*flags) | |
| grad = jax.tree_util.tree_map(np.zeros_like, self._params) | |
| aux.update({"loss":loss, "grad":grad}) | |
| aux_all.append(aux) | |
| # average results | |
| self.aux = jax.tree_util.tree_map(lambda *x:np.stack(x).mean(0), *aux_all) | |
| self.aux["model_num"] = model_num | |
| def step(self, backprop=True, callback=None, save_best=True, verbose=1): | |
| self.run(backprop=backprop) | |
| if callback is not None: callback(self) | |
| # modify gradients | |
| if self.opt["norm_seq_grad"]: self._norm_seq_grad() | |
| self._state, self.aux["grad"] = self._optimizer(self._state, self.aux["grad"], self._params) | |
| # apply gradients | |
| lr = self.opt["learning_rate"] | |
| self._params = jax.tree_util.tree_map(lambda x,g:x-lr*g, self._params, self.aux["grad"]) | |
| # increment | |
| self._k += 1 | |
| # save results | |
| if save_best: | |
| if "aux" not in self._tmp["best"] or self.aux["loss"] < self._tmp["best"]["aux"]["loss"]: | |
| self._tmp["best"]["aux"] = self.aux | |
| if verbose and (self._k % verbose) == 0: | |
| x = self.get_loss(get_best=False) | |
| x["models"] = self.aux["model_num"] | |
| print(dict_to_str(x, print_str=f"{self._k}", keys=["models"])) | |
| def predict(self, seq=None, models=0): | |
| self.set_opt(dropout=False) | |
| if seq is not None: | |
| self.set_seq(seq=seq, set_state=False) | |
| self.run(backprop=False) | |
| def design(self, iters=100, opt=None, weights=None, save_best=True, verbose=1): | |
| self.set_opt(opt) | |
| self.set_weights(weights) | |
| for _ in range(iters): | |
| self.step(save_best=save_best, verbose=verbose) | |
| def plot(self, mode="preds", dpi=100, get_best=True): | |
| '''plot predictions''' | |
| assert mode in ["preds","feats","bkg_feats"] | |
| if mode == "preds": | |
| aux = self._tmp["best"]["aux"] if (get_best and "aux" in self._tmp["best"]) else self.aux | |
| x = aux["outputs"] | |
| elif mode == "feats": | |
| x = self._inputs["6D"] | |
| elif mode == "bkg_feats": | |
| x = self._inputs["6D_bkg"] | |
| x = jax.tree_util.tree_map(np.asarray, x) | |
| plt.figure(figsize=(4*4,4), dpi=dpi) | |
| for n,k in enumerate(["theta","phi","dist","omega"]): | |
| v = x[k] | |
| plt.subplot(1,4,n+1) | |
| plt.title(k) | |
| plt.imshow(v.argmax(-1),cmap="binary") | |
| plt.show() | |
| def get_loss(self, k=None, get_best=True): | |
| aux = self._tmp["best"]["aux"] if (get_best and "aux" in self._tmp["best"]) else self.aux | |
| if k is None: | |
| return {k:self.get_loss(k, get_best=get_best) for k in aux["losses"].keys()} | |
| losses = aux["losses"][k] | |
| weights = aux["opt"]["weights"][k] | |
| weighted_losses = jax.tree_util.tree_map(lambda l,w:l*w, losses, weights) | |
| return float(sum(jax.tree_util.tree_leaves(weighted_losses))) | |
| def af_callback(self, weight=1.0, seed=None): | |
| def callback(af_model): | |
| # copy [opt]ions from afdesign | |
| for k,v in af_model.opt.items(): | |
| if k in self.opt and k not in ["weights"]: | |
| self.opt[k] = af_model.opt[k] | |
| # update sequence input | |
| self._params["seq"] = af_model._params["seq"] | |
| # run trdesign | |
| self.run(backprop = weight > 0) | |
| # add gradients | |
| af_model.aux["grad"]["seq"] += weight * self.aux["grad"]["seq"] | |
| # add loss | |
| af_model.aux["loss"] += weight * self.aux["loss"] | |
| # for verbose printout | |
| if self.protocol in ["hallucination","partial"]: | |
| af_model.aux["losses"]["TrD_bkg"] = self.get_loss("bkg", get_best=False) | |
| if self.protocol in ["fixbb","partial"]: | |
| af_model.aux["losses"]["TrD_cce"] = self.get_loss("cce", get_best=False) | |
| self.restart(seed=seed) | |
| return callback |