已合并
modified md files(for readability improvement) #5496
gitee-duhuiping创建于 5月30日
modified md files(for readability improvement) #5496
已合并
gitee-duhuiping创建于 5月30日
30 个文件变更+98-91
@@ -117,7 +117,7 @@ aclnnStatus aclnnForeachAddcdivScalarV2(
117 <td>-</td>117 <td>-</td>
118 </tr>118 </tr>
119 <tr>119 <tr>
120- <td>y(aclTensorList*)</td>120+ <td>out(aclTensorList*)</td>
121 <td>输出</td>121 <td>输出</td>
122 <td>表示混合运算的输出张量列表。对应公式中的`y`。</td>122 <td>表示混合运算的输出张量列表。对应公式中的`y`。</td>
123 <td><ul><li>不支持空Tensor。</li><li>该参数中所有Tensor的数据类型保持一致。</li><li>数据类型和数据格式与入参`x1`的数据类型和数据格式一致,shape size大于等于入参`x1`的shape size。</li></ul></td>123 <td><ul><li>不支持空Tensor。</li><li>该参数中所有Tensor的数据类型保持一致。</li><li>数据类型和数据格式与入参`x1`的数据类型和数据格式一致,shape size大于等于入参`x1`的shape size。</li></ul></td>
@@ -167,7 +167,7 @@ aclnnStatus aclnnForeachMaximumScalarV2(
167 </tr>167 </tr>
168 </tbody></table>168 </tbody></table>
169 169 
170-## aclnnForeachDivScalarV2170+## aclnnForeachMaximumScalarV2
171 171 
172- **参数说明**172- **参数说明**
173 173 
@@ -153,7 +153,7 @@ aclnnStatus aclnnForeachSubScalarList(
153 <tr>153 <tr>
154 <td rowspan="2">ACLNN_ERR_PARAM_INVALID</td>154 <td rowspan="2">ACLNN_ERR_PARAM_INVALID</td>
155 <td rowspan="2">161002</td>155 <td rowspan="2">161002</td>
156- <td>x、scalar、out的数据类型不在支持的范围之内。</td>156+ <td>x、scalars、out的数据类型不在支持的范围之内。</td>
157 </tr>157 </tr>
158 <tr>158 <tr>
159 <td>x和out的数据类型不一致。</td></tr>159 <td>x和out的数据类型不一致。</td></tr>
@@ -85,7 +85,7 @@
85 <tr>85 <tr>
86 <td>epsilon</td>86 <td>epsilon</td>
87 <td>可选属性</td>87 <td>可选属性</td>
88- <td><ul><li>添加到分母中的值,以确保数值稳定,对应公式中的`epsilon`。</li><li>默认值为1e-5f。</li></ul></td>88+ <td><ul><li>添加到分母中的值,以确保数值稳定,对应公式中的`epsilon`。</li><li>默认值为1e-5。</li></ul></td>
89 <td>FLOAT32</td>89 <td>FLOAT32</td>
90 <td>-</td>90 <td>-</td>
91 </tr>91 </tr>
@@ -112,7 +112,7 @@
112 <tr>112 <tr>
113 <td>epsilon</td>113 <td>epsilon</td>
114 <td>可选属性</td>114 <td>可选属性</td>
115- <td><ul><li>表示添加到分母中的值,以确保数值稳定。对应公式中的`epsilon`。</li><li>默认值为1e-5f。</li></ul></td>115+ <td><ul><li>表示添加到分母中的值,以确保数值稳定。对应公式中的`epsilon`。</li><li>默认值为1e-5。</li></ul></td>
116 <td>FLOAT32</td>116 <td>FLOAT32</td>
117 <td>-</td>117 <td>-</td>
118 </tr>118 </tr>
@@ -92,7 +92,7 @@
92 <tr>92 <tr>
93 <td>epsilon</td>93 <td>epsilon</td>
94 <td>可选属性</td>94 <td>可选属性</td>
95- <td><ul><li>表示添加到分母中的值,以确保数值稳定。对应公式中的`epsilon`。</li><li>默认值为1e-5f。</li></ul></td>95+ <td><ul><li>表示添加到分母中的值,以确保数值稳定。对应公式中的`epsilon`。</li><li>默认值为1e-5。</li></ul></td>
96 <td>FLOAT32</td>96 <td>FLOAT32</td>
97 <td>-</td>97 <td>-</td>
98 </tr>98 </tr>
@@ -101,7 +101,7 @@
101 <tr>101 <tr>
102 <td>epsilon</td>102 <td>epsilon</td>
103 <td>可选属性</td>103 <td>可选属性</td>
104- <td><ul><li>添加到分母中的值,以确保数值稳定,对应公式中的`epsilon`。</li><li>默认值为1e-5f。</li></ul></td>104+ <td><ul><li>添加到分母中的值,以确保数值稳定,对应公式中的`epsilon`。</li><li>默认值为1e-5。</li></ul></td>
105 <td>FLOAT</td>105 <td>FLOAT</td>
106 <td>-</td>106 <td>-</td>
107 </tr>107 </tr>
@@ -122,7 +122,7 @@
122 <tr>122 <tr>
123 <td>mean</td>123 <td>mean</td>
124 <td>输出</td>124 <td>输出</td>
125- <td>输出LayerNorm算过程中(x1 + x2 + bias)的结果的均值,对应公式中的`x`的平均值。shape需要与`x1`满足broadcast关系(前几维的维度和`x1`前几维的维度相同,后面的维度为1,总维度与`x1`维度相同,前几维指`x1`的维度减去`gamma`的维度,表示不需要norm的维度)。</td>125+ <td>输出LayerNorm算过程中(x1 + x2 + bias)的结果的均值,对应公式中的`x`的平均值。shape需要与`x1`满足broadcast关系(前几维的维度和`x1`前几维的维度相同,后面的维度为1,总维度与`x1`维度相同,前几维指`x1`的维度减去`gamma`的维度,表示不需要norm的维度)。</td>
126 <td>FLOAT32</td>126 <td>FLOAT32</td>
127 <td>ND</td>127 <td>ND</td>
128 </tr>128 </tr>
@@ -179,7 +179,7 @@
179 <tr>179 <tr>
180 <td>epsilon</td>180 <td>epsilon</td>
181 <td>可选属性</td>181 <td>可选属性</td>
182- <td><ul><li>添加到分母中的值,以确保数值稳定,用于防止除0错误,对应公式中的`epsilon`。</li><li>默认值为1e-5f。</li></ul></td>182+ <td><ul><li>添加到分母中的值,以确保数值稳定,用于防止除0错误,对应公式中的`epsilon`。</li><li>默认值为1e-5。</li></ul></td>
183 <td>FLOAT</td>183 <td>FLOAT</td>
184 <td>-</td>184 <td>-</td>
185 </tr>185 </tr>
@@ -68,7 +68,7 @@
68 <tr>68 <tr>
69 <td>epsilon</td>69 <td>epsilon</td>
70 <td>可选属性</td>70 <td>可选属性</td>
71- <td><ul><li>添加到分母中的值,以确保数值稳定,用于防止除0错误,值需要大于等于零,对应公式中的eps。</li><li>默认值为1e-6f。</li></ul></td>71+ <td><ul><li>添加到分母中的值,以确保数值稳定,用于防止除0错误,值需要大于等于零,对应公式中的eps。</li><li>默认值为1e-6。</li></ul></td>
72 <td>FLOAT</td>72 <td>FLOAT</td>
73 <td>-</td>73 <td>-</td>
74 </tr>74 </tr>
@@ -72,7 +72,7 @@
72 <tr>72 <tr>
73 <td>epsilon</td>73 <td>epsilon</td>
74 <td>可选属性</td>74 <td>可选属性</td>
75- <td><ul><li>添加到分母中的值,以确保数值稳定,用于防止除0错误,对应公式中的eps。</li><li>默认值为1e-6f。</li></ul></td>75+ <td><ul><li>添加到分母中的值,以确保数值稳定,用于防止除0错误,对应公式中的eps。</li><li>默认值为1e-6。</li></ul></td>
76 <td>FLOAT32</td>76 <td>FLOAT32</td>
77 <td>-</td>77 <td>-</td>
78 </tr>78 </tr>
@@ -82,14 +82,14 @@
82 <tr>82 <tr>
83 <td>epsilon</td>83 <td>epsilon</td>
84 <td>可选属性</td>84 <td>可选属性</td>
85- <td><ul><li>添加到方差中的小值以避免除以零,对应公式中的`ε`。</li><li>默认值为1e-5f。</li></ul></td>85+ <td><ul><li>添加到方差中的小值以避免除以零,对应公式中的`ε`。</li><li>默认值为1e-5。</li></ul></td>
86 <td>FLOAT32</td>86 <td>FLOAT32</td>
87 <td>-</td>87 <td>-</td>
88 </tr>88 </tr>
89 <tr>89 <tr>
90 <td>momentum</td>90 <td>momentum</td>
91 <td>可选属性</td>91 <td>可选属性</td>
92- <td><ul><li>动量参数,用于更新训练期间的均值和方差。</li><li>默认值为0.1f。</li></ul></td>92+ <td><ul><li>动量参数,用于更新训练期间的均值和方差。</li><li>默认值为0.1。</li></ul></td>
93 <td>FLOAT32</td>93 <td>FLOAT32</td>
94 <td>-</td>94 <td>-</td>
95 </tr>95 </tr>
@@ -55,7 +55,7 @@
55 <tr>55 <tr>
56 <td>gx</td>56 <td>gx</td>
57 <td>输入</td>57 <td>输入</td>
58- <td>输入数据的梯度,用于反向传播,公式中的输入`gx`。</td><ul><li>输入数据的梯度,用于反向传播,公式中的输入`gx`。</li><li>数据类型与输入`x`的数据类型保持一致。</li></ul>58+ <td>输入数据的梯度,用于反向传播,公式中的输入`gx`。<ul><li>输入数据的梯度,用于反向传播,公式中的输入`gx`。</li><li>数据类型与输入`x`的数据类型保持一致。</li></ul></td>
59 <td>FLOAT32、FLOAT16、BFLOAT16</td>59 <td>FLOAT32、FLOAT16、BFLOAT16</td>
60 <td>ND</td>60 <td>ND</td>
61 </tr>61 </tr>
@@ -135,7 +135,7 @@
135 <tr>135 <tr>
136 <td>epsilon</td>136 <td>epsilon</td>
137 <td>可选属性</td>137 <td>可选属性</td>
138- <td><ul><li>添加到分母中的值,以确保数值稳定,对应公式中的`epsilon`。</li><li>默认值为1e-5f。</li></ul></td>138+ <td><ul><li>添加到分母中的值,以确保数值稳定,对应公式中的`epsilon`。</li><li>默认值为1e-5。</li></ul></td>
139 <td>FLOAT32</td>139 <td>FLOAT32</td>
140 <td>-</td>140 <td>-</td>
141 </tr>141 </tr>
@@ -57,7 +57,7 @@
57 <tr>57 <tr>
58 <td>epsilon</td>58 <td>epsilon</td>
59 <td>可选属性</td>59 <td>可选属性</td>
60- <td><ul><li>添加到分母中的值,以确保数值稳定,用于防止除0错误,对应公式中的`eps`。</li><li>默认值为1e-6f。</li></ul></td>60+ <td><ul><li>添加到分母中的值,以确保数值稳定,用于防止除0错误,对应公式中的`eps`。</li><li>默认值为1e-6。</li></ul></td>
61 <td>FLOAT32</td>61 <td>FLOAT32</td>
62 <td>-</td>62 <td>-</td>
63 </tr>63 </tr>
@@ -86,7 +86,7 @@
86 <tr>86 <tr>
87 <td>epsilon</td>87 <td>epsilon</td>
88 <td>可选属性</td>88 <td>可选属性</td>
89- <td><ul><li>添加到分母中的值,以确保数值稳定,对应公式中的eps。</li><li>默认值为1e-5f。</li></ul></td>89+ <td><ul><li>添加到分母中的值,以确保数值稳定,对应公式中的eps。</li><li>默认值为1e-5。</li></ul></td>
90 <td>FLOAT32</td>90 <td>FLOAT32</td>
91 <td>-</td>91 <td>-</td>
92 </tr>92 </tr>
@@ -68,7 +68,7 @@
68 <tr>68 <tr>
69 <td>epsilon</td>69 <td>epsilon</td>
70 <td>可选属性</td>70 <td>可选属性</td>
71- <td><ul><li>添加到分母中的值,以确保数值稳定,用于防止除0错误,对应公式中的`eps`。</li><li>默认值为1e-6f。</li></ul></td>71+ <td><ul><li>添加到分母中的值,以确保数值稳定,用于防止除0错误,对应公式中的`eps`。</li><li>默认值为1e-6。</li></ul></td>
72 <td>FLOAT32</td>72 <td>FLOAT32</td>
73 <td>-</td>73 <td>-</td>
74 </tr>74 </tr>
@@ -69,7 +69,7 @@
69 <tr>69 <tr>
70 <td>epsilon</td>70 <td>epsilon</td>
71 <td>可选属性</td>71 <td>可选属性</td>
72- <td><ul><li>表示添加到方差中的值,以避免出现除以零的情况。对应公式中的`ε`。</li><li>默认值为1e-6f。</li></ul></td>72+ <td><ul><li>表示添加到方差中的值,以避免出现除以零的情况。对应公式中的`ε`。</li><li>默认值为1e-6。</li></ul></td>
73 <td>FLOAT32</td>73 <td>FLOAT32</td>
74 <td>-</td>74 <td>-</td>
75 </tr>75 </tr>
@@ -112,7 +112,7 @@
112 <tr>112 <tr>
113 <td>epsilon</td>113 <td>epsilon</td>
114 <td>可选属性</td>114 <td>可选属性</td>
115- <td><ul><li>表示对应LayerNorm中的epsilon,添加到分母中的值,以确保数值稳定。对应公式中的`epsilon`。</li><li>默认值为1e-5f。</li></ul></td>115+ <td><ul><li>表示对应LayerNorm中的epsilon,添加到分母中的值,以确保数值稳定。对应公式中的`epsilon`。</li><li>默认值为1e-5。</li></ul></td>
116 <td>FLOAT</td>116 <td>FLOAT</td>
117 <td>-</td>117 <td>-</td>
118 </tr>118 </tr>
@@ -147,4 +147,4 @@
147 147 
148| 调用方式 | 样例代码 | 说明 |148| 调用方式 | 样例代码 | 说明 |
149| ---------------- | --------------------------- | --------------------------------------------------- |149| ---------------- | --------------------------- | --------------------------------------------------- |
150-| aclnn接口 | [test_aclnn_layer_norm_quant](examples/arch35/test_aclnn_layer_norm_quant.cpp) | 通过[aclnnLayerNormQuant](docs/aclnnLayerNormQuant.md)接口方式调用LayerNormQuant算子。 |150+| aclnn接口 | [test_aclnn_layer_norm_quant](examples/test_aclnn_layer_norm_quant.cpp) | 通过[aclnnLayerNormQuant](docs/aclnnLayerNormQuant.md)接口方式调用LayerNormQuant算子。 |
@@ -53,7 +53,7 @@ CANN 算子开发者需要同时维护 Host 端和 Kernel 端代码。
53 53 
54**TILING_KEY 的生成流程:**54**TILING_KEY 的生成流程:**
55 55 
56-```56+```mermaid
57输入 Tensor Shape + DataType57输入 Tensor Shape + DataType
5858
59 Host 端 Tiling 模板匹配(按优先级)59 Host 端 Tiling 模板匹配(按优先级)
@@ -135,7 +135,7 @@ Tiling 模板注册时附带一个优先级数值(如 Welford 为 4000,TwoPa
135 135 
136**double buffer** 是一种经典的 **流水线优化技术**:在 UB 中分配两块相同大小的缓冲区(Buffer A 和 Buffer B),使 **数据搬运****计算** 可以重叠执行:136**double buffer** 是一种经典的 **流水线优化技术**:在 UB 中分配两块相同大小的缓冲区(Buffer A 和 Buffer B),使 **数据搬运****计算** 可以重叠执行:
137 137 
138-```138+```mermaid
139时间线 → [搬入A] [计算A+搬入B] [计算B+搬入A] [计算A+...] ...139时间线 → [搬入A] [计算A+搬入B] [计算B+搬入A] [计算A+...] ...
140 ↑ ↑ ↑140 ↑ ↑ ↑
141 A就绪 搬运与计算并行 搬运与计算并行141 A就绪 搬运与计算并行 搬运与计算并行
@@ -152,18 +152,20 @@ Tiling 模板注册时附带一个优先级数值(如 Welford 为 4000,TwoPa
152由于浮点数的有限精度(float32 约为 7 位有效十进制数字,FP16 约 3 位,BF16 约 2 位),加法操作不满足数学上的结合律。即 `(a + b) + c ≠ a + (b + c)` 在浮点运算下可能成立。152由于浮点数的有限精度(float32 约为 7 位有效十进制数字,FP16 约 3 位,BF16 约 2 位),加法操作不满足数学上的结合律。即 `(a + b) + c ≠ a + (b + c)` 在浮点运算下可能成立。
153 153 
154**具体例子:**154**具体例子:**
155-```155+ 
156+```Cpp
156(1.0e10 + 1.0) - 1.0e10 → 0.0 (1.0 被"吃掉"了)157(1.0e10 + 1.0) - 1.0e10 → 0.0 (1.0 被"吃掉"了)
1571.0e10 + (1.0 - 1.0e10) → 0.0 (同样丢失)1581.0e10 + (1.0 - 1.0e10) → 0.0 (同样丢失)
1581.0 + (1.0e10 - 1.0e10) → 1.0 (正确)1591.0 + (1.0e10 - 1.0e10) → 1.0 (正确)
159```160```
161+ 
160当大数与小数相加时,小数的低位有效数字会被舍入误差湮没。在深度学习中,这种误差在梯度累加、归一化统计量计算等场景中尤为常见,因为 batch 内元素值的量级差异往往很大。162当大数与小数相加时,小数的低位有效数字会被舍入误差湮没。在深度学习中,这种误差在梯度累加、归一化统计量计算等场景中尤为常见,因为 batch 内元素值的量级差异往往很大。
161 163 
162#### 0.4.2 catastrophic cancellation(灾难性抵消 / 相消误差)164#### 0.4.2 catastrophic cancellation(灾难性抵消 / 相消误差)
163 165 
164当两个数值相近的大数相减时,结果的有效数字位数急剧减少,相对误差被极度放大。例如:166当两个数值相近的大数相减时,结果的有效数字位数急剧减少,相对误差被极度放大。例如:
165 167 
166-```168+```Cpp
167√(1.0000001 × 10^10) - √(1.0000000 × 10^10)169√(1.0000001 × 10^10) - √(1.0000000 × 10^10)
168 ↑ 真值约为 0.0005 ↑170 ↑ 真值约为 0.0005 ↑
169但由于浮点表示精度有限,这两个大数在浮点寄存器中可能完全相同,相减结果直接为 0,相对误差为 100%。171但由于浮点表示精度有限,这两个大数在浮点寄存器中可能完全相同,相减结果直接为 0,相对误差为 100%。
@@ -177,7 +179,7 @@ Tiling 模板注册时附带一个优先级数值(如 Welford 为 4000,TwoPa
177 179 
178Layer Normalization 是深度学习中的基础算子,其核心计算为:180Layer Normalization 是深度学习中的基础算子,其核心计算为:
179 181 
180-```182+```Cpp
181mean = E[x] // 均值183mean = E[x] // 均值
182rstd = 1 / sqrt(Var[x] + eps) // 标准差的倒数 (reciprocal standard deviation)184rstd = 1 / sqrt(Var[x] + eps) // 标准差的倒数 (reciprocal standard deviation)
183y = gamma * ((x - mean) * rstd) + beta // 归一化 + 仿射变换185y = gamma * ((x - mean) * rstd) + beta // 归一化 + 仿射变换
@@ -229,7 +231,7 @@ Welford 算法 [1] 是一种**在线算法(online algorithm)**,能够以**
229 231 
230朴素方差计算公式为:232朴素方差计算公式为:
231 233 
232-```234+```Cpp
233σ² = Sum[(x_i - μ)²] / n for i = 1 to n235σ² = Sum[(x_i - μ)²] / n for i = 1 to n
234```236```
235 237 
@@ -239,7 +241,7 @@ Welford 的洞察是使用递推式更新,设前 n-1 个元素的均值和二
239 241 
240**Step 1 — 增量更新均值:**242**Step 1 — 增量更新均值:**
241 243 
242-```244+```Cpp
243μ_n = μ_{n-1} + (x_n - μ_{n-1}) / n245μ_n = μ_{n-1} + (x_n - μ_{n-1}) / n
244```246```
245 247 
@@ -249,19 +251,19 @@ Welford 的洞察是使用递推式更新,设前 n-1 个元素的均值和二
249 251 
250二阶中心矩定义为 `M_{2,n} = Sum[(x_i - μ_n)²]`(i=1 to n)。其递推式为:252二阶中心矩定义为 `M_{2,n} = Sum[(x_i - μ_n)²]`(i=1 to n)。其递推式为:
251 253 
252-```254+```Cpp
253M_{2,n} = M_{2,n-1} + (x_n - μ_{n-1}) * (x_n - μ_n)255M_{2,n} = M_{2,n-1} + (x_n - μ_{n-1}) * (x_n - μ_n)
254```256```
255 257 
256代入 delta 后得到:258代入 delta 后得到:
257 259 
258-```260+```Cpp
259M_{2,n} = M_{2,n-1} + delta * (x_n - μ_n)261M_{2,n} = M_{2,n-1} + delta * (x_n - μ_n)
260```262```
261 263 
262**Step 3 — 方差归一化:**264**Step 3 — 方差归一化:**
263 265 
264-```266+```Cpp
265σ_n² = M_{2,n} / n267σ_n² = M_{2,n} / n
266```268```
267 269 
@@ -269,7 +271,7 @@ M_{2,n} = M_{2,n-1} + delta * (x_n - μ_n)
269 271 
270**最终更新公式(代码中使用的形式):**272**最终更新公式(代码中使用的形式):**
271 273 
272-```274+```Cpp
273delta = x - mean // mean 为旧均值 μ_{n-1}275delta = x - mean // mean 为旧均值 μ_{n-1}
274mean += delta / n // 更新为新均值 μ_n276mean += delta / n // 更新为新均值 μ_n
275M2 += delta * (x - mean_new) // mean_new 即更新后的 μ_n277M2 += delta * (x - mean_new) // mean_new 即更新后的 μ_n
@@ -280,7 +282,7 @@ variance = M2 / n // 当前方差
280 282 
281**M2** 是 "**Second Central Moment**"(二阶中心矩)的缩写。它是**未经归一化的方差**283**M2** 是 "**Second Central Moment**"(二阶中心矩)的缩写。它是**未经归一化的方差**
282 284 
283-```285+```Cpp
284M2 = Sum[(x_i - μ_n)²] = n * σ² for i = 1 to n286M2 = Sum[(x_i - μ_n)²] = n * σ² for i = 1 to n
285```287```
286 288 
@@ -294,7 +296,7 @@ M2 = Sum[(x_i - μ_n)²] = n * σ² for i = 1 to n
294 296 
295**核心流程** (`Process()` 方法, 第109-173行):297**核心流程** (`Process()` 方法, 第109-173行):
296 298 
297-```299+```Cpp
298对每一行 i (0 ~ currentBlockFactor):300对每一行 i (0 ~ currentBlockFactor):
299 WelfordInitialize(mean, variance) // 将均值累加器清零、M2 累加器清零301 WelfordInitialize(mean, variance) // 将均值累加器清零、M2 累加器清零
300 for each tile:302 for each tile:
@@ -310,7 +312,7 @@ M2 = Sum[(x_i - μ_n)²] = n * σ² for i = 1 to n
310 312 
311`WelfordUpdate` 是 CANN 框架在 AI Core 上提供的 **向量级内置函数**(vector intrinsic),用于在一条 SIMD 指令中完成对 64 个 float32 元素的 Welford 递推更新。它不是用户自定义函数,而是硬件指令的编译器封装。其逻辑等价于:313`WelfordUpdate` 是 CANN 框架在 AI Core 上提供的 **向量级内置函数**(vector intrinsic),用于在一条 SIMD 指令中完成对 64 个 float32 元素的 Welford 递推更新。它不是用户自定义函数,而是硬件指令的编译器封装。其逻辑等价于:
312 314 
313-```315+```Cpp
314for i in range(VL_F32): // VL_F32 = 64, SIMD 并行执行316for i in range(VL_F32): // VL_F32 = 64, SIMD 并行执行
315 delta = data[i] - mean_acc317 delta = data[i] - mean_acc
316 mean_acc += delta / current_n318 mean_acc += delta / current_n
@@ -323,7 +325,7 @@ for i in range(VL_F32): // VL_F32 = 64, SIMD 并行执行
323 325 
324**Tile 划分策略** (`op_host/layer_norm_v4_welford_tiling.cpp`):326**Tile 划分策略** (`op_host/layer_norm_v4_welford_tiling.cpp`):
325 327 
326-```328+```Cpp
327tileLength = 最大的满足 UB 空间约束的 64 对齐值329tileLength = 最大的满足 UB 空间约束的 64 对齐值
328welfordUpdateTimes = N / tileLength330welfordUpdateTimes = N / tileLength
329welfordUpdateTail = N % tileLength331welfordUpdateTail = N % tileLength
@@ -367,7 +369,7 @@ Tile 大小通过 `IsValidTileLength()` 函数校验,确保 **double buffer(
367 369 
368**算法步骤:**370**算法步骤:**
369 371 
370-```372+```Cpp
371输入: R 个元素, 计算 PowerOfTwoForR = 大于等于 R 的最小 2 的幂373输入: R 个元素, 计算 PowerOfTwoForR = 大于等于 R 的最小 2 的幂
372Step 1: 在 binaryAddQuotient 处拆分374Step 1: 在 binaryAddQuotient 处拆分
373 left_part = x[0 : binaryAddQuotient - 1] // 2 的幂次部分375 left_part = x[0 : binaryAddQuotient - 1] // 2 的幂次部分
@@ -387,10 +389,11 @@ Step 3: 二叉归并树 (代码第548-558行)
387 // 归并完成后 tmp[0] 即为最终总和389 // 归并完成后 tmp[0] 即为最终总和
388```390```
389 391 
390-**为什么二分累加精度更高?** 二叉归并树的每一层中,相加的两个数数量级接近(因为来自等长的子块),避免了"大数吃小数"。信息从叶子到根 O(log n) 层传播,舍入误差仅随 log n 增长(而非线性累加的 O(n) 增长)。392+**为什么二分累加精度更高?** 二叉归并树的每一层中,相加的两个数数量级接近(因为来自等长的子块),避免了"大数吃小数"。信息从叶子到根 O(log n) 层传播,舍入误差仅随log n增长(而非线性累加的 O(n) 增长)。
391 393 
392**图示(以 8 个元素为例):**394**图示(以 8 个元素为例):**
393-```395+ 
396+```mermaid
394Level 0 (叶子): a0 a1 a2 a3 a4 a5 a6 a7397Level 0 (叶子): a0 a1 a2 a3 a4 a5 a6 a7
395 \ / \ / \ / \ /398 \ / \ / \ / \ /
396Level 1: s0+s1 s2+s3 s4+s5 s6+s7 ← 两两数量级相近399Level 1: s0+s1 s2+s3 s4+s5 s6+s7 ← 两两数量级相近
@@ -441,18 +444,22 @@ while (curBinaryAddNum < binaryAddNum) {
441与朴素线性累加相比,二分累加将 O(n) 的链式加法转化为 O(log n) 的多层树状加法:444与朴素线性累加相比,二分累加将 O(n) 的链式加法转化为 O(log n) 的多层树状加法:
442 445 
443**线性累加(精度差):**446**线性累加(精度差):**
444-```447+ 
448+```Cpp
445sum = ((...(a0 + a1) + a2) + ...) + a_n449sum = ((...(a0 + a1) + a2) + ...) + a_n
446```450```
451+ 
447误差随 N 线性增长。每步都可能发生大数吃小数。452误差随 N 线性增长。每步都可能发生大数吃小数。
448 453 
449**二分累加(精度优):**454**二分累加(精度优):**
450-```455+ 
456+```Cpp
451level 0: a0+a1, a2+a3, a4+a5, ... // 子块内部累加457level 0: a0+a1, a2+a3, a4+a5, ... // 子块内部累加
452level 1: (a0+a1)+(a2+a3), ... // 子块之间归并458level 1: (a0+a1)+(a2+a3), ... // 子块之间归并
453...459...
454level k: 最终总和460level k: 最终总和
455```461```
462+ 
456每级加法中操作数数量级相近,避免了大数吃小数。误差仅随 log N 增长。463每级加法中操作数数量级相近,避免了大数吃小数。误差仅随 log N 增长。
457 464 
458(二分累加将线性累加 O(n) 的误差增长降低为 O(log n),这一结论源自浮点加法舍入误差的标准分析,详见拓展阅读 Higham 教材第 4 章。)465(二分累加将线性累加 O(n) 的误差增长降低为 O(log n),这一结论源自浮点加法舍入误差的标准分析,详见拓展阅读 Higham 教材第 4 章。)
@@ -471,13 +478,16 @@ level k: 最终总和
471 478 
472**3. 行配对 + FusedMulDstAdd 融合指令:**479**3. 行配对 + FusedMulDstAdd 融合指令:**
473`CalculateNormalizeVF` 中每次处理 **两行**(row pairing),使用 `FusedMulDstAdd` 指令。`FusedMulDstAdd` 是 CANN 的 **向量融合指令**,在单个指令周期内同时完成:480`CalculateNormalizeVF` 中每次处理 **两行**(row pairing),使用 `FusedMulDstAdd` 指令。`FusedMulDstAdd` 是 CANN 的 **向量融合指令**,在单个指令周期内同时完成:
474-```481+ 
482+```Cpp
475dst = src * gamma + beta483dst = src * gamma + beta
476```484```
485+ 
477即将 `(x - mean) * rstd * gamma + beta` 中的 `* gamma + beta` 两步融合为一条指令,减少指令发射次数和寄存器占用。486即将 `(x - mean) * rstd * gamma + beta` 中的 `* gamma + beta` 两步融合为一条指令,减少指令发射次数和寄存器占用。
478 487 
479**4. 多级归约(lastLoopNums 模板参数):**488**4. 多级归约(lastLoopNums 模板参数):**
480`LAST_LOOP_NUMS` 模板参数控制二分累加树**最后一级**的归并策略:489`LAST_LOOP_NUMS` 模板参数控制二分累加树**最后一级**的归并策略:
490+ 
481- `LAST_LOOP_NUMS = 1`:最后一级使用单次 `ReduceSum` 指令完成归约(简洁直接)491- `LAST_LOOP_NUMS = 1`:最后一级使用单次 `ReduceSum` 指令完成归约(简洁直接)
482- `LAST_LOOP_NUMS = 2`:最后一级先使用 `ShiftLefts` 将元素左移(对齐累加器的指数位),再执行 `Add` 完成归约。`ShiftLefts` 是向量移位指令,在此上下文中用于将多个累加器的数值指数对齐后再相加,减少因指数不匹配引入的舍入误差。这种方式在 R 维度足够长时能获得更高的累加精度,代价是增加少量指令。492- `LAST_LOOP_NUMS = 2`:最后一级先使用 `ShiftLefts` 将元素左移(对齐累加器的指数位),再执行 `Add` 完成归约。`ShiftLefts` 是向量移位指令,在此上下文中用于将多个累加器的数值指数对齐后再相加,减少因指数不匹配引入的舍入误差。这种方式在 R 维度足够长时能获得更高的累加精度,代价是增加少量指令。
483 493 
@@ -523,7 +533,7 @@ CANN 框架在 Host 端按照注册的优先级顺序依次尝试匹配 Tiling
523 533 
524**条件之间的逻辑关系(由严格到宽松):**534**条件之间的逻辑关系(由严格到宽松):**
525 535 
526-```536+```Cpp
527Welford 的 IsCapable 条件: isRegBase == true ← 最宽松537Welford 的 IsCapable 条件: isRegBase == true ← 最宽松
528TwoPass 的 IsCapable 条件: isRegBase + UB约束 ← 中等538TwoPass 的 IsCapable 条件: isRegBase + UB约束 ← 中等
529TwoPassPerf 的 IsCapable 条件: isRegBase + UB约束 + R≤8192 + 非超大A/R ← 最严格539TwoPassPerf 的 IsCapable 条件: isRegBase + UB约束 + R≤8192 + 非超大A/R ← 最严格
@@ -544,6 +554,7 @@ TwoPassPerf 的 IsCapable 条件: isRegBase + UB约束 + R≤8192 + 非超大A/R
544**设计哲学:** 三套算法不是互相替代的关系,而是对 **"shape 空间"** 的一个覆盖性划分。框架通过优先级(4000→200→150)对常见 case 做最优配置,通过 IsCapable 条件在极端 case 下做兜底保证。554**设计哲学:** 三套算法不是互相替代的关系,而是对 **"shape 空间"** 的一个覆盖性划分。框架通过优先级(4000→200→150)对常见 case 做最优配置,通过 IsCapable 条件在极端 case 下做兜底保证。
545 555 
546对于软件领域的读者来说,这种"多算法覆盖 shape 空间"的模式类似于:556对于软件领域的读者来说,这种"多算法覆盖 shape 空间"的模式类似于:
557+ 
547- **编译器的多级优化 pass**——不同优化等级适用不同代码特征558- **编译器的多级优化 pass**——不同优化等级适用不同代码特征
548- **数据库查询优化器的多计划候选**——基于 cost model 选择最优执行计划559- **数据库查询优化器的多计划候选**——基于 cost model 选择最优执行计划
549- **BLAS 库的 kernel 选择**——MKL/OpenBLAS 根据矩阵大小动态切换不同 kernel560- **BLAS 库的 kernel 选择**——MKL/OpenBLAS 根据矩阵大小动态切换不同 kernel
@@ -565,6 +576,7 @@ TwoPassPerf 的 IsCapable 条件: isRegBase + UB约束 + R≤8192 + 非超大A/R
565| `op_host/layer_norm_v4_def.cpp:192` | Host CPU | 950 平台配置,注册为 `layer_norm_v4_apt` |576| `op_host/layer_norm_v4_def.cpp:192` | Host CPU | 950 平台配置,注册为 `layer_norm_v4_apt` |
566 577 
567代码分层小结:578代码分层小结:
579+ 
568- `op_kernel/arch35/` 下的 `.h` 文件是 **架构相关的设备端 kernel 代码**,运行在 Ascend 950 的 AI Core 上580- `op_kernel/arch35/` 下的 `.h` 文件是 **架构相关的设备端 kernel 代码**,运行在 Ascend 950 的 AI Core 上
569- `op_host/` 下的 `.cpp`/`.h` 文件是 **Host 端控制代码**,运行在服务器 CPU 上,负责参数计算和调度581- `op_host/` 下的 `.cpp`/`.h` 文件是 **Host 端控制代码**,运行在服务器 CPU 上,负责参数计算和调度
570- 开发一个算子需要同时维护两端代码——Host 端做"策略计算",Kernel 端做"策略执行"582- 开发一个算子需要同时维护两端代码——Host 端做"策略计算",Kernel 端做"策略执行"
@@ -123,7 +123,7 @@
123 <tr>123 <tr>
124 <td>epsilon</td>124 <td>epsilon</td>
125 <td>可选属性</td>125 <td>可选属性</td>
126- <td><ul><li>添加到分母中的值,以确保数值稳定,对应公式中的`eps`。</li><li>默认值为1e-5f。</li></ul></td>126+ <td><ul><li>添加到分母中的值,以确保数值稳定,对应公式中的`eps`。</li><li>默认值为1e-5。</li></ul></td>
127 <td>FLOAT32</td>127 <td>FLOAT32</td>
128 <td>-</td>128 <td>-</td>
129 </tr>129 </tr>
@@ -100,7 +100,7 @@
100 <td>output_zero_point</td>100 <td>output_zero_point</td>
101 <td>输入</td>101 <td>输入</td>
102 <td>输入标量,模型输出数据的偏置,对应公式中的`outputZeroPoint`。传入值不能超过input对应数据类型的上下边界,例如INT8上下边界为[-128,127]。</td>102 <td>输入标量,模型输出数据的偏置,对应公式中的`outputZeroPoint`。传入值不能超过input对应数据类型的上下边界,例如INT8上下边界为[-128,127]。</td>
103- <td>FLOAT32</td>103+ <td>INT32</td>
104 <td>ND</td>104 <td>ND</td>
105 </tr>105 </tr>
106 <tr>106 <tr>
@@ -57,7 +57,7 @@
57 <tr>57 <tr>
58 <td>epsilon</td>58 <td>epsilon</td>
59 <td>可选属性</td>59 <td>可选属性</td>
60- <td><ul><li>添加到分母中的值,以确保数值稳定,对应公式中的eps。</li><li>默认值为1e-6f。</li></ul></td>60+ <td><ul><li>添加到分母中的值,以确保数值稳定,对应公式中的eps。</li><li>默认值为1e-6。</li></ul></td>
61 <td>FLOAT32</td>61 <td>FLOAT32</td>
62 <td>-</td>62 <td>-</td>
63 </tr>63 </tr>
@@ -171,4 +171,4 @@
171| 调用方式 | 样例代码 | 说明 |171| 调用方式 | 样例代码 | 说明 |
172| ---------------- | --------------------------- | --------------------------------------------------- |172| ---------------- | --------------------------- | --------------------------------------------------- |
173| aclnn接口 | [test_aclnn_rms_norm_quant_v3](../rms_norm_quant_v3/examples/test_aclnn_rms_norm_quant_v3.cpp) | 通过[aclnnRmsNormQuantV3](../rms_norm_quant_v3/docs/aclnnRmsNormQuantV3.md)接口方式调用RmsNormQuantV3算子。 |173| aclnn接口 | [test_aclnn_rms_norm_quant_v3](../rms_norm_quant_v3/examples/test_aclnn_rms_norm_quant_v3.cpp) | 通过[aclnnRmsNormQuantV3](../rms_norm_quant_v3/docs/aclnnRmsNormQuantV3.md)接口方式调用RmsNormQuantV3算子。 |
174-| 图模式 | - | 通过[算子IR](op_graph/rms_norm_quant_v3_proto.h)构图方式调用RmsNormQuantV3算子。 |174+| 图模式 | - | 通过[算子IR](op_graph/rms_norm_quant_v3_proto.h)构图方式调用RmsNormQuantV3算子。 |
@@ -18,19 +18,21 @@
18- 接口功能:RmsNorm算子是大模型常用的标准化操作,相比LayerNorm算子,其去掉了减去均值的部分。RmsNormQuantV3算子将RmsNorm算子以及RmsNorm后的Quantize算子融合起来,减少搬入搬出操作。同时在RmsNormQuantV2算子的基础上新增了Rstd的输出。18- 接口功能:RmsNorm算子是大模型常用的标准化操作,相比LayerNorm算子,其去掉了减去均值的部分。RmsNormQuantV3算子将RmsNorm算子以及RmsNorm后的Quantize算子融合起来,减少搬入搬出操作。同时在RmsNormQuantV2算子的基础上新增了Rstd的输出。
19- 计算公式:19- 计算公式:
20 20 
21-$$21+ $$
22-quant\_in_i=\frac{x_i}{\operatorname{Rms}(\mathbf{x})} gamma_i + beta_i, \quad \text { where } \operatorname{Rms}(\mathbf{x})=\sqrt{\frac{1}{n} \sum_{i=1}^n x_i^2+epsilon}22+ quant\_in_i=\frac{x_i}{\operatorname{Rms}(\mathbf{x})} gamma_i + beta_i, \quad \text { where } \operatorname{Rms}(\mathbf{x})=\sqrt{\frac{1}{n} \sum_{i=1}^n x_i^2+epsilon}
23-$$23+ $$
24 24 
25-- divMode为True时:25+ - divMode为True时:
26-$$26+
27-y=round((quant\_in/scale)+offset)27+ $$
28-$$28+ y=round((quant\_in/scale)+offset)
29- 29+ $$
30-- divMode为False时:30+
31-$$31+ - divMode为False时:
32-y=round((quant\_in*scale)+offset)32+
33-$$33+ $$
34+ y=round((quant\_in*scale)+offset)
35+ $$
34 36 
35## 函数原型37## 函数原型
36 38 
@@ -207,7 +209,6 @@ aclnnStatus aclnnRmsNormQuantV3(
207 </tr>209 </tr>
208 </tbody>210 </tbody>
209 </table>211 </table>
210-
211 212 
212- **返回值**213- **返回值**
213 214 
@@ -244,8 +245,6 @@ aclnnStatus aclnnRmsNormQuantV3(
244 <tr>245 <tr>
245 <td>gamma不满足维度数为1-2维,scale不满足维度数为1维。</td>246 <td>gamma不满足维度数为1-2维,scale不满足维度数为1维。</td>
246 </tr>247 </tr>
247- 
248- 
249 <tr>248 <tr>
250 <td rowspan="7">ACLNN_ERR_INNER_TILING_ERROR</td>249 <td rowspan="7">ACLNN_ERR_INNER_TILING_ERROR</td>
251 <td rowspan="7">561002</td>250 <td rowspan="7">561002</td>
@@ -17,9 +17,9 @@
17 17 
18- 计算公式:18- 计算公式:
19 19 
20-$$20+ $$
21-gradInput = ({gradOut} - {meanDy}) - ((input - mean) * (invstd^{2} * {meanDyXmu})) * invstd * weight21+ gradInput = ({gradOut} - {meanDy}) - ((input - mean) * (invstd^{2} * {meanDyXmu})) * invstd * weight
22-$$22+ $$
23 23 
24## 参数说明24## 参数说明
25 25 
@@ -27,15 +27,12 @@ $$
27 27 
28## 参数说明28## 参数说明
29 29 
30- <table style="undefined;table-layout: fixed; width: 1550px"><colgroup>30+ <table style="undefined;table-layout: fixed; width: 1005px"><colgroup>
31 <col style="width: 170px">31 <col style="width: 170px">
32- <col style="width: 120px">32+ <col style="width: 170px">
33- <col style="width: 271px">33+ <col style="width: 352px">
34- <col style="width: 330px">34+ <col style="width: 213px">
35- <col style="width: 223px">35+ <col style="width: 100px">
36- <col style="width: 101px">
37- <col style="width: 190px">
38- <col style="width: 145px">
39 </colgroup>36 </colgroup>
40 <thead>37 <thead>
41 <tr>38 <tr>
@@ -18,20 +18,17 @@
18- 计算公式:18- 计算公式:
19 19 
20$$20$$
21-runningMeanUpdate = (mean * momentum) + runningMean * (1 - momentum)21+running_mean_update = (mean * momentum) + running_mean * (1 - momentum)
22$$22$$
23 23 
24## 参数说明24## 参数说明
25 25 
26- <table style="undefined;table-layout: fixed; width: 1550px"><colgroup>26+ <table style="undefined;table-layout: fixed; width: 1005px"><colgroup>
27 <col style="width: 170px">27 <col style="width: 170px">
28 <col style="width: 170px">28 <col style="width: 170px">
29- <col style="width: 271px">29+ <col style="width: 352px">
30- <col style="width: 330px">30+ <col style="width: 213px">
31- <col style="width: 223px">31+ <col style="width: 100px">
32- <col style="width: 101px">
33- <col style="width: 190px">
34- <col style="width: 145px">
35 </colgroup>32 </colgroup>
36 <thead>33 <thead>
37 <tr>34 <tr>
@@ -52,7 +49,7 @@ $$
52 <tr>49 <tr>
53 <td>running_mean</td>50 <td>running_mean</td>
54 <td>输入</td>51 <td>输入</td>
55- <td>表示计算过程中的均值,对应公式中的`runningMean`。</td>52+ <td>表示计算过程中的均值,对应公式中的`running_mean`。</td>
56 <td>FLOAT32、FLOAT16、BFLOAT16</td>53 <td>FLOAT32、FLOAT16、BFLOAT16</td>
57 <td>ND</td>54 <td>ND</td>
58 </tr>55 </tr>
@@ -66,7 +63,7 @@ $$
66 <tr>63 <tr>
67 <td>running_mean_update</td>64 <td>running_mean_update</td>
68 <td>输出</td>65 <td>输出</td>
69- <td>更新后的均值,对应公式中的`runningMeanUpdate`。</td>66+ <td>更新后的均值,对应公式中的`running_mean_update`。</td>
70 <td>FLOAT32、FLOAT16、BFLOAT16</td>67 <td>FLOAT32、FLOAT16、BFLOAT16</td>
71 <td>ND</td>68 <td>ND</td>
72 </tr>69 </tr>
@@ -81,5 +78,5 @@ $$
81 78 
82| 调用方式 | 样例代码 | 说明 |79| 调用方式 | 样例代码 | 说明 |
83| ---------------- | --------------------------- | --------------------------------------------------- |80| ---------------- | --------------------------- | --------------------------------------------------- |
84-| aclnn接口 | [test_aclnn_BatchNormGatherStatsWithCounts](../sync_batch_norm_gather_stats_with_counts/examples/test_aclnn_BatchNormGatherStatsWithCounts.cpp) | 通过[aclnnBatchNormGatherStatsWithCounts](../sync_batch_norm_gather_stats_with_counts/docs/aclnnBatchNormGatherStatsWithCounts.md)接口方式调用SyncBNTrainingUpdate。 |81+| aclnn接口 | [test_aclnn_BatchNormGatherStatsWithCounts](../sync_batch_norm_gather_stats_with_counts/examples/test_aclnn_BatchNormGatherStatsWithCounts.cpp) | 通过[aclnnBatchNormGatherStatsWithCounts](../sync_batch_norm_gather_stats_with_counts/docs/aclnnBatchNormGatherStatsWithCounts.md)接口方式调用SyncBNTrainingUpdate算子。 |
85| 图模式 | [test_geir_sync_bn_training_update](../sync_bn_training_update/examples/test_geir_sync_bn_training_update.cpp) | 通过[算子IR](op_graph/sync_bn_training_update_proto.h)构图方式调用SyncBNTrainingUpdate算子。 |82| 图模式 | [test_geir_sync_bn_training_update](../sync_bn_training_update/examples/test_geir_sync_bn_training_update.cpp) | 通过[算子IR](op_graph/sync_bn_training_update_proto.h)构图方式调用SyncBNTrainingUpdate算子。 |
@@ -92,7 +92,7 @@
92 - 将输入x在axis维度上按k = blocksize个数分组,一组k个数 $\{\{V_i\}_{i=1}^{k}\}$ 动态量化为 $\{mxscale1, \{P_i\}_{i=1}^{k}\}$, k = blocksize:92 - 将输入x在axis维度上按k = blocksize个数分组,一组k个数 $\{\{V_i\}_{i=1}^{k}\}$ 动态量化为 $\{mxscale1, \{P_i\}_{i=1}^{k}\}$, k = blocksize:
93 93 
94 $$94 $$
95- shared\_exp = \begin{cases} ceil(log_2(max_i(|V_i|))) - emax, & \text{如果} 尾数位的高比特前一/两位 \text{为1,且尾数不全为0} \\ floor(log_2(max_i(|V_i|))) - emax, & \text{其} \end{cases} \\95+ shared\_exp = \begin{cases} ceil(log_2(max_i(|V_i|))) - emax, & \text{如果} 尾数位的高比特前一/两位 \text{为1,且尾数不全为0} \\ floor(log_2(max_i(|V_i|))) - emax, & \text{其} \end{cases} \\
96 $$96 $$
97 97 
98 $$98 $$
@@ -28,9 +28,10 @@
28 $$28 $$
29 3. `scale`按bit位取高19位截断,存储于`out`的bit位32位处,并将46位修改为1。29 3. `scale`按bit位取高19位截断,存储于`out`的bit位32位处,并将46位修改为1。
30 30 
31- $$31+ $$
32- out = out\ |\ (scale\ \&\ 0xFFFFE000)\ |\ (1\ll46)32+ out = out\ |\ (scale\ \&\ 0xFFFFE000)\ |\ (1\ll46)
33- $$33+ $$
34+ 
34 4. 根据`offset`取值进行后续计算:35 4. 根据`offset`取值进行后续计算:
35 36 
36 -`offset`不存在,不再进行后续计算。37 -`offset`不存在,不再进行后续计算。
@@ -21,9 +21,10 @@
21 21 
22 2. `scale`按bit位取高19位截断,存储于`out`的bit位32位处,并将46位修改为1。22 2. `scale`按bit位取高19位截断,存储于`out`的bit位32位处,并将46位修改为1。
23 23 
24- $$24+ $$
25- out = out\ |\ (scale\ \&\ 0xFFFFE000)\ |\ (1\ll46)25+ out = out\ |\ (scale\ \&\ 0xFFFFE000)\ |\ (1\ll46)
26- $$26+ $$
27+ 
27 3. 根据`offset`取值进行后续计算:28 3. 根据`offset`取值进行后续计算:
28 -`offset`不存在,不再进行后续计算。29 -`offset`不存在,不再进行后续计算。
29 -`offset`存在:30 -`offset`存在: