flashinfer.norm.gemma_rmsnorm

flashinfer.norm.gemma_rmsnorm(input: Tensor, weight: Tensor, eps: float = 1e-06, out: Tensor | None = None, enable_pdl: bool | None = None) Tensor

Gemma 风格的均方根归一化。

out[i] = (input[i] / RMS(input)) * (weight[i] + 1)

参数:
  • input (torch.Tensor) – 输入张量,形状为 (batch_size, hidden_size)。

  • weight (torch.Tensor) – 权重张量,形状为 (hidden_size,)。

  • eps (float) – 用于数值稳定的 epsilon。

  • out (Optional[torch.Tensor]) – 输出张量,如果指定,则内核将就地更新此张量。

  • enable_pdl (bool) – 是否启用 程序依赖启动

返回值:

输出 – Gemma 归一化张量,形状为 (batch_size, hidden_size)。

返回值类型:

torch.Tensor