| 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'] |