| 138 | |
| 139 | |
| 140 | class VecLayerNorm(nn.Module): |
| 141 | def __init__(self, hidden_channels, trainable, norm_type="max_min"): |
| 142 | super(VecLayerNorm, self).__init__() |
| 143 | |
| 144 | self.hidden_channels = hidden_channels |
| 145 | self.eps = 1e-12 |
| 146 | |
| 147 | weight = torch.ones(self.hidden_channels) |
| 148 | if trainable: |
| 149 | self.register_parameter("weight", nn.Parameter(weight)) |
| 150 | else: |
| 151 | self.register_buffer("weight", weight) |
| 152 | |
| 153 | if norm_type == "rms": |
| 154 | self.norm = self.rms_norm |
| 155 | elif norm_type == "max_min": |
| 156 | self.norm = self.max_min_norm |
| 157 | else: |
| 158 | self.norm = self.none_norm |
| 159 | |
| 160 | self.reset_parameters() |
| 161 | |
| 162 | def reset_parameters(self): |
| 163 | weight = torch.ones(self.hidden_channels) |
| 164 | self.weight.data.copy_(weight) |
| 165 | |
| 166 | def none_norm(self, vec): |
| 167 | return vec |
| 168 | |
| 169 | def rms_norm(self, vec): |
| 170 | # vec: (num_atoms, 3 or 5, hidden_channels) |
| 171 | dist = torch.norm(vec, dim=1) |
| 172 | |
| 173 | if (dist == 0).all(): |
| 174 | return torch.zeros_like(vec) |
| 175 | |
| 176 | dist = dist.clamp(min=self.eps) |
| 177 | dist = torch.sqrt(torch.mean(dist ** 2, dim=-1)) |
| 178 | return vec / F.relu(dist).unsqueeze(-1).unsqueeze(-1) |
| 179 | |
| 180 | def max_min_norm(self, vec): |
| 181 | # vec: (num_atoms, 3 or 5, hidden_channels) |
| 182 | dist = torch.norm(vec, dim=1, keepdim=True) |
| 183 | |
| 184 | if (dist == 0).all(): |
| 185 | return torch.zeros_like(vec) |
| 186 | |
| 187 | dist = dist.clamp(min=self.eps) |
| 188 | direct = vec / dist |
| 189 | |
| 190 | max_val, _ = torch.max(dist, dim=-1) |
| 191 | min_val, _ = torch.min(dist, dim=-1) |
| 192 | delta = (max_val - min_val).view(-1) |
| 193 | delta = torch.where(delta == 0, torch.ones_like(delta), delta) |
| 194 | dist = (dist - min_val.view(-1, 1, 1)) / delta.view(-1, 1, 1) |
| 195 | |
| 196 | return F.relu(dist) * direct |
| 197 | |