(self, original: nn.Linear, kernel_fn: Callable)
| 409 | """Wraps nn.Linear to use an optimized matmul kernel_fn.""" |
| 410 | |
| 411 | def __init__(self, original: nn.Linear, kernel_fn: Callable): |
| 412 | super().__init__() |
| 413 | self.original = original |
| 414 | self.kernel_fn = kernel_fn |
| 415 | self.weight = original.weight |
| 416 | self.bias = original.bias |
| 417 | |
| 418 | def forward(self, x: torch.Tensor) -> torch.Tensor: |
| 419 | # Reshape to 2D for kernel_fn, then reshape back |