(
self,
representation_model,
output_model,
prior_model=None,
reduce_op="add",
mean=None,
std=None,
derivative=False,
)
| 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() |
nothing calls this directly
no test coverage detected