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

Method __init__

diffpack/task.py:41–59  ·  view source on GitHub ↗
(self, sigma_embedding: nn.Module,
                 model: nn.Module,
                 torsion_mlp_hidden_dims: list,
                 schedule_1pi_periodic: SO2VESchedule,
                 schedule_2pi_periodic: SO2VESchedule,
                 graph_construction_model: Optional[Any] = None,
                 verbose: int = 0,
                 train_chi_id: Optional[Any] = None, )

Source from the content-addressed store, hash-verified

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)

Callers 1

__init__Method · 0.45

Calls

no outgoing calls

Tested by

no test coverage detected