File size: 3,149 Bytes
d766458 | 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 | from colabdesign.af.model import mk_af_model
from colabdesign.tr.model import mk_tr_model
class mk_af_tr_model:
def __init__(self, protocol="fixbb", use_templates=False,
recycle_mode="last", num_recycles=0):
assert protocol in ["fixbb","partial","hallucination","binder"]
self.af = mk_af_model(protocol=protocol, use_templates=use_templates,
recycle_mode=recycle_mode, num_recycles=num_recycles)
if protocol == "binder":
def _prep_inputs(pdb_filename, chain, binder_len=50, binder_chain=None,
ignore_missing=True, **kwargs):
self.af.prep_inputs(pdb_filename=pdb_filename, chain=chain,
binder_len=binder_len, binder_chain=binder_chain,
ignore_missing=ignore_missing, **kwargs)
flags = dict(ignore_missing=ignore_missing)
if binder_chain is None:
self.tr = mk_tr_model(protocol="hallucination")
self.tr.prep_inputs(length=binder_len, **flags)
else:
self.tr = mk_tr_model(protocol="fixbb")
self.tr.prep_inputs(pdb_filename=pdb_filename, chain=binder_chain, **flags)
else:
self.tr = mk_tr_model(protocol=protocol)
if protocol == "fixbb":
def _prep_inputs(pdb_filename, chain, fix_pos=None,
ignore_missing=True, **kwargs):
flags = dict(pdb_filename=pdb_filename, chain=chain,
fix_pos=fix_pos, ignore_missing=ignore_missing)
self.af.prep_inputs(**flags, **kwargs)
self.tr.prep_inputs(**flags, chain=chain)
if protocol == "partial":
def _prep_inputs(pdb_filename, chain, pos=None, length=None,
fix_pos=None, use_sidechains=False, atoms_to_exclude=None,
ignore_missing=True, **kwargs):
if use_sidechains: fix_seq = True
flags = dict(pdb_filename=pdb_filename, chain=chain,
length=length, pos=pos, fix_pos=fix_pos,
ignore_missing=ignore_missing)
af_a2e = kwargs.pop("af_atoms_to_exclude",atoms_to_exclude)
tr_a2e = kwargs.pop("tr_atoms_to_exclude",atoms_to_exclude)
self.af.prep_inputs(**flags, use_sidechains=use_sidechains, atoms_to_exclude=af_a2e, **kwargs)
self.tr.prep_inputs(**flags, atoms_to_exclude=tr_a2e)
def _rewire(order=None, offset=0, loops=0):
self.af.rewire(order=order, offset=offset, loops=loops)
self.tr.rewire(order=order, offset=offset, loops=loops)
self.rewire = _rewire
if protocol == "hallucintion":
def _prep_inputs(length=None, **kwargs):
self.af.prep_inputs(length=length, **kwargs)
self.tr.prep_inputs(length=length)
self.prep_inputs = _prep_inputs
def set_opt(self,*args,**kwargs):
self.af.set_opt(*args,**kwargs)
self.tr.set_opt(*args,**kwargs)
def joint_design(self, iters=100, tr_weight=1.0, tr_seed=None, **kwargs):
self.af.design(iters, callback=self.tr.af_callback(weight=tr_weight, seed=tr_seed), **kwargs)
|