已开启
feat:add _layer_norm_fwd_1pass_kernel #110
gcw_K5CvmS79创建于 7月23日
feat:add _layer_norm_fwd_1pass_kernel #110
已开启
共 7 个文件变更+886-0
| @@ -0,0 +1,265 @@ | |||
| 1 | +# _layer_norm_fwd_1pass_kernel 算子 | ||
| 2 | + | ||
| 3 | +## 概述 | ||
| 4 | + | ||
| 5 | +`_layer_norm_fwd_1pass_kernel` 是一个基于 Triton 实现的高效 LayerNorm/RMSNorm 前向算子,源自 state-spaces/mamba 开源仓库。该算子支持标准 LayerNorm 和 RMSNorm 两种模式,并支持门控机制(SiLU gating via Z branch),适用于 Mamba 状态空间模型中的归一化计算。 | ||
| 6 | + | ||
| 7 | +当前实现面向 Ascend arch32(910C),通过 1D 物理核网格 + 任务分发循环 + 多 Token 批量处理,充分利用昇腾 NPU 的 Vector Core 并行计算能力。 | ||
| 8 | + | ||
| 9 | +--- | ||
| 10 | + | ||
| 11 | +## 函数签名 | ||
| 12 | + | ||
| 13 | +```python | ||
| 14 | +from mindspeed_ops.api.triton.layer_norm_fwd_1pass import layer_norm_fwd_1pass | ||
| 15 | + | ||
| 16 | +def layer_norm_fwd_1pass( | ||
| 17 | + x: torch.Tensor, | ||
| 18 | + weight: torch.Tensor, | ||
| 19 | + bias: torch.Tensor = None, | ||
| 20 | + z: torch.Tensor = None, | ||
| 21 | + eps: float = 1e-5, | ||
| 22 | + norm_before_gate: bool = True, | ||
| 23 | + is_rms_norm: bool = False, | ||
| 24 | + ngroups: int = 1, | ||
| 25 | +) -> tuple: | ||
| 26 | + """LayerNorm/RMSNorm forward with optional gating.""" | ||
| 27 | +``` | ||
| 28 | + | ||
| 29 | +返回顺序固定为 `(y, mean, rstd)`。 | ||
| 30 | + | ||
| 31 | +--- | ||
| 32 | + | ||
| 33 | +## 参数说明 | ||
| 34 | + | ||
| 35 | +### 输入 | ||
| 36 | + | ||
| 37 | +| 参数 | 形状 | dtype | 描述 | | ||
| 38 | +|------|------|-------|------| | ||
| 39 | +| `x` | `[M, N*ngroups]` | `float16` / `float32` / `bfloat16` | 输入张量 | | ||
| 40 | +| `weight` | `[N*ngroups]`(`ngroups=1` 时即 `[N]`) | 同输入 | 归一化权重,按组切分为每组 `[N]` | | ||
| 41 | +| `bias` | `[N*ngroups]` 或 None(`ngroups=1` 时即 `[N]`) | 同输入 | 归一化偏置,按组切分(可选) | | ||
| 42 | +| `z` | `[M, N*ngroups]` 或 None | 同输入 | 门控分支张量(可选) | | ||
| 43 | +| `eps` | scalar | `float` | 防止除零的 epsilon | | ||
| 44 | +| `norm_before_gate` | scalar | `bool` | True: 先归一化再门控; False: 先门控再归一化 | | ||
| 45 | +| `is_rms_norm` | scalar | `bool` | True: RMSNorm; False: LayerNorm | | ||
| 46 | +| `ngroups` | scalar | `int` | 分组数 | | ||
| 47 | + | ||
| 48 | +### 输出 | ||
| 49 | + | ||
| 50 | +| 输出 | 形状 | dtype | 描述 | | ||
| 51 | +|------|------|-------|------| | ||
| 52 | +| `y` | 同 x | 同输入 | 归一化输出 | | ||
| 53 | +| `mean` | `[ngroups, M]` 或 None | `float32` | 均值(RMSNorm 时为 None) | | ||
| 54 | +| `rstd` | `[ngroups, M]` | `float32` | 逆标准差 | | ||
| 55 | + | ||
| 56 | +--- | ||
| 57 | + | ||
| 58 | +## 支持的数据类型 | ||
| 59 | + | ||
| 60 | +| 参数 | 支持的 dtype | 说明 | | ||
| 61 | +|------|--------------|------| | ||
| 62 | +| `x` / `weight` / `bias` / `z` | `float16` / `float32` / `bfloat16` | 内部以 float32 计算 | | ||
| 63 | + | ||
| 64 | +--- | ||
| 65 | + | ||
| 66 | +## 实现原理 | ||
| 67 | + | ||
| 68 | +### 核心数学公式 | ||
| 69 | + | ||
| 70 | +**LayerNorm:** | ||
| 71 | + | ||
| 72 | +```text | ||
| 73 | +mean = sum(x) / N | ||
| 74 | +var = sum((x - mean)^2) / N | ||
| 75 | +rstd = 1 / sqrt(var + eps) | ||
| 76 | +x_hat = (x - mean) * rstd | ||
| 77 | +y = x_hat * weight [+ bias] # bias 可选 | ||
| 78 | +``` | ||
| 79 | + | ||
| 80 | +**RMSNorm:** | ||
| 81 | + | ||
| 82 | +```text | ||
| 83 | +var = sum(x^2) / N | ||
| 84 | +rstd = 1 / sqrt(var + eps) | ||
| 85 | +x_hat = x * rstd | ||
| 86 | +y = x_hat * weight [+ bias] # bias 可选 | ||
| 87 | +``` | ||
| 88 | + | ||
| 89 | +**门控 (SiLU Gating):** | ||
| 90 | + | ||
| 91 | +```text | ||
| 92 | +gate(z) = z * sigmoid(z) | ||
| 93 | +# NORM_BEFORE_GATE=True: y = norm(x) * gate(z) | ||
| 94 | +# NORM_BEFORE_GATE=False: y = norm(x * gate(z)) | ||
| 95 | +``` | ||
| 96 | + | ||
| 97 | +### NPU 优化策略 | ||
| 98 | + | ||
| 99 | +1. **1D 物理核网格 + 任务分发**:将原始 2D Grid `(M, ngroups)` 转换为 `(num_core,)` 1D 网格,核内用 `tl.range(core_id, task_num, num_core)` 分发任务 | ||
| 100 | +2. **多 Token 批量处理 (BT)**:每个 task 处理 BT 行,减少任务粒度与调度次数 | ||
| 101 | +3. **2D 批量归约(核心优化)**:BT 行组织成 `[BT, BLOCK_N]` 二维 tile,一次 `tl.load` + `tl.sum(axis=1)` 完成归约,消除逐行标量循环 | ||
| 102 | +4. **UB 容量规划**:192KB 总容量取 50% 作 Double Buffering(约 85KB),按每 token `BLOCK_N*4*3`(含门控再加一份)动态计算最大 BT(`_compute_bt`) | ||
| 103 | +5. **索引/偏移使用 int32**:`cols`、`rows`、`task_id` 及指针偏移统一为 int32,避免 int64 scalar 退化 | ||
| 104 | +6. **权重/偏置广播复用**:weight/bias 每个 task 只加载一次,并在 2D 计算中按行广播 | ||
| 105 | +7. **自适应 warp/stages 配置**:根据 N 维度动态调整 `num_warps` 和 `num_stages` | ||
| 106 | + | ||
| 107 | +### Grid 与 Tiling | ||
| 108 | + | ||
| 109 | +| 参数 | 值 | 说明 | | ||
| 110 | +|------|------|------| | ||
| 111 | +| `grid` | `(num_core,)` | 1D 物理核网格(运行时通过 `get_vector_num()` 获取核数) | | ||
| 112 | +| `task_num` | `ceil(M/BT) * ngroups` | 总任务数 | | ||
| 113 | +| `BT` | 动态计算 | 每个 task 处理的行数 | | ||
| 114 | +| `BLOCK_N` | `next_power_of_2(N)` | 列方向 block 大小 | | ||
| 115 | + | ||
| 116 | +--- | ||
| 117 | + | ||
| 118 | +## 使用示例 | ||
| 119 | + | ||
| 120 | +```python | ||
| 121 | +import torch | ||
| 122 | +from mindspeed_ops.api.triton.layer_norm_fwd_1pass import layer_norm_fwd_1pass | ||
| 123 | + | ||
| 124 | +M, N = 32, 1024 | ||
| 125 | +dtype = torch.float16 | ||
| 126 | +device = "npu" | ||
| 127 | + | ||
| 128 | +x = torch.randn(M, N, device=device, dtype=dtype) | ||
| 129 | +weight = torch.randn(N, device=device, dtype=dtype) | ||
| 130 | +bias = torch.randn(N, device=device, dtype=dtype) | ||
| 131 | +z = torch.randn(M, N, device=device, dtype=dtype) | ||
| 132 | + | ||
| 133 | +# LayerNorm with gating | ||
| 134 | +y, mean, rstd = layer_norm_fwd_1pass(x, weight, bias, z, eps=1e-5, | ||
| 135 | + norm_before_gate=True, is_rms_norm=False) | ||
| 136 | + | ||
| 137 | +# RMSNorm without gating | ||
| 138 | +y, mean, rstd = layer_norm_fwd_1pass(x, weight, bias, None, eps=1e-5, | ||
| 139 | + norm_before_gate=True, is_rms_norm=True) | ||
| 140 | + | ||
| 141 | +print(f"y shape: {y.shape}") | ||
| 142 | +print(f"rstd shape: {rstd.shape}") | ||
| 143 | +``` | ||
| 144 | + | ||
| 145 | +--- | ||
| 146 | + | ||
| 147 | +## 测试结果 | ||
| 148 | + | ||
| 149 | +### UT 测试 | ||
| 150 | + | ||
| 151 | +测试文件:`tests/unit_tests/triton/test_layer_norm_fwd_1pass.py` | ||
| 152 | + | ||
| 153 | +- 测试用例数:147(6 shapes × 3 dtypes × 8 configs + 3 ngroups) | ||
| 154 | +- 结果:**147 passed**,耗时 9.23s | ||
| 155 | +- 覆盖配置:`is_rms_norm` × `has_z` × `norm_before_gate` 全组合,额外 `ngroups=2` 测试 | ||
| 156 | + | ||
| 157 | +#### UT 测试 Shape | ||
| 158 | + | ||
| 159 | +| M | N | dtype | is_rms_norm | has_z | norm_before_gate | | ||
| 160 | +|---|---|-------|-------------|-------|-----------------| | ||
| 161 | +| 2 | 64 | fp32/bf16/fp16 | True/False | True/False | True/False | | ||
| 162 | +| 4 | 128 | fp32/bf16/fp16 | True/False | True/False | True/False | | ||
| 163 | +| 8 | 256 | fp32/bf16/fp16 | True/False | True/False | True/False | | ||
| 164 | +| 16 | 512 | fp32/bf16/fp16 | True/False | True/False | True/False | | ||
| 165 | +| 32 | 1024 | fp32/bf16/fp16 | True/False | True/False | True/False | | ||
| 166 | +| 64 | 2048 | fp32/bf16/fp16 | True/False | True/False | True/False | | ||
| 167 | + | ||
| 168 | +### ATK 精度测试 | ||
| 169 | + | ||
| 170 | +测试文件:`tests/atk_tests/triton/layer_norm_fwd_1pass/` | ||
| 171 | + | ||
| 172 | +- 测试用例数:23(多 shapes × 3 dtypes 混合) | ||
| 173 | +- 精度基准:`single_bm`(Triton vs CPU golden) | ||
| 174 | +- 结果:**23/23 passed**,全部 SUCCESS | ||
| 175 | +- 最大绝对误差:fp32 ≤ 3.6e-7,fp16 ≤ 2.4e-4,bf16 ≤ 2.0e-3 | ||
| 176 | + | ||
| 177 | +### ATK 性能测试(Device 性能,Triton vs torch_npu,Ascend 910C) | ||
| 178 | + | ||
| 179 | +| Shape (M, N) | dtype | Triton (us) | torch_npu (us) | 加速比 | | ||
| 180 | +|--------------|-------|-------------|----------------|--------| | ||
| 181 | +| (2, 64) | fp16 | 4.44 | 20.41 | 4.60x | | ||
| 182 | +| (4, 128) | bf16 | 4.45 | 23.90 | 5.37x | | ||
| 183 | +| (8, 256) | fp16 | 4.30 | 35.57 | 8.28x | | ||
| 184 | +| (16, 256) | fp32 | 4.10 | 39.21 | 9.57x | | ||
| 185 | +| (16, 512) | bf16 | 4.20 | 49.61 | 11.81x | | ||
| 186 | +| (32, 128) | bf16 | 4.23 | 53.35 | 12.60x | | ||
| 187 | +| (32, 256) | fp32 | 4.56 | 56.12 | 12.31x | | ||
| 188 | +| (32, 1024) | fp32 | 5.69 | 55.12 | 9.68x | | ||
| 189 | +| (32, 2048) | fp32 | 5.83 | 64.89 | 11.12x | | ||
| 190 | +| (64, 64) | bf16 | 4.08 | 51.11 | 12.51x | | ||
| 191 | +| (64, 128) | bf16 | 4.48 | 54.83 | 12.24x | | ||
| 192 | +| (64, 256) | fp16 | 4.89 | 85.86 | 17.55x | | ||
| 193 | +| (64, 512) | fp16 | 5.57 | 96.62 | 17.36x | | ||
| 194 | +| (64, 2048) | fp16 | 5.88 | 107.26 | 18.23x | | ||
| 195 | + | ||
| 196 | +**ATK 性能汇总**: | ||
| 197 | + | ||
| 198 | +- 全部 23 用例 Device 性能测试通过 | ||
| 199 | +- 平均加速比:约 10.7x(Triton vs torch_npu) | ||
| 200 | +- 最高加速比:18.23x(M=64, N=2048, fp16) | ||
| 201 | +- 性能提升随 M×N 增大而显著提升 | ||
| 202 | + | ||
| 203 | +--- | ||
| 204 | + | ||
| 205 | +## 优化前后性能对比(msprof kernel 级测量) | ||
| 206 | + | ||
| 207 | +| 方法 | 说明 | | ||
| 208 | +|------|------| | ||
| 209 | +| 原始 Triton 实现(2D Grid) | 基线:直接使用 `grid=(M,)` 每行一个 block | | ||
| 210 | +| Triton NPU 优化实现 | 1D 物理核网格 + 多 Token 并行 + UB 容量规划 | | ||
| 211 | + | ||
| 212 | +### 性能测试结果(msprof op 测量,fp16,Ascend 910C) | ||
| 213 | + | ||
| 214 | +| Shape (M, N) | 基线耗时 (us) | 优化后耗时 (us) | 加速比 | | ||
| 215 | +|--------------|-------------|---------------|--------| | ||
| 216 | +| (256, 1024) | 18.56 | 7.44 | 2.50x | | ||
| 217 | +| (1024, 2048) | 77.42 | 24.40 | 3.17x | | ||
| 218 | +| (2048, 1024) | 156.18 | 26.82 | 5.82x | | ||
| 219 | +| (4096, 1024) | 310.10 | 44.84 | 6.92x | | ||
| 220 | +| (4096, 2048) | 303.62 | 95.18 | 3.19x | | ||
| 221 | +| (8192, 1024) | 618.84 | 85.14 | 7.27x | | ||
| 222 | +| (16384, 1024) | 1204.58 | 161.48 | 7.46x | | ||
| 223 | + | ||
| 224 | +**汇总统计(对比原 Triton 算子,满足任务书"提升 20%"要求)**: | ||
| 225 | + | ||
| 226 | +- 整体平均加速比:约 5.19x(即平均提升 419%,远超 20% 门槛) | ||
| 227 | +- 达到 20% 以上性能提升的测试用例:**100%(7/7)** | ||
| 228 | +- 最小加速比:2.50x(M=256, N=1024),最大加速比:7.46x(M=16384, N=1024) | ||
| 229 | +- 性能提升随 M 增大而显著提升,大 batch 场景下可达 7.46x | ||
| 230 | +- 小 batch(256×1024)由逐行标量循环改为 2D 批量归约后,从 1.18x 提升至 2.50x | ||
| 231 | + | ||
| 232 | +### 测试方法 | ||
| 233 | + | ||
| 234 | +使用 `msprof op` 工具测量纯 kernel 执行时间(Task Duration),排除 host 侧 Python 开销: | ||
| 235 | + | ||
| 236 | +```bash | ||
| 237 | +msprof op --output=./prof_output --kernel-name="_layer_norm_fwd_1pass_kernel" \ | ||
| 238 | + --warm-up=5 --launch-count=5 --kill=on python test_perf_optimized.py <M> <N> | ||
| 239 | +``` | ||
| 240 | + | ||
| 241 | +### 优化策略 | ||
| 242 | + | ||
| 243 | +1. **1D 物理核网格 + 任务分发**:以 `grid=(num_core,)` 固定物理核数,核内用 `tl.range(core_id, task_num, num_core)` 分发任务,消除 2D 逻辑网格调度开销,充分利用全部 Vector Core | ||
| 244 | +2. **多 Token 批量处理 (BT)**:每个 task 负责 BT 行,减少任务粒度与调度次数 | ||
| 245 | +3. **UB 容量规划**:根据 BLOCK_N 动态计算最大 BT(`_compute_bt`),确保 `[BT, BLOCK_N]` tile 不溢出 UB 并预留 Double Buffering | ||
| 246 | +4. **2D 批量归约(核心优化)**:将每个 task 的 BT 行组织成 `[BT, BLOCK_N]` 二维 tile,一次 `tl.load` + 一次 `tl.sum(axis=1)` 完成归约,彻底消除逐行标量循环,大幅提升 Vector 单元利用率 | ||
| 247 | +5. **索引/偏移使用 int32**:`cols`、`rows`、`task_id` 及指针偏移统一为 int32,避免 int64 导致的 scalar 退化 | ||
| 248 | +6. **权重偏置广播复用**:BT 行共享一次 weight/bias 加载,并在 2D 计算中按行广播,减少全局内存访问 | ||
| 249 | +7. **边计算边写入(多写入流)**:Mean/Rstd 提前写出,不阻塞后续归一化计算 | ||
| 250 | +8. **门控与归一化融合**:sigmoid gating 与归一化在同一 tile 内完成,隐藏访存延迟 | ||
| 251 | +9. **自适应 warp/stages 配置**:根据 N 维度动态选择 `num_warps`/`num_stages` | ||
| 252 | + | ||
| 253 | +--- | ||
| 254 | + | ||
| 255 | +## 注意事项 | ||
| 256 | + | ||
| 257 | +1. 该算子在 arch35 平台上暂不支持,会抛出 `NotImplementedError` | ||
| 258 | +2. 输入张量建议保持 contiguous,算子内部会强制调用 `contiguous()` | ||
| 259 | +3. 内部计算使用 float32 精度,输出保持输入 dtype | ||
| 260 | +4. `BLOCK_N` 自动设为 `next_power_of_2(N)`,N 过大时可能触发 UB 溢出 | ||
| 261 | +5. 不同 dtype 的精度阈值: | ||
| 262 | + - `float32`: `ratio=1e-4, atol=1e-5` | ||
| 263 | + - `bfloat16`: `ratio=1e-2, atol=5e-3` | ||
| 264 | + - `float16`: `ratio=1e-3, atol=5e-4` | ||
| 265 | +6. `ngroups > 1` 时,weight/bias 形状需为 `[N*ngroups]`,kernel 按组切分(每组访问 `W[group*N : group*N+N]`),各 group 使用各自独立的权重与偏置 | ||
| @@ -0,0 +1,155 @@ | |||
| 1 | +# Copyright (c) 2023, Tri Dao, Albert Gu | ||
| 2 | +# Copyright (c) 2026, Huawei Technologies Co., Ltd. | ||
| 3 | +# | ||
| 4 | +# Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 5 | +# you may not use this file except in compliance with the License. | ||
| 6 | + | ||
| 7 | +import torch | ||
| 8 | +import triton | ||
| 9 | + | ||
| 10 | +from mindspeed_ops.arch32.triton.layer_norm_fwd_1pass import _layer_norm_fwd_1pass_kernel | ||
| 11 | +from mindspeed_ops.api.triton.utils import get_vector_num | ||
| 12 | +from mindspeed_ops.utils import is_arch35 | ||
| 13 | + | ||
| 14 | +__all__ = ["layer_norm_fwd_1pass"] | ||
| 15 | + | ||
| 16 | +# UB total capacity on Ascend 910 series (192KB) | ||
| 17 | +_UB_TOTAL_BYTES = 192 * 1024 | ||
| 18 | +# Use 50% for double buffering | ||
| 19 | +_UB_USABLE_RATIO = 0.5 | ||
| 20 | +# Safety margin to avoid edge-case overflow | ||
| 21 | +_UB_SAFETY_FACTOR = 170 / 192 | ||
| 22 | + | ||
| 23 | + | ||
| 24 | +def _compute_bt(block_n, has_z, dtype_bytes=4): | ||
| 25 | + """Compute optimal BT (rows per task) based on UB capacity. | ||
| 26 | + | ||
| 27 | + UB budget = 192KB * (170/192) * 50% = 85KB | ||
| 28 | + Per-token peak UB usage: | ||
| 29 | + - x input: BLOCK_N * 4 bytes (fp32) | ||
| 30 | + - y output: BLOCK_N * 4 bytes (fp32) | ||
| 31 | + - xbar: BLOCK_N * 4 bytes (fp32, intermediate) | ||
| 32 | + - z gating: BLOCK_N * 4 bytes (fp32, if HAS_Z) | ||
| 33 | + Weight/bias are shared across BT rows, not counted per-token. | ||
| 34 | + """ | ||
| 35 | + usable_bytes = int(_UB_TOTAL_BYTES * _UB_SAFETY_FACTOR * _UB_USABLE_RATIO) | ||
| 36 | + | ||
| 37 | + # Per-token UB footprint: input + output + intermediate (xbar) | ||
| 38 | + s_token = block_n * dtype_bytes * 3 | ||
| 39 | + if has_z: | ||
| 40 | + s_token += block_n * dtype_bytes | ||
| 41 | + | ||
| 42 | + # Use integer division to avoid overflow | ||
| 43 | + max_bt = usable_bytes // s_token if s_token > 0 else 1 | ||
| 44 | + max_bt = max(1, max_bt) | ||
| 45 | + | ||
| 46 | + # Round down to power of 2 | ||
| 47 | + bt = 1 | ||
| 48 | + while bt * 2 <= max_bt: | ||
| 49 | + bt *= 2 | ||
| 50 | + | ||
| 51 | + return bt | ||
| 52 | + | ||
| 53 | + | ||
| 54 | +def _select_launch_params(n): | ||
| 55 | + """Select num_warps and num_stages based on hidden dimension N. | ||
| 56 | + | ||
| 57 | + Larger N benefits from more warps for parallel reduction. | ||
| 58 | + """ | ||
| 59 | + if n <= 64: | ||
| 60 | + return 2, 2 | ||
| 61 | + elif n <= 128: | ||
| 62 | + return 4, 2 | ||
| 63 | + elif n <= 512: | ||
| 64 | + return 4, 3 | ||
| 65 | + elif n <= 1024: | ||
| 66 | + return 8, 3 | ||
| 67 | + else: | ||
| 68 | + return 8, 4 | ||
| 69 | + | ||
| 70 | + | ||
| 71 | +def layer_norm_fwd_1pass( | ||
| 72 | + x: torch.Tensor, | ||
| 73 | + weight: torch.Tensor, | ||
| 74 | + bias: torch.Tensor = None, | ||
| 75 | + z: torch.Tensor = None, | ||
| 76 | + eps: float = 1e-5, | ||
| 77 | + norm_before_gate: bool = True, | ||
| 78 | + is_rms_norm: bool = False, | ||
| 79 | + ngroups: int = 1, | ||
| 80 | +) -> tuple: | ||
| 81 | + if is_arch35(): | ||
| 82 | + raise NotImplementedError("layer_norm_fwd_1pass is not supported on arch35") | ||
| 83 | + | ||
| 84 | + x = x.contiguous() | ||
| 85 | + weight = weight.contiguous() | ||
| 86 | + if bias is not None: | ||
| 87 | + bias = bias.contiguous() | ||
| 88 | + if z is not None: | ||
| 89 | + z = z.contiguous() | ||
| 90 | + | ||
| 91 | + x_shape_og = x.shape | ||
| 92 | + x = x.reshape(-1, x.shape[-1]) | ||
| 93 | + if z is not None: | ||
| 94 | + z = z.reshape(-1, z.shape[-1]) | ||
| 95 | + | ||
| 96 | + M, N_total = x.shape | ||
| 97 | + N = N_total // ngroups | ||
| 98 | + | ||
| 99 | + assert weight.numel() >= N * ngroups, f"weight size mismatch: expected at least {N * ngroups}, got {weight.numel()}" | ||
| 100 | + if bias is not None: | ||
| 101 | + assert bias.numel() >= N * ngroups, f"bias size mismatch: expected at least {N * ngroups}, got {bias.numel()}" | ||
| 102 | + | ||
| 103 | + y = torch.empty_like(x) | ||
| 104 | + if not is_rms_norm: | ||
| 105 | + mean = torch.empty((ngroups, M), dtype=torch.float32, device=x.device) | ||
| 106 | + else: | ||
| 107 | + mean = None | ||
| 108 | + rstd = torch.empty((ngroups, M), dtype=torch.float32, device=x.device) | ||
| 109 | + | ||
| 110 | + BLOCK_N = triton.next_power_of_2(N) | ||
| 111 | + | ||
| 112 | + num_core = get_vector_num() | ||
| 113 | + ngroups_step = ngroups | ||
| 114 | + | ||
| 115 | + # Dynamic BT computation based on actual UB constraints | ||
| 116 | + BT = _compute_bt(BLOCK_N, has_z=(z is not None)) | ||
| 117 | + | ||
| 118 | + # Task partitioning | ||
| 119 | + row_blocks = triton.cdiv(M, BT) | ||
| 120 | + task_num = row_blocks * ngroups_step | ||
| 121 | + | ||
| 122 | + grid = (num_core,) | ||
| 123 | + | ||
| 124 | + # Dynamic launch parameter selection | ||
| 125 | + num_warps, num_stages = _select_launch_params(N) | ||
| 126 | + | ||
| 127 | + _layer_norm_fwd_1pass_kernel[grid]( | ||
| 128 | + x, | ||
| 129 | + y, | ||
| 130 | + weight, | ||
| 131 | + bias if bias is not None else x, | ||
| 132 | + z if z is not None else x, | ||
| 133 | + mean if mean is not None else rstd, | ||
| 134 | + rstd, | ||
| 135 | + x.stride(0), | ||
| 136 | + y.stride(0), | ||
| 137 | + z.stride(0) if z is not None else x.stride(0), | ||
| 138 | + M, | ||
| 139 | + N, | ||
| 140 | + eps, | ||
| 141 | + BLOCK_N=BLOCK_N, | ||
| 142 | + HAS_BIAS=bias is not None, | ||
| 143 | + HAS_Z=z is not None, | ||
| 144 | + NORM_BEFORE_GATE=norm_before_gate, | ||
| 145 | + IS_RMS_NORM=is_rms_norm, | ||
| 146 | + BT=BT, | ||
| 147 | + ngroups_step=ngroups_step, | ||
| 148 | + task_num=task_num, | ||
| 149 | + num_core=num_core, | ||
| 150 | + num_warps=num_warps, | ||
| 151 | + num_stages=num_stages, | ||
| 152 | + ) | ||
| 153 | + | ||
| 154 | + y = y.reshape(x_shape_og) | ||
| 155 | + return y, mean, rstd | ||
| @@ -0,0 +1,113 @@ | |||
| 1 | +# Copyright (c) 2023, Tri Dao, Albert Gu | ||
| 2 | +# Copyright (c) 2026, Huawei Technologies Co., Ltd. | ||
| 3 | +# | ||
| 4 | +# Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 5 | +# you may not use this file except in compliance with the License. | ||
| 6 | + | ||
| 7 | +import triton | ||
| 8 | +import triton.language as tl | ||
| 9 | + | ||
| 10 | + | ||
| 11 | + | ||
| 12 | +def _layer_norm_fwd_1pass_kernel( | ||
| 13 | + X, | ||
| 14 | + Y, | ||
| 15 | + W, | ||
| 16 | + B, | ||
| 17 | + Z, | ||
| 18 | + Mean, | ||
| 19 | + Rstd, | ||
| 20 | + stride_x_row, | ||
| 21 | + stride_y_row, | ||
| 22 | + stride_z_row, | ||
| 23 | + M, | ||
| 24 | + N, | ||
| 25 | + eps, | ||
| 26 | + BLOCK_N: tl.constexpr, | ||
| 27 | + HAS_BIAS: tl.constexpr, | ||
| 28 | + HAS_Z: tl.constexpr, | ||
| 29 | + NORM_BEFORE_GATE: tl.constexpr, | ||
| 30 | + IS_RMS_NORM: tl.constexpr, | ||
| 31 | + BT: tl.constexpr, | ||
| 32 | + ngroups_step: tl.constexpr, | ||
| 33 | + task_num: tl.constexpr, | ||
| 34 | + num_core: tl.constexpr, | ||
| 35 | +): | ||
| 36 | + # Optimization 1: 1D physical core grid with task dispatch | ||
| 37 | + core_id = tl.program_id(0) | ||
| 38 | + | ||
| 39 | + # Optimization 7: precompute tile indices outside the task loop | ||
| 40 | + # Optimization (int32): keep index tensors in int32 to avoid int64 scalar degradation | ||
| 41 | + cols = tl.arange(0, BLOCK_N).to(tl.int32) | ||
| 42 | + col_mask = cols < N | ||
| 43 | + row_ids = tl.arange(0, BT).to(tl.int32) | ||
| 44 | + | ||
| 45 | + for task_id in tl.range(core_id, task_num, num_core): | ||
| 46 | + # Reconstruct original 2D indices from task_id (int32) | ||
| 47 | + task_id_i32 = task_id.to(tl.int32) | ||
| 48 | + row_block = task_id_i32 // ngroups_step | ||
| 49 | + group = task_id_i32 % ngroups_step | ||
| 50 | + | ||
| 51 | + # Optimization 3: precompute group offsets | ||
| 52 | + group_offset = group * N | ||
| 53 | + mean_base = group * M | ||
| 54 | + row_start = row_block * BT | ||
| 55 | + | ||
| 56 | + # Optimization 2: process BT rows as a single [BT, BLOCK_N] tile | ||
| 57 | + rows = row_start + row_ids | ||
| 58 | + row_mask = rows < M | ||
| 59 | + # 2D mask for load/store boundary; reduction relies on zeroed padding | ||
| 60 | + rc_mask = row_mask[:, None] & col_mask[None, :] | ||
| 61 | + | ||
| 62 | + # Optimization 4: load weight/bias once per task, broadcast across BT rows | ||
| 63 | + w_val = tl.load(W + group_offset + cols, mask=col_mask).to(tl.float32) | ||
| 64 | + if HAS_BIAS: | ||
| 65 | + b_val = tl.load(B + group_offset + cols, mask=col_mask).to(tl.float32) | ||
| 66 | + | ||
| 67 | + # Optimization 5: 2D pointer offsets kept in int32 | ||
| 68 | + x_off = rows[:, None] * stride_x_row + group_offset + cols[None, :] | ||
| 69 | + y_off = rows[:, None] * stride_y_row + group_offset + cols[None, :] | ||
| 70 | + | ||
| 71 | + x = tl.load(X + x_off, mask=rc_mask).to(tl.float32) | ||
| 72 | + | ||
| 73 | + # Optimization 8: gating applied before norm | ||
| 74 | + if HAS_Z and not NORM_BEFORE_GATE: | ||
| 75 | + z_off = rows[:, None] * stride_z_row + group_offset + cols[None, :] | ||
| 76 | + z_val = tl.load(Z + z_off, mask=rc_mask).to(tl.float32) | ||
| 77 | + x *= z_val * tl.sigmoid(z_val) | ||
| 78 | + | ||
| 79 | + # Optimization 1 (vectorized): single 2D reduction along the hidden axis | ||
| 80 | + if not IS_RMS_NORM: | ||
| 81 | + mean_val = tl.sum(x, axis=1) / N | ||
| 82 | + # zero out padded columns so they do not pollute variance | ||
| 83 | + xbar = tl.where(col_mask[None, :], x - mean_val[:, None], 0.0) | ||
| 84 | + var = tl.sum(xbar * xbar, axis=1) / N | ||
| 85 | + # Optimization 9: write mean early (separate write stream) | ||
| 86 | + tl.store(Mean + mean_base + rows, mean_val, mask=row_mask) | ||
| 87 | + else: | ||
| 88 | + var = tl.sum(x * x, axis=1) / N | ||
| 89 | + | ||
| 90 | + rstd_val = 1 / tl.sqrt(var + eps) | ||
| 91 | + # Optimization 9: write rstd early (separate write stream) | ||
| 92 | + tl.store(Rstd + mean_base + rows, rstd_val, mask=row_mask) | ||
| 93 | + | ||
| 94 | + # Normalize and apply linear transformation (broadcast over BT rows) | ||
| 95 | + if not IS_RMS_NORM: | ||
| 96 | + x_hat = xbar * rstd_val[:, None] | ||
| 97 | + else: | ||
| 98 | + x_hat = x * rstd_val[:, None] | ||
| 99 | + | ||
| 100 | + # Optimization 10: fuse multiply-add with weight/bias | ||
| 101 | + if HAS_BIAS: | ||
| 102 | + y_val = x_hat * w_val[None, :] + b_val[None, :] | ||
| 103 | + else: | ||
| 104 | + y_val = x_hat * w_val[None, :] | ||
| 105 | + | ||
| 106 | + # Optimization 8: gating applied after norm | ||
| 107 | + if HAS_Z and NORM_BEFORE_GATE: | ||
| 108 | + z_off = rows[:, None] * stride_z_row + group_offset + cols[None, :] | ||
| 109 | + z_val = tl.load(Z + z_off, mask=rc_mask).to(tl.float32) | ||
| 110 | + y_val *= z_val * tl.sigmoid(z_val) | ||
| 111 | + | ||
| 112 | + # Write output tile | ||
| 113 | + tl.store(Y + y_off, y_val, mask=rc_mask) | ||
| @@ -0,0 +1,50 @@ | |||
| 1 | +from atk.case_generator.generator.generate_types import GENERATOR_REGISTRY | ||
| 2 | +from atk.case_generator.generator.base_generator import CaseGenerator | ||
| 3 | +from atk.configs.case_config import CaseConfig | ||
| 4 | + | ||
| 5 | + | ||
| 6 | +CASES = [ | ||
| 7 | + (2, 64), | ||
| 8 | + (4, 128), | ||
| 9 | + (8, 256), | ||
| 10 | + (16, 512), | ||
| 11 | + (32, 1024), | ||
| 12 | + (64, 2048), | ||
| 13 | + (64, 64), | ||
| 14 | + (32, 128), | ||
| 15 | + (16, 256), | ||
| 16 | + (8, 512), | ||
| 17 | + (4, 1024), | ||
| 18 | + (2, 2048), | ||
| 19 | + (32, 256), | ||
| 20 | + (16, 1024), | ||
| 21 | + (64, 512), | ||
| 22 | + (8, 2048), | ||
| 23 | + (64, 128), | ||
| 24 | + (4, 512), | ||
| 25 | + (32, 2048), | ||
| 26 | + (16, 64), | ||
| 27 | + (8, 1024), | ||
| 28 | + (64, 256), | ||
| 29 | + (2, 512), | ||
| 30 | + (4, 2048), | ||
| 31 | +] | ||
| 32 | + | ||
| 33 | + | ||
| 34 | + | ||
| 35 | +class LayerNormFwd1passGenerator(CaseGenerator): | ||
| 36 | + _case_index = 0 | ||
| 37 | + | ||
| 38 | + def after_case_config(self, case_config: CaseConfig) -> CaseConfig: | ||
| 39 | + m, n = CASES[LayerNormFwd1passGenerator._case_index % len(CASES)] | ||
| 40 | + LayerNormFwd1passGenerator._case_index += 1 | ||
| 41 | + | ||
| 42 | + case_config.inputs[0].shape = [m, n] | ||
| 43 | + case_config.inputs[1].shape = [n] | ||
| 44 | + case_config.inputs[2].shape = [n] | ||
| 45 | + | ||
| 46 | + x_dtype = case_config.inputs[0].dtype | ||
| 47 | + case_config.inputs[1].dtype = x_dtype | ||
| 48 | + case_config.inputs[2].dtype = x_dtype | ||
| 49 | + | ||
| 50 | + return case_config | ||
| @@ -0,0 +1,52 @@ | |||
| 1 | +api: pytorch | ||
| 2 | +version: v2.1 | ||
| 3 | +name: torch_layer_norm_fwd_1pass | ||
| 4 | +triton_name: triton_layer_norm_fwd_1pass.TritonLayerNormFwd1passFunctionApi | ||
| 5 | +api_type: torch_layer_norm_fwd_1pass | ||
| 6 | +triton_api_type: triton_layer_norm_fwd_1pass | ||
| 7 | +generate: generate_layer_norm_fwd_1pass | ||
| 8 | +dtype_numbers: 8 | ||
| 9 | +standard: | ||
| 10 | + acc: single_bm | ||
| 11 | + perf: not_key | ||
| 12 | +inputs: | ||
| 13 | + - name: x | ||
| 14 | + type: tensor | ||
| 15 | + required: true | ||
| 16 | + dtypes: | ||
| 17 | + values: [ fp32, fp16, bf16 ] | ||
| 18 | + ranges: | ||
| 19 | + valid: | ||
| 20 | + values: [ [-1, 1] ] | ||
| 21 | + shapes: | ||
| 22 | + dim_numbers: | ||
| 23 | + values: [ 2 ] | ||
| 24 | + dim_values: | ||
| 25 | + values: [ 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 2048 ] | ||
| 26 | + max_length: 196608 | ||
| 27 | + - name: weight | ||
| 28 | + type: tensor | ||
| 29 | + required: true | ||
| 30 | + dtypes: | ||
| 31 | + values: [ fp32, fp16, bf16 ] | ||
| 32 | + ranges: | ||
| 33 | + valid: | ||
| 34 | + values: [ [-1, 1] ] | ||
| 35 | + shapes: | ||
| 36 | + dim_numbers: | ||
| 37 | + values: [ 1 ] | ||
| 38 | + dim_values: | ||
| 39 | + values: [ 64, 128, 256, 512, 1024, 2048 ] | ||
| 40 | + - name: bias | ||
| 41 | + type: tensor | ||
| 42 | + required: true | ||
| 43 | + dtypes: | ||
| 44 | + values: [ fp32, fp16, bf16 ] | ||
| 45 | + ranges: | ||
| 46 | + valid: | ||
| 47 | + values: [ [-1, 1] ] | ||
| 48 | + shapes: | ||
| 49 | + dim_numbers: | ||
| 50 | + values: [ 1 ] | ||
| 51 | + dim_values: | ||
| 52 | + values: [ 64, 128, 256, 512, 1024, 2048 ] | ||
| @@ -0,0 +1,56 @@ | |||
| 1 | +import torch | ||
| 2 | +from atk.configs.dataset_config import InputDataset | ||
| 3 | +from atk.tasks.api_execute import register | ||
| 4 | +from atk.tasks.api_execute.base_api import BaseApi | ||
| 5 | +from atk.tasks.api_execute.triton_base_api import TritonBaseApi | ||
| 6 | +from mindspeed_ops.api.triton.layer_norm_fwd_1pass import layer_norm_fwd_1pass | ||
| 7 | + | ||
| 8 | + | ||
| 9 | +def _layer_norm_cpu(x, weight, bias, eps=1e-5): | ||
| 10 | + x_float = x.float() | ||
| 11 | + mean = x_float.mean(dim=-1, keepdim=True) | ||
| 12 | + var = ((x_float - mean) ** 2).mean(dim=-1, keepdim=True) | ||
| 13 | + rstd = 1.0 / torch.sqrt(var + eps) | ||
| 14 | + x_hat = (x_float - mean) * rstd | ||
| 15 | + y = x_hat * weight.float() | ||
| 16 | + if bias is not None: | ||
| 17 | + y = y + bias.float() | ||
| 18 | + return y.to(x.dtype) | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | +class TorchLayerNormFwd1passFunctionApi(BaseApi): | ||
| 23 | + def __call__(self, input_data: InputDataset, with_output: bool = False): | ||
| 24 | + kwargs = input_data.kwargs | ||
| 25 | + x = kwargs.get("x") | ||
| 26 | + weight = kwargs.get("weight") | ||
| 27 | + bias = kwargs.get("bias", None) | ||
| 28 | + | ||
| 29 | + y = _layer_norm_cpu(x, weight, bias, eps=1e-5) | ||
| 30 | + | ||
| 31 | + if not with_output: | ||
| 32 | + return None | ||
| 33 | + return y | ||
| 34 | + | ||
| 35 | + | ||
| 36 | + | ||
| 37 | +class TritonLayerNormFwd1passFunctionApi(TritonBaseApi): | ||
| 38 | + def __call__(self, input_data: InputDataset, with_output: bool = False): | ||
| 39 | + kwargs = input_data.kwargs | ||
| 40 | + x = kwargs.get("x") | ||
| 41 | + weight = kwargs.get("weight") | ||
| 42 | + bias = kwargs.get("bias", None) | ||
| 43 | + | ||
| 44 | + y, mean, rstd = layer_norm_fwd_1pass( | ||
| 45 | + x, | ||
| 46 | + weight, | ||
| 47 | + bias, | ||
| 48 | + None, | ||
| 49 | + eps=1e-5, | ||
| 50 | + norm_before_gate=True, | ||
| 51 | + is_rms_norm=False, | ||
| 52 | + ) | ||
| 53 | + | ||
| 54 | + if not with_output: | ||
| 55 | + return None | ||
| 56 | + return y | ||
| @@ -0,0 +1,195 @@ | |||
| 1 | +# Copyright (c) 2026, HUAWEI CORPORATION. All rights reserved. | ||
| 2 | +# pylint: disable=duplicate-code | ||
| 3 | + | ||
| 4 | +import pytest | ||
| 5 | +import torch | ||
| 6 | + | ||
| 7 | +from mindspeed_ops.api.triton.layer_norm_fwd_1pass import layer_norm_fwd_1pass | ||
| 8 | +from mindspeed_ops.api.triton.utils import get_available_device | ||
| 9 | +from tests.utils import print_diff, assert_close | ||
| 10 | + | ||
| 11 | + | ||
| 12 | +def _dtype_thresholds(dtype): | ||
| 13 | + if dtype == torch.float32: | ||
| 14 | + return 1e-4, 1e-5 | ||
| 15 | + if dtype == torch.bfloat16: | ||
| 16 | + return 1e-2, 5e-3 | ||
| 17 | + if dtype == torch.float16: | ||
| 18 | + return 1e-3, 5e-4 | ||
| 19 | + raise ValueError(f"unsupported dtype {dtype}") | ||
| 20 | + | ||
| 21 | + | ||
| 22 | +TEST_SHAPES = [ | ||
| 23 | + (2, 64), | ||
| 24 | + (4, 128), | ||
| 25 | + (8, 256), | ||
| 26 | + (16, 512), | ||
| 27 | + (32, 1024), | ||
| 28 | + (64, 2048), | ||
| 29 | +] | ||
| 30 | + | ||
| 31 | +DTYPES = [torch.float32, torch.bfloat16, torch.float16] | ||
| 32 | + | ||
| 33 | + | ||
| 34 | +def _case_id(shape, dtype, is_rms, has_z, norm_before_gate): | ||
| 35 | + M, N = shape | ||
| 36 | + dtype_name = {torch.float32: "fp32", torch.bfloat16: "bf16", torch.float16: "fp16"} | ||
| 37 | + rms_str = "rms" if is_rms else "ln" | ||
| 38 | + z_str = "z" if has_z else "noz" | ||
| 39 | + gate_str = "nbg" if norm_before_gate else "gnb" | ||
| 40 | + return f"M{M}_N{N}_{dtype_name[dtype]}_{rms_str}_{z_str}_{gate_str}" | ||
| 41 | + | ||
| 42 | + | ||
| 43 | +TEST_CONFIGS = [ | ||
| 44 | + (is_rms, has_z, norm_before_gate) | ||
| 45 | + for is_rms in [False, True] | ||
| 46 | + for has_z in [False, True] | ||
| 47 | + for norm_before_gate in [False, True] | ||
| 48 | +] | ||
| 49 | + | ||
| 50 | +TEST_CASES = [ | ||
| 51 | + (shape, dtype, is_rms, has_z, nbg) | ||
| 52 | + for shape in TEST_SHAPES | ||
| 53 | + for dtype in DTYPES | ||
| 54 | + for (is_rms, has_z, nbg) in TEST_CONFIGS | ||
| 55 | +] | ||
| 56 | + | ||
| 57 | + | ||
| 58 | +def cpu_golden(x, weight, bias, z, eps, norm_before_gate, is_rms_norm): | ||
| 59 | + x_float = x.float() | ||
| 60 | + | ||
| 61 | + if z is not None and not norm_before_gate: | ||
| 62 | + z_float = z.float() | ||
| 63 | + x_float = x_float * z_float * torch.sigmoid(z_float) | ||
| 64 | + | ||
| 65 | + if not is_rms_norm: | ||
| 66 | + mean = x_float.mean(dim=-1, keepdim=True) | ||
| 67 | + var = ((x_float - mean) ** 2).mean(dim=-1, keepdim=True) | ||
| 68 | + else: | ||
| 69 | + mean = None | ||
| 70 | + var = (x_float**2).mean(dim=-1, keepdim=True) | ||
| 71 | + | ||
| 72 | + rstd = 1.0 / torch.sqrt(var + eps) | ||
| 73 | + | ||
| 74 | + if not is_rms_norm: | ||
| 75 | + x_hat = (x_float - mean) * rstd | ||
| 76 | + else: | ||
| 77 | + x_hat = x_float * rstd | ||
| 78 | + | ||
| 79 | + w = weight.float() | ||
| 80 | + y = x_hat * w | ||
| 81 | + if bias is not None: | ||
| 82 | + y = y + bias.float() | ||
| 83 | + | ||
| 84 | + if z is not None and norm_before_gate: | ||
| 85 | + z_float = z.float() | ||
| 86 | + y = y * z_float * torch.sigmoid(z_float) | ||
| 87 | + | ||
| 88 | + return y, mean, rstd | ||
| 89 | + | ||
| 90 | + | ||
| 91 | +class TestLayerNormFwd1PassOperator: | ||
| 92 | + def setup_method(self): | ||
| 93 | + torch.manual_seed(42) | ||
| 94 | + | ||
| 95 | + | ||
| 96 | + ("shape", "dtype", "is_rms", "has_z", "norm_before_gate"), | ||
| 97 | + [pytest.param(*test, id=_case_id(*test)) for test in TEST_CASES], | ||
| 98 | + ) | ||
| 99 | + def test_layer_norm_fwd_1pass_npu( | ||
| 100 | + self, | ||
| 101 | + shape, | ||
| 102 | + dtype: torch.dtype, | ||
| 103 | + is_rms: bool, | ||
| 104 | + has_z: bool, | ||
| 105 | + norm_before_gate: bool, | ||
| 106 | + ): | ||
| 107 | + M, N = shape | ||
| 108 | + device = get_available_device() | ||
| 109 | + | ||
| 110 | + x = torch.randn(M, N, dtype=dtype) | ||
| 111 | + weight = torch.randn(N, dtype=dtype) | ||
| 112 | + bias = torch.randn(N, dtype=dtype) | ||
| 113 | + z = torch.randn(M, N, dtype=dtype) if has_z else None | ||
| 114 | + | ||
| 115 | + y_golden_f, mean_golden, rstd_golden = cpu_golden(x, weight, bias, z, 1e-5, norm_before_gate, is_rms) | ||
| 116 | + y_golden = y_golden_f.to(dtype).to(device) | ||
| 117 | + | ||
| 118 | + x_npu = x.to(device) | ||
| 119 | + weight_npu = weight.to(device) | ||
| 120 | + bias_npu = bias.to(device) | ||
| 121 | + z_npu = z.to(device) if z is not None else None | ||
| 122 | + | ||
| 123 | + y_npu, mean_npu, rstd_npu = layer_norm_fwd_1pass( | ||
| 124 | + x_npu, | ||
| 125 | + weight_npu, | ||
| 126 | + bias_npu, | ||
| 127 | + z_npu, | ||
| 128 | + eps=1e-5, | ||
| 129 | + norm_before_gate=norm_before_gate, | ||
| 130 | + is_rms_norm=is_rms, | ||
| 131 | + ) | ||
| 132 | + | ||
| 133 | + ratio, atol = _dtype_thresholds(dtype) | ||
| 134 | + print_diff("cpu_and_triton_y", y_golden, y_npu, atol) | ||
🟡 Medium Priority 测试文件 建议:在 CPU golden 中同时计算参考的 mean 和 rstd,并对 ![]() ![]() 不准确? | |||
| 135 | + assert_close("cpu_and_triton_y", y_golden, y_npu, ratio, err_atol=atol) | ||
| 136 | + | ||
| 137 | + # mean/rstd are always computed in float32 internally; validate against golden | ||
| 138 | + rstd_ratio, rstd_atol = 1e-3, 1e-3 | ||
| 139 | + rstd_golden_d = rstd_golden.reshape(1, M).to(device) | ||
| 140 | + print_diff("cpu_and_triton_rstd", rstd_golden_d, rstd_npu, rstd_atol) | ||
| 141 | + assert_close("cpu_and_triton_rstd", rstd_golden_d, rstd_npu, rstd_ratio, err_atol=rstd_atol) | ||
| 142 | + if not is_rms: | ||
| 143 | + mean_golden_d = mean_golden.reshape(1, M).to(device) | ||
| 144 | + print_diff("cpu_and_triton_mean", mean_golden_d, mean_npu, rstd_atol) | ||
| 145 | + assert_close("cpu_and_triton_mean", mean_golden_d, mean_npu, rstd_ratio, err_atol=rstd_atol) | ||
| 146 | + | ||
| 147 | + | ||
| 148 | + def test_layer_norm_fwd_1pass_ngroups(self, dtype): | ||
| 149 | + M, N, ngroups = 8, 128, 2 | ||
| 150 | + device = get_available_device() | ||
| 151 | + | ||
| 152 | + x = torch.randn(M, N * ngroups, dtype=dtype) | ||
| 153 | + weight = torch.randn(N * ngroups, dtype=dtype) | ||
| 154 | + bias = torch.randn(N * ngroups, dtype=dtype) | ||
| 155 | + | ||
| 156 | + # CPU golden per group | ||
| 157 | + y_parts = [] | ||
| 158 | + mean_parts = [] | ||
| 159 | + rstd_parts = [] | ||
| 160 | + for g in range(ngroups): | ||
| 161 | + x_g = x[:, g * N : (g + 1) * N] | ||
| 162 | + w_g = weight[g * N : (g + 1) * N] | ||
| 163 | + b_g = bias[g * N : (g + 1) * N] | ||
| 164 | + y_g, mean_g, rstd_g = cpu_golden(x_g, w_g, b_g, None, 1e-5, True, False) | ||
| 165 | + y_parts.append(y_g) | ||
| 166 | + mean_parts.append(mean_g.reshape(1, M)) | ||
| 167 | + rstd_parts.append(rstd_g.reshape(1, M)) | ||
| 168 | + y_golden = torch.cat(y_parts, dim=-1).to(dtype).to(device) | ||
| 169 | + mean_golden = torch.cat(mean_parts, dim=0).to(device) | ||
| 170 | + rstd_golden = torch.cat(rstd_parts, dim=0).to(device) | ||
| 171 | + | ||
| 172 | + x_npu = x.to(device) | ||
| 173 | + weight_npu = weight.to(device) | ||
| 174 | + bias_npu = bias.to(device) | ||
| 175 | + | ||
| 176 | + y_npu, mean_npu, rstd_npu = layer_norm_fwd_1pass( | ||
| 177 | + x_npu, | ||
| 178 | + weight_npu, | ||
| 179 | + bias_npu, | ||
| 180 | + None, | ||
| 181 | + eps=1e-5, | ||
| 182 | + norm_before_gate=True, | ||
| 183 | + is_rms_norm=False, | ||
| 184 | + ngroups=ngroups, | ||
| 185 | + ) | ||
| 186 | + | ||
| 187 | + ratio, atol = _dtype_thresholds(dtype) | ||
| 188 | + print_diff("ngroups_y", y_golden, y_npu, atol) | ||
| 189 | + assert_close("ngroups_y", y_golden, y_npu, ratio, err_atol=atol) | ||
| 190 | + | ||
| 191 | + rstd_ratio, rstd_atol = 1e-3, 1e-3 | ||
| 192 | + print_diff("ngroups_mean", mean_golden, mean_npu, rstd_atol) | ||
| 193 | + assert_close("ngroups_mean", mean_golden, mean_npu, rstd_ratio, err_atol=rstd_atol) | ||
| 194 | + print_diff("ngroups_rstd", rstd_golden, rstd_npu, rstd_atol) | ||
| 195 | + assert_close("ngroups_rstd", rstd_golden, rstd_npu, rstd_ratio, err_atol=rstd_atol) | ||


🟡 Medium Priority
文档第 40 行将 weight 形状标注为
[N],第 42 行将 bias 形状标注为[N]。但 API 代码(layer_norm_fwd_1pass.py:99)断言weight.numel() >= N * ngroups,kernel 中通过group_offset = group * N按组切分 weight,每组访问W[group*N .. group*N+N-1]。当ngroups > 1时,若用户按文档传入形状[N]的 weight,将触发断言失败或越界访问。正确的形状应为[N * ngroups](或等价形状如[ngroups, N])。建议:将 weight 形状从
[N]改为[N * ngroups](或注明 ngroups=1 时为[N],ngroups>1 时为[N * ngroups]),bias 同理。