MCPcopy Create free account
hub / github.com/RightNow-AI/autokernel / _LayerNormWrapper

Class _LayerNormWrapper

verify.py:441–475  ·  view source on GitHub ↗

Wraps nn.LayerNorm to use an optimized kernel_fn.

Source from the content-addressed store, hash-verified

439
440
441class _LayerNormWrapper(nn.Module):
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])
455 orig_shape = x.shape
456 if x.dim() > 2:
457 x_2d = x.reshape(-1, x.shape[-1])
458 else:
459 x_2d = x
460
461 try:
462 # Try full signature: kernel_fn(x, weight, bias, eps)
463 out = self.kernel_fn(x_2d, self.weight, self.bias, self.eps)
464 except TypeError:
465 try:
466 # Try without eps: kernel_fn(x, weight, bias)
467 out = self.kernel_fn(x_2d, self.weight, self.bias)
468 except TypeError:
469 # Fallback: just x
470 out = self.kernel_fn(x_2d)
471
472 if len(orig_shape) > 2:
473 out = out.reshape(orig_shape)
474
475 return out
476
477
478class _RMSNormWrapper(nn.Module):

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected