| 543 | |
| 544 | |
| 545 | class ViSNet(nn.Module): |
| 546 | def __init__( |
| 547 | self, |
| 548 | representation_model, |
| 549 | output_model, |
| 550 | prior_model=None, |
| 551 | reduce_op="add", |
| 552 | mean=None, |
| 553 | std=None, |
| 554 | derivative=False, |
| 555 | ): |
| 556 | super(ViSNet, self).__init__() |
| 557 | self.representation_model = representation_model |
| 558 | self.output_model = output_model |
| 559 | |
| 560 | self.prior_model = prior_model |
| 561 | if not output_model.allow_prior_model and prior_model is not None: |
| 562 | self.prior_model = None |
| 563 | rank_zero_warn( |
| 564 | "Prior model was given but the output model does " |
| 565 | "not allow prior models. Dropping the prior model." |
| 566 | ) |
| 567 | |
| 568 | self.reduce_op = reduce_op |
| 569 | self.derivative = derivative |
| 570 | |
| 571 | mean = torch.scalar_tensor(0) if mean is None else mean |
| 572 | self.register_buffer("mean", mean) |
| 573 | std = torch.scalar_tensor(1) if std is None else std |
| 574 | self.register_buffer("std", std) |
| 575 | |
| 576 | self.reset_parameters() |
| 577 | |
| 578 | def reset_parameters(self): |
| 579 | self.representation_model.reset_parameters() |
| 580 | self.output_model.reset_parameters() |
| 581 | if self.prior_model is not None: |
| 582 | self.prior_model.reset_parameters() |
| 583 | |
| 584 | def forward(self, data: Data) -> Tuple[Tensor, Optional[Tensor]]: |
| 585 | |
| 586 | if self.derivative: |
| 587 | data.pos.requires_grad_(True) |
| 588 | |
| 589 | x, v = self.representation_model(data) |
| 590 | x = self.output_model.pre_reduce(x, v, data.z, data.pos, data.batch) |
| 591 | x = x * self.std |
| 592 | |
| 593 | if self.prior_model is not None: |
| 594 | x = self.prior_model(x, data.z) |
| 595 | |
| 596 | out = scatter(x, data.batch, dim=0, reduce=self.reduce_op) |
| 597 | out = self.output_model.post_reduce(out) |
| 598 | |
| 599 | out = out + self.mean |
| 600 | |
| 601 | # compute gradients with respect to coordinates |
| 602 | if self.derivative: |