已开启
feat:add _layer_norm_fwd_1pass_kernel #110
gcw_K5CvmS79创建于 7月23日
feat:add _layer_norm_fwd_1pass_kernel #110
已开启
gcw_K5CvmS79创建于 7月23日
共 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 | 同输入 | 门控分支张量(可选) |
atomgit-bot
atomgit-botatomgit-bot7月23日

🟡 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 同理。

likedislike
不准确?
gcw_K5CvmS79
7月23日 评论:
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+@triton.jit
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+@GENERATOR_REGISTRY.register("generate_layer_norm_fwd_1pass")
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+@register("torch_layer_norm_fwd_1pass")
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+@register("triton_layer_norm_fwd_1pass")
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+ @pytest.mark.parametrize(
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)
atomgit-bot
atomgit-botatomgit-bot7月23日

🟡 Medium Priority

测试文件 test_layer_norm_fwd_1pass.py:131-132 仅对 y_npu 调用了 assert_close,kernel 返回的 mean_npu 和 rstd_npu(第 123 行解包)完全未被使用。kernel 中 mean/rstd 的计算(arch32:81-92)涉及归约、掩码处理和除 N 操作,若存在逻辑错误(如除 N 误写为除 BLOCK_N、掩码未生效导致 padding 污染),将不会被现有测试捕获。

建议:在 CPU golden 中同时计算参考的 mean 和 rstd,并对 mean_npu(非 RMSNorm 时)和 rstd_npu 使用合适的容差调用 assert_close 校验。

likedislike
不准确?
gcw_K5CvmS79
7月23日 评论:
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+ @pytest.mark.parametrize("dtype", DTYPES)
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)