| 39 | _option_members = {"task", "criterion", "metric"} |
| 40 | |
| 41 | def __init__(self, sigma_embedding: nn.Module, |
| 42 | model: nn.Module, |
| 43 | torsion_mlp_hidden_dims: list, |
| 44 | schedule_1pi_periodic: SO2VESchedule, |
| 45 | schedule_2pi_periodic: SO2VESchedule, |
| 46 | graph_construction_model: Optional[Any] = None, |
| 47 | verbose: int = 0, |
| 48 | train_chi_id: Optional[Any] = None, ): |
| 49 | super(TorsionalDiffusion, self).__init__() |
| 50 | self.torsion_mlp_hidden_dims = torsion_mlp_hidden_dims |
| 51 | self.model_list = nn.ModuleList([deepcopy(model) for _ in range(self.NUM_CHI_ANGLES)]) |
| 52 | self.sigma_embedding_list = nn.ModuleList([deepcopy(sigma_embedding) for _ in range(self.NUM_CHI_ANGLES)]) |
| 53 | self.torsion_mlp_list = nn.ModuleList([layers.MLP(self.model_list[i].output_dim, torsion_mlp_hidden_dims |
| 54 | + [4,]) for i in range(self.NUM_CHI_ANGLES)]) |
| 55 | self.schedule_2pi_periodic = schedule_2pi_periodic |
| 56 | self.schedule_1pi_periodic = schedule_1pi_periodic |
| 57 | self.graph_construction_model = graph_construction_model |
| 58 | self.verbose = verbose |
| 59 | self.train_chi_id = train_chi_id |
| 60 | |
| 61 | def forward(self, batch): |
| 62 | all_loss = torch.tensor(0, dtype=torch.float32, device=self.device) |