已合并
modified md files(for readability improvement) #5496
gitee-duhuiping创建于 5月30日
modified md files(for readability improvement) #5496
已合并
共 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 | -## aclnnForeachDivScalarV2 | 170 | +## 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 + DataType | 57 | 输入 Tensor Shape + DataType |
| 58 | ↓ | 58 | ↓ |
| 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 被"吃掉"了) |
| 157 | 1.0e10 + (1.0 - 1.0e10) → 0.0 (同样丢失) | 158 | 1.0e10 + (1.0 - 1.0e10) → 0.0 (同样丢失) |
| 158 | 1.0 + (1.0e10 - 1.0e10) → 1.0 (正确) | 159 | 1.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 | ||
| 178 | Layer Normalization 是深度学习中的基础算子,其核心计算为: | 180 | Layer Normalization 是深度学习中的基础算子,其核心计算为: |
| 179 | 181 | ||
| 180 | -``` | 182 | +```Cpp |
| 181 | mean = E[x] // 均值 | 183 | mean = E[x] // 均值 |
| 182 | rstd = 1 / sqrt(Var[x] + eps) // 标准差的倒数 (reciprocal standard deviation) | 184 | rstd = 1 / sqrt(Var[x] + eps) // 标准差的倒数 (reciprocal standard deviation) |
| 183 | y = gamma * ((x - mean) * rstd) + beta // 归一化 + 仿射变换 | 185 | y = 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 n | 235 | σ² = 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}) / n | 245 | μ_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 |
| 253 | M_{2,n} = M_{2,n-1} + (x_n - μ_{n-1}) * (x_n - μ_n) | 255 | M_{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 |
| 259 | M_{2,n} = M_{2,n-1} + delta * (x_n - μ_n) | 261 | M_{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} / n | 267 | σ_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 |
| 273 | delta = x - mean // mean 为旧均值 μ_{n-1} | 275 | delta = x - mean // mean 为旧均值 μ_{n-1} |
| 274 | mean += delta / n // 更新为新均值 μ_n | 276 | mean += delta / n // 更新为新均值 μ_n |
| 275 | M2 += delta * (x - mean_new) // mean_new 即更新后的 μ_n | 277 | M2 += 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 |
| 284 | M2 = Sum[(x_i - μ_n)²] = n * σ² for i = 1 to n | 286 | M2 = 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 |
| 314 | for i in range(VL_F32): // VL_F32 = 64, SIMD 并行执行 | 316 | for i in range(VL_F32): // VL_F32 = 64, SIMD 并行执行 |
| 315 | delta = data[i] - mean_acc | 317 | delta = data[i] - mean_acc |
| 316 | mean_acc += delta / current_n | 318 | 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 |
| 327 | tileLength = 最大的满足 UB 空间约束的 64 对齐值 | 329 | tileLength = 最大的满足 UB 空间约束的 64 对齐值 |
| 328 | welfordUpdateTimes = N / tileLength | 330 | welfordUpdateTimes = N / tileLength |
| 329 | welfordUpdateTail = N % tileLength | 331 | welfordUpdateTail = 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 的幂 |
| 372 | Step 1: 在 binaryAddQuotient 处拆分 | 374 | Step 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 | ||
| 394 | Level 0 (叶子): a0 a1 a2 a3 a4 a5 a6 a7 | 397 | Level 0 (叶子): a0 a1 a2 a3 a4 a5 a6 a7 |
| 395 | \ / \ / \ / \ / | 398 | \ / \ / \ / \ / |
| 396 | Level 1: s0+s1 s2+s3 s4+s5 s6+s7 ← 两两数量级相近 | 399 | Level 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 | ||
| 445 | sum = ((...(a0 + a1) + a2) + ...) + a_n | 449 | sum = ((...(a0 + a1) + a2) + ...) + a_n |
| 446 | ``` | 450 | ``` |
| 451 | + | ||
| 447 | 误差随 N 线性增长。每步都可能发生大数吃小数。 | 452 | 误差随 N 线性增长。每步都可能发生大数吃小数。 |
| 448 | 453 | ||
| 449 | **二分累加(精度优):** | 454 | **二分累加(精度优):** |
| 450 | -``` | 455 | + |
| 456 | +```Cpp | ||
| 451 | level 0: a0+a1, a2+a3, a4+a5, ... // 子块内部累加 | 457 | level 0: a0+a1, a2+a3, a4+a5, ... // 子块内部累加 |
| 452 | level 1: (a0+a1)+(a2+a3), ... // 子块之间归并 | 458 | level 1: (a0+a1)+(a2+a3), ... // 子块之间归并 |
| 453 | ... | 459 | ... |
| 454 | level k: 最终总和 | 460 | level 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 | ||
| 475 | dst = src * gamma + beta | 483 | dst = 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 |
| 527 | Welford 的 IsCapable 条件: isRegBase == true ← 最宽松 | 537 | Welford 的 IsCapable 条件: isRegBase == true ← 最宽松 |
| 528 | TwoPass 的 IsCapable 条件: isRegBase + UB约束 ← 中等 | 538 | TwoPass 的 IsCapable 条件: isRegBase + UB约束 ← 中等 |
| 529 | TwoPassPerf 的 IsCapable 条件: isRegBase + UB约束 + R≤8192 + 非超大A/R ← 最严格 | 539 | TwoPassPerf 的 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 根据矩阵大小动态切换不同 kernel | 560 | - **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 * weight | 21 | + 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`存在: |