forked from gelnesr/dEVA
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathsampler.py
More file actions
48 lines (41 loc) · 1.58 KB
/
Copy pathsampler.py
File metadata and controls
48 lines (41 loc) · 1.58 KB
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
import gc
import os
import time
import torch
import subprocess
import numpy as np
from evolve.individual import Individual
class Sampler(object):
def __init__(self, models):
super(Sampler, self).__init__()
self.rem_models = models
self.seq_model = self.rem_models.pop('seq_model')
self.fixed_residues = None
if hasattr(self.seq_model, 'init_seq') and callable(getattr(self.seq_model, 'init_seq')):
pass
else:
raise ValueError("Sequence design model does not have init_seq function")
pass
def init_seq(self, individual: Individual):
self.get_fixed_residues()
self.seq_model.init_seq(individual)
for k, m in self.rem_models.items():
m.score(individual)
self._maybe_rescore_pmpnn(k, m, individual)
def step(self, individual, num_mutations=1):
self.seq_model.score(individual, num_mutations=num_mutations)
for k, m in self.rem_models.items():
m.score(individual)
self._maybe_rescore_pmpnn(k, m, individual)
def _maybe_rescore_pmpnn(self, model_key, model, individual):
"""If relax rewrote the backbone, refresh pmpnn on that PDB."""
if model_key != 'relax':
return
if not getattr(model, 'rescore_pmpnn_after', False):
return
if not hasattr(self.seq_model, 'rescore'):
return
self.seq_model.rescore(individual)
def get_fixed_residues(self):
self.fixed_residues = self.seq_model.fixed_resis()
return self.fixed_residues