lite_boost.layers.rms_norm
- lite_boost.layers.rms_norm(x, gamma, eps=1e-6)[源代码]
对 x 的最后一维执行逐行RMS归一化。
计算 \(y = x / RMS(x) * gamma\),其中 \(RMS(x) = \sqrt{mean(x^2)}\),\(mean(x^2)\) 是 x 在最后一维上的平方均值。通常用于替代Wan系列VAE
RMS_norm层中展开的F.normalize(x, dim) * sqrt(dim) * gamma计算链。- 参数:
x (Tensor) - shape为 \((N, C)\) 的二维输入张量,归一化维度是最后一维。支持float16、float32,在A2等平台上额外支持bfloat16。
gamma (Tensor) - 逐列缩放因子,shape为 \((C,)\),dtype与 x 保持一致,与 x 的最后一维匹配。
eps (float, 可选) - 为数值稳定性添加到分母的值。默认值:
1e-6。
- 返回:
Tensor, 与 x 相同shape和dtype的归一化结果。
- 异常:
ValueError - 输入不是二维(仅支持 (N, C) 输入),或 gamma 不是一维且与 x 最后一维不匹配时抛出。
ValueError - 在300I Duo上 x 为bfloat16,或 x 的最后一维小于16时抛出。
- 支持平台:
Ascend
样例:
>>> import torch >>> import torch_npu >>> from lite_boost.layers import rms_norm >>> torch.npu.set_device(0) >>> x = torch.full((2, 16), 2.0, device="npu") >>> gamma = torch.ones(16, device="npu") >>> y = rms_norm(x, gamma) >>> print(y.shape) torch.Size([2, 16])