| 222 | |
| 223 | @R.register("tasks.ConfidencePrediction") |
| 224 | class 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 |
nothing calls this directly
no outgoing calls
no test coverage detected