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