(
self,
num_heads,
hidden_channels,
activation,
attn_activation,
cutoff,
vecnorm_type,
trainable_vecnorm,
last_layer=False,
)
| 153 | |
| 154 | class ViS_MP(MessagePassing): |
| 155 | def __init__( |
| 156 | self, |
| 157 | num_heads, |
| 158 | hidden_channels, |
| 159 | activation, |
| 160 | attn_activation, |
| 161 | cutoff, |
| 162 | vecnorm_type, |
| 163 | trainable_vecnorm, |
| 164 | last_layer=False, |
| 165 | ): |
| 166 | super(ViS_MP, self).__init__(aggr="add", node_dim=0) |
| 167 | assert hidden_channels % num_heads == 0, ( |
| 168 | f"The number of hidden channels ({hidden_channels}) " |
| 169 | f"must be evenly divisible by the number of " |
| 170 | f"attention heads ({num_heads})" |
| 171 | ) |
| 172 | |
| 173 | self.num_heads = num_heads |
| 174 | self.hidden_channels = hidden_channels |
| 175 | self.head_dim = hidden_channels // num_heads |
| 176 | self.last_layer = last_layer |
| 177 | |
| 178 | self.layernorm = nn.LayerNorm(hidden_channels) |
| 179 | self.vec_layernorm = VecLayerNorm(hidden_channels, trainable=trainable_vecnorm, norm_type=vecnorm_type) |
| 180 | |
| 181 | self.act = act_class_mapping[activation]() |
| 182 | self.attn_activation = act_class_mapping[attn_activation]() |
| 183 | |
| 184 | self.cutoff = CosineCutoff(cutoff) |
| 185 | |
| 186 | self.vec_proj = nn.Linear(hidden_channels, hidden_channels * 3, bias=False) |
| 187 | |
| 188 | self.q_proj = nn.Linear(hidden_channels, hidden_channels) |
| 189 | self.k_proj = nn.Linear(hidden_channels, hidden_channels) |
| 190 | self.v_proj = nn.Linear(hidden_channels, hidden_channels) |
| 191 | self.dk_proj = nn.Linear(hidden_channels, hidden_channels) |
| 192 | self.dv_proj = nn.Linear(hidden_channels, hidden_channels) |
| 193 | |
| 194 | self.s_proj = nn.Linear(hidden_channels, hidden_channels * 2) |
| 195 | if not self.last_layer: |
| 196 | self.f_proj = nn.Linear(hidden_channels, hidden_channels) |
| 197 | self.w_src_proj = nn.Linear(hidden_channels, hidden_channels, bias=False) |
| 198 | self.w_trg_proj = nn.Linear(hidden_channels, hidden_channels, bias=False) |
| 199 | |
| 200 | self.o_proj = nn.Linear(hidden_channels, hidden_channels * 3) |
| 201 | |
| 202 | self.reset_parameters() |
| 203 | |
| 204 | @staticmethod |
| 205 | def vector_rejection(vec, d_ij): |
nothing calls this directly
no test coverage detected