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

Class VecLayerNorm

tasks/AutoMolecule3D/code/visnet/models/utils.py:140–210  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

138
139
140class 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

Callers 2

__init__Method · 0.90
__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected