| 86 | |
| 87 | |
| 88 | class FP32LayerNorm(nn.LayerNorm): |
| 89 | def forward(self, inputs: torch.Tensor) -> torch.Tensor: |
| 90 | origin_dtype = inputs.dtype |
| 91 | return F.layer_norm( |
| 92 | inputs.float(), |
| 93 | self.normalized_shape, |
| 94 | self.weight.float() if self.weight is not None else None, |
| 95 | self.bias.float() if self.bias is not None else None, |
| 96 | self.eps, |
| 97 | ).to(origin_dtype) |
| 98 | |
| 99 | |
| 100 | class AdaLayerNormZero(nn.Module): |