| 6 | from rfdiffusion.util_module import init_lecun_normal |
| 7 | |
| 8 | class FeedForwardLayer(nn.Module): |
| 9 | def __init__(self, d_model, r_ff, p_drop=0.1): |
| 10 | super(FeedForwardLayer, self).__init__() |
| 11 | self.norm = nn.LayerNorm(d_model) |
| 12 | self.linear1 = nn.Linear(d_model, d_model*r_ff) |
| 13 | self.dropout = nn.Dropout(p_drop) |
| 14 | self.linear2 = nn.Linear(d_model*r_ff, d_model) |
| 15 | |
| 16 | self.reset_parameter() |
| 17 | |
| 18 | def reset_parameter(self): |
| 19 | # initialize linear layer right before ReLu: He initializer (kaiming normal) |
| 20 | nn.init.kaiming_normal_(self.linear1.weight, nonlinearity='relu') |
| 21 | nn.init.zeros_(self.linear1.bias) |
| 22 | |
| 23 | # initialize linear layer right before residual connection: zero initialize |
| 24 | nn.init.zeros_(self.linear2.weight) |
| 25 | nn.init.zeros_(self.linear2.bias) |
| 26 | |
| 27 | def forward(self, src): |
| 28 | src = self.norm(src) |
| 29 | src = self.linear2(self.dropout(F.relu_(self.linear1(src)))) |
| 30 | return src |
| 31 | |
| 32 | class Attention(nn.Module): |
| 33 | # calculate multi-head attention |