MCPcopy Create free account
hub / github.com/InternScience/InternAgent / ViSNet

Class ViSNet

tasks/AutoMolecule3D/code/experiment.py:545–614  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

543
544
545class 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:

Callers 1

create_modelFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected