MCPcopy Create free account
hub / github.com/DeepGraphLearning/DiffPack / ConfidencePrediction

Class ConfidencePrediction

diffpack/task.py:224–283  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

222
223@R.register("tasks.ConfidencePrediction")
224class ConfidencePrediction(TorsionalDiffusion):
225 eps = 1e-10
226 _option_members = {"task", "criterion", "metric"}
227
228 def __init__(self, sigma_embedding: nn.Module,
229 model: nn.Module,
230 confidence_model: nn.Module,
231 torsion_mlp_hidden_dims: list,
232 schedule_1pi_periodic: SO2VESchedule,
233 schedule_2pi_periodic: SO2VESchedule,
234 num_sample: int = 5,
235 num_mlp_layer: int = 1,
236 graph_construction_model: Optional[Any] = None,
237 verbose: int = 0,
238 train_chi_id: Optional[Any] = None):
239 super().__init__(sigma_embedding,
240 model,
241 torsion_mlp_hidden_dims,
242 schedule_1pi_periodic,
243 schedule_2pi_periodic,
244 graph_construction_model,
245 verbose,
246 train_chi_id)
247 self.confidence_model = confidence_model
248 self.num_sample = num_sample
249 self.mlp = layers.MLP(self.confidence_model.output_dim,
250 [self.confidence_model.output_dim] * num_mlp_layer + [1])
251
252 def predict_rmsd(self, batch, all_loss=None, metric=None):
253 protein = batch['graph']
254 if self.graph_construction_model:
255 protein = self.graph_construction_model(protein)
256 atom_feature = self.confidence_model(protein, protein.node_feature.float())["node_feature"]
257 residue_feature = scatter_mean(atom_feature, protein.atom2residue, dim=0,
258 dim_size=protein.num_residue) # [num_residue, feature_dim]
259 pred = self.mlp(residue_feature).squeeze(-1) # [num_residue]
260 return pred
261
262 @torch.no_grad()
263 def generate(self, batch, randomize=True):
264 protein = batch['graph']
265 if randomize:
266 protein = rotamer.randomize(protein)
267
268 best_protein = protein.clone()
269 best_rmsd = torch.zeros(protein.num_residue, device=self.device) + 1e6
270 for _ in tqdm(range(self.num_sample), desc="Confidence sampling"):
271 batch = super().generate(batch, randomize=True) # TODO: do we need to randomize?
272 protein = batch['graph']
273 rmsd = self.predict_rmsd(batch)
274 residue_update_mask = rmsd < best_rmsd # [num_residue]
275 atom_update_mask = residue_update_mask[protein.atom2residue] # [num_atom]
276 best_protein.node_position[atom_update_mask] = protein.node_position[atom_update_mask]
277 best_rmsd[residue_update_mask] = rmsd[residue_update_mask]
278
279 best_batch = {
280 "graph": best_protein,
281 "rmsd": best_rmsd

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected