(
self,
x: torch.Tensor,
)
| 358 | self.ffn_norm.reset_parameters() |
| 359 | |
| 360 | def forward( |
| 361 | self, |
| 362 | x: torch.Tensor, |
| 363 | ): |
| 364 | |
| 365 | x_attn = x + self.ls1(self.attention(self.attention_norm(x), self.is_causal)) |
| 366 | x_ffn = x_attn + self.ls2(self.feed_forward(self.ffn_norm(x_attn))) |
| 367 | return x_ffn |
| 368 | |
| 369 | |
| 370 | class ResidualAttentionBlock(nn.Module): |