flashinfer.norm.rmsnorm¶
- flashinfer.norm.rmsnorm(input: Tensor, weight: Tensor, eps: float = 1e-06, out: Tensor | None = None, enable_pdl: bool | None = None) Tensor¶
均方根归一化。
out[i] = (input[i] / RMS(input)) * weight[i]- 参数:
input (torch.Tensor) – 输入张量,2D 形状 (batch_size, hidden_size) 或 3D 形状 (batch_size, num_heads, hidden_size)。
weight (torch.Tensor) – 权重张量,形状为 (hidden_size,)。
eps (float) – 用于数值稳定的 epsilon。
out (Optional[torch.Tensor]) – 输出张量,如果指定,则内核将就地更新此张量。
enable_pdl (bool) – 是否启用 程序依赖启动
- 返回值:
output – 归一化张量,2D 形状 (batch_size, hidden_size) 或 3D 形状 (batch_size, num_heads, hidden_size)。
- 返回值类型:
torch.Tensor