(
self,
lmax=2,
vecnorm_type='none',
trainable_vecnorm=False,
num_heads=8,
num_layers=9,
hidden_channels=256,
num_rbf=32,
rbf_type="expnorm",
trainable_rbf=False,
activation="silu",
attn_activation="silu",
max_z=100,
cutoff=5.0,
max_num_neighbors=32,
vertex_type="Edge",
)
| 50 | class ViSNetBlock(nn.Module): |
| 51 | |
| 52 | def __init__( |
| 53 | self, |
| 54 | lmax=2, |
| 55 | vecnorm_type='none', |
| 56 | trainable_vecnorm=False, |
| 57 | num_heads=8, |
| 58 | num_layers=9, |
| 59 | hidden_channels=256, |
| 60 | num_rbf=32, |
| 61 | rbf_type="expnorm", |
| 62 | trainable_rbf=False, |
| 63 | activation="silu", |
| 64 | attn_activation="silu", |
| 65 | max_z=100, |
| 66 | cutoff=5.0, |
| 67 | max_num_neighbors=32, |
| 68 | vertex_type="Edge", |
| 69 | ): |
| 70 | super(ViSNetBlock, self).__init__() |
| 71 | self.lmax = lmax |
| 72 | self.vecnorm_type = vecnorm_type |
| 73 | self.trainable_vecnorm = trainable_vecnorm |
| 74 | self.num_heads = num_heads |
| 75 | self.num_layers = num_layers |
| 76 | self.hidden_channels = hidden_channels |
| 77 | self.num_rbf = num_rbf |
| 78 | self.rbf_type = rbf_type |
| 79 | self.trainable_rbf = trainable_rbf |
| 80 | self.activation = activation |
| 81 | self.attn_activation = attn_activation |
| 82 | self.max_z = max_z |
| 83 | self.cutoff = cutoff |
| 84 | self.max_num_neighbors = max_num_neighbors |
| 85 | |
| 86 | self.embedding = nn.Embedding(max_z, hidden_channels) |
| 87 | self.distance = Distance(cutoff, max_num_neighbors=max_num_neighbors, loop=True) |
| 88 | self.sphere = Sphere(l=lmax) |
| 89 | self.distance_expansion = rbf_class_mapping[rbf_type](cutoff, num_rbf, trainable_rbf) |
| 90 | self.neighbor_embedding = NeighborEmbedding(hidden_channels, num_rbf, cutoff, max_z).jittable() |
| 91 | self.edge_embedding = EdgeEmbedding(num_rbf, hidden_channels).jittable() |
| 92 | |
| 93 | self.vis_mp_layers = nn.ModuleList() |
| 94 | vis_mp_kwargs = dict( |
| 95 | num_heads=num_heads, |
| 96 | hidden_channels=hidden_channels, |
| 97 | activation=activation, |
| 98 | attn_activation=attn_activation, |
| 99 | cutoff=cutoff, |
| 100 | vecnorm_type=vecnorm_type, |
| 101 | trainable_vecnorm=trainable_vecnorm |
| 102 | ) |
| 103 | vis_mp_class = VIS_MP_MAP.get(vertex_type, ViS_MP) |
| 104 | for _ in range(num_layers - 1): |
| 105 | layer = vis_mp_class(last_layer=False, **vis_mp_kwargs).jittable() |
| 106 | self.vis_mp_layers.append(layer) |
| 107 | self.vis_mp_layers.append(vis_mp_class(last_layer=True, **vis_mp_kwargs).jittable()) |
| 108 | |
| 109 | self.out_norm = nn.LayerNorm(hidden_channels) |
nothing calls this directly
no test coverage detected