(self, original: nn.LayerNorm, kernel_fn: Callable)
| 442 | """Wraps nn.LayerNorm to use an optimized kernel_fn.""" |
| 443 | |
| 444 | def __init__(self, original: nn.LayerNorm, kernel_fn: Callable): |
| 445 | super().__init__() |
| 446 | self.original = original |
| 447 | self.kernel_fn = kernel_fn |
| 448 | self.weight = original.weight |
| 449 | self.bias = original.bias |
| 450 | self.eps = original.eps |
| 451 | self.normalized_shape = original.normalized_shape |
| 452 | |
| 453 | def forward(self, x: torch.Tensor) -> torch.Tensor: |
| 454 | # Reshape if needed: kernel_fn expects (x, weight, bias[, eps]) |