| 355 | |
| 356 | |
| 357 | class T5LayerFF(torch.nn.Module): |
| 358 | def __init__(self, model_dim, ff_dim, dtype, device): |
| 359 | super().__init__() |
| 360 | self.DenseReluDense = T5DenseGatedActDense(model_dim, ff_dim, dtype, device) |
| 361 | self.layer_norm = T5LayerNorm(model_dim, dtype=dtype, device=device) |
| 362 | |
| 363 | def forward(self, x): |
| 364 | forwarded_states = self.layer_norm(x) |
| 365 | forwarded_states = self.DenseReluDense(forwarded_states) |
| 366 | x += forwarded_states |
| 367 | return x |
| 368 | |
| 369 | |
| 370 | class T5Attention(torch.nn.Module): |