@CustomOp.register("nemotron_layer_norm")
class NemotronLayerNorm1P(CustomOp):
"""LayerNorm variant used by Nemotron: `x * (1 + w)` instead of `x * w`."""
def __init__(self, hidden_size: int, eps: float = 1e-5) -> None:
super().__init__()
self.normalized_shape = (hidden_size,)
self.eps = eps
self.weight = nn.Parameter(torch.zeros(hidden_size))
self.bias = nn.Parameter(torch.zeros(hidden_size))
def forward_native(
self,
x: torch.Tensor,
residual: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
if residual is not None:
x = x + residual
residual = x
device_type = x.device.type
args = _cast_if_autocast_enabled(
device_type, x, self.normalized_shape, self.weight + 1, self.bias, self.eps
)
with torch.amp.autocast(device_type, enabled=False):
x = torch.nn.functional.layer_norm(*args)
return x if residual is None else (x, residual)
def forward_cuda(
self,
x: torch.Tensor,
residual: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
return self.forward_native(x, residual)
def forward_xpu(
self,
x: torch.Tensor,
residual: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
import vllm._xpu_ops # noqa: F401 registers torch.ops.vllm.xpu_nemotron_layer_norm
if not hasattr(torch.ops._C, "nemotron_layer_norm"):
return self.forward_native(x, residual)
if residual is not None:
if not hasattr(torch.ops._C, "fused_add_nemotron_layer_norm"):
return self.forward_native(x, residual)
torch.ops.vllm.xpu_fused_add_nemotron_layer_norm(
x, residual, self.weight, self.bias, self.eps
)
return x, residual
# empty_like preserves x's strides, but the kernel requires a
# contiguous out (unlike x, which it can handle non-contiguous).
out = torch.empty(x.shape, device=x.device, dtype=x.dtype)
torch.ops.vllm.xpu_nemotron_layer_norm(out, x, self.weight, self.bias, self.eps)
return out