已合并
delete quant_all_reduce & quant_reduce_scatter aclnn #12003
yifuxiong创建于 22 天前
delete quant_all_reduce & quant_reduce_scatter aclnn #12003
已合并
yifuxiong创建于 22 天前
从已删除 :pr_920_beta2合入到cann/ops-transformer9.2.0-beta.2
共 12 个文件变更+9-3181
@@ -200,7 +200,6 @@
200- [aclnnPromptFlashAttentionV3](../../attention/prompt_flash_attention/docs/aclnnPromptFlashAttentionV3.md)200- [aclnnPromptFlashAttentionV3](../../attention/prompt_flash_attention/docs/aclnnPromptFlashAttentionV3.md)
201- [aclnnQkvRmsNormRopeCache](../../posembedding/qkv_rms_norm_rope_cache/docs/aclnnQkvRmsNormRopeCache.md)201- [aclnnQkvRmsNormRopeCache](../../posembedding/qkv_rms_norm_rope_cache/docs/aclnnQkvRmsNormRopeCache.md)
202- [aclnnQkvRmsNormRopeCacheWithKScale](../../posembedding/qkv_rms_norm_rope_cache_with_k_scale/docs/aclnnQkvRmsNormRopeCacheWithKScale.md)202- [aclnnQkvRmsNormRopeCacheWithKScale](../../posembedding/qkv_rms_norm_rope_cache_with_k_scale/docs/aclnnQkvRmsNormRopeCacheWithKScale.md)
203-- [aclnnQuantAllReduce](../../mc2/quant_all_reduce/docs/aclnnQuantAllReduce.md)
204- [aclnnQuantCompressor](../../attention/quant_compressor/docs/aclnnQuantCompressor.md)203- [aclnnQuantCompressor](../../attention/quant_compressor/docs/aclnnQuantCompressor.md)
205- [aclnnQuantFlashAttentionScore](../../attention/flash_attention_score/docs/aclnnQuantFlashAttentionScore.md)204- [aclnnQuantFlashAttentionScore](../../attention/flash_attention_score/docs/aclnnQuantFlashAttentionScore.md)
206- [aclnnQuantGroupedMatMulAlltoAllv](../../mc2/quant_grouped_mat_mul_allto_allv/docs/aclnnQuantGroupedMatMulAlltoAllv.md)205- [aclnnQuantGroupedMatMulAlltoAllv](../../mc2/quant_grouped_mat_mul_allto_allv/docs/aclnnQuantGroupedMatMulAlltoAllv.md)
@@ -219,7 +218,6 @@
219- [aclnnQuantMatmulAllReduceV5](../../mc2/matmul_all_reduce/docs/aclnnQuantMatmulAllReduceV5.md)218- [aclnnQuantMatmulAllReduceV5](../../mc2/matmul_all_reduce/docs/aclnnQuantMatmulAllReduceV5.md)
220- [aclnnQuantMatmulAlltoAll](../../mc2/matmul_allto_all/docs/aclnnQuantMatmulAlltoAll.md)219- [aclnnQuantMatmulAlltoAll](../../mc2/matmul_allto_all/docs/aclnnQuantMatmulAlltoAll.md)
221- [aclnnQuantMatmulAlltoAllV2](../../mc2/matmul_allto_all/docs/aclnnQuantMatmulAlltoAllV2.md)220- [aclnnQuantMatmulAlltoAllV2](../../mc2/matmul_allto_all/docs/aclnnQuantMatmulAlltoAllV2.md)
222-- [aclnnQuantReduceScatter](../../mc2/quant_reduce_scatter/docs/aclnnQuantReduceScatter.md)
223- [aclnnRainFusionAttention](../../attention/rain_fusion_attention/docs/aclnnRainFusionAttention.md)221- [aclnnRainFusionAttention](../../attention/rain_fusion_attention/docs/aclnnRainFusionAttention.md)
224- [aclnnRecurrentGatedDeltaRule](../../attention/recurrent_gated_delta_rule/docs/aclnnRecurrentGatedDeltaRule.md)222- [aclnnRecurrentGatedDeltaRule](../../attention/recurrent_gated_delta_rule/docs/aclnnRecurrentGatedDeltaRule.md)
225- [aclnnRingAttentionUpdate](../../attention/ring_attention_update/docs/aclnnRingAttentionUpdate.md)223- [aclnnRingAttentionUpdate](../../attention/ring_attention_update/docs/aclnnRingAttentionUpdate.md)
@@ -209,7 +209,6 @@
209|[aclnnPromptFlashAttentionV3](../../attention/prompt_flash_attention/docs/aclnnPromptFlashAttentionV3.md)|全量推理场景的FlashAttention算子。|默认确定性实现| - |209|[aclnnPromptFlashAttentionV3](../../attention/prompt_flash_attention/docs/aclnnPromptFlashAttentionV3.md)|全量推理场景的FlashAttention算子。|默认确定性实现| - |
210|[aclnnQkvRmsNormRopeCache](../../posembedding/qkv_rms_norm_rope_cache/docs/aclnnQkvRmsNormRopeCache.md)|输入qkv融合张量,通过SplitVD拆分q、k、v张量,执行RmsNorm、ApplyRotaryPosEmb、Quant、Scatter融合操作,输出qOut、kCache、vCache、qBeforeQuant(可选)、kBeforeQuant(可选)、vBeforeQuant(可选)。|默认确定性实现| - |210|[aclnnQkvRmsNormRopeCache](../../posembedding/qkv_rms_norm_rope_cache/docs/aclnnQkvRmsNormRopeCache.md)|输入qkv融合张量,通过SplitVD拆分q、k、v张量,执行RmsNorm、ApplyRotaryPosEmb、Quant、Scatter融合操作,输出qOut、kCache、vCache、qBeforeQuant(可选)、kBeforeQuant(可选)、vBeforeQuant(可选)。|默认确定性实现| - |
211|[aclnnQkvRmsNormRopeCacheWithKScale](../../posembedding/qkv_rms_norm_rope_cache_with_k_scale/docs/aclnnQkvRmsNormRopeCacheWithKScale.md)|输入Q/K/V融合张量,拆分Q、K、V后对Q/K执行RMSNorm、RoPE、共享rotation矩阵乘和FP8动态量化,输出qOut和qScale;K/V分支按slotMapping更新kCacheRef、kScaleCacheRef和vCacheRef。|-|默认确定性实现|211|[aclnnQkvRmsNormRopeCacheWithKScale](../../posembedding/qkv_rms_norm_rope_cache_with_k_scale/docs/aclnnQkvRmsNormRopeCacheWithKScale.md)|输入Q/K/V融合张量,拆分Q、K、V后对Q/K执行RMSNorm、RoPE、共享rotation矩阵乘和FP8动态量化,输出qOut和qScale;K/V分支按slotMapping更新kCacheRef、kScaleCacheRef和vCacheRef。|-|默认确定性实现|
212-|[aclnnQuantAllReduce](../../mc2/quant_all_reduce/docs/aclnnQuantAllReduce.md)|实现quant + allReduce融合计算。|- | 默认非确定性说明,支持配置开启 |
213|[aclnnQuantCompressor](../../attention/quant_compressor/docs/aclnnQuantCompressor.md)|Compressor的量化版本,将每4或128个token的KV cache压缩成一个,然后每个token与这些压缩的KV cache进行DSA计算。|- | 默认确定性实现 |212|[aclnnQuantCompressor](../../attention/quant_compressor/docs/aclnnQuantCompressor.md)|Compressor的量化版本,将每4或128个token的KV cache压缩成一个,然后每个token与这些压缩的KV cache进行DSA计算。|- | 默认确定性实现 |
214|[aclnnQuantFlashAttentionScore](../../attention/flash_attention_score/docs/aclnnQuantFlashAttentionScore.md)| 量化的训练场景下,使用FlashAttention算法实现self-attention(自注意力)的计算。|- | 默认确定性说明 |213|[aclnnQuantFlashAttentionScore](../../attention/flash_attention_score/docs/aclnnQuantFlashAttentionScore.md)| 量化的训练场景下,使用FlashAttention算法实现self-attention(自注意力)的计算。|- | 默认确定性说明 |
215|[aclnnQuantGroupedMatmulDequantWeightNZ](../../gmm/quant_grouped_matmul_dequant/docs/aclnnQuantGroupedMatmulDequantWeightNZ.md)|对输入x进行量化,分组矩阵乘以及反量化,输入权重Weight会被强制视为NZ格式。| - | - |214|[aclnnQuantGroupedMatmulDequantWeightNZ](../../gmm/quant_grouped_matmul_dequant/docs/aclnnQuantGroupedMatmulDequantWeightNZ.md)|对输入x进行量化,分组矩阵乘以及反量化,输入权重Weight会被强制视为NZ格式。| - | - |
@@ -224,7 +223,6 @@
224|[aclnnQuantMatmulAlltoAll](../../mc2/matmul_allto_all/docs/aclnnQuantMatmulAlltoAll.md)|对量化后的入参x1、x2进行MatMul计算后,接着进行Dequant计算,最后做AlltoAll通信。|默认确定性实现| 默认确定性实现 |223|[aclnnQuantMatmulAlltoAll](../../mc2/matmul_allto_all/docs/aclnnQuantMatmulAlltoAll.md)|对量化后的入参x1、x2进行MatMul计算后,接着进行Dequant计算,最后做AlltoAll通信。|默认确定性实现| 默认确定性实现 |
225|[aclnnQuantMatmulAlltoAllV2](../../mc2/matmul_allto_all/docs/aclnnQuantMatmulAlltoAllV2.md)|兼容[aclnnQuantMatmulAlltoAll](../../mc2/matmul_allto_all/docs/aclnnQuantMatmulAlltoAll.md)支持的功能,在此基础上新增commMode参数,供用户指定通信引擎参数。|默认确定性实现| 默认确定性实现 |224|[aclnnQuantMatmulAlltoAllV2](../../mc2/matmul_allto_all/docs/aclnnQuantMatmulAlltoAllV2.md)|兼容[aclnnQuantMatmulAlltoAll](../../mc2/matmul_allto_all/docs/aclnnQuantMatmulAlltoAll.md)支持的功能,在此基础上新增commMode参数,供用户指定通信引擎参数。|默认确定性实现| 默认确定性实现 |
226|[aclnnQuantGroupedMatmulDequant](../../gmm/quant_grouped_matmul_dequant/docs/aclnnQuantGroupedMatmulDequant.md)|对输入x进行量化,分组矩阵乘以及反量化。|默认确定性实现| 默认确定性实现 |225|[aclnnQuantGroupedMatmulDequant](../../gmm/quant_grouped_matmul_dequant/docs/aclnnQuantGroupedMatmulDequant.md)|对输入x进行量化,分组矩阵乘以及反量化。|默认确定性实现| 默认确定性实现 |
227-|[aclnnQuantReduceScatter](../../mc2/quant_reduce_scatter/docs/aclnnQuantReduceScatter.md)|实现quant + reduceScatter融合计算。|默认确定性实现| 默认确定性实现 |
228|[aclnnRainFusionAttention](../../attention/rain_fusion_attention/docs/aclnnRainFusionAttention.md)|RainFusionAttention稀疏注意力计算,支持灵活的块级稀疏模式,通过selectIdx指定每个Q块选择的KV块,实现高效的稀疏注意力计算。|默认确定性实现| - |226|[aclnnRainFusionAttention](../../attention/rain_fusion_attention/docs/aclnnRainFusionAttention.md)|RainFusionAttention稀疏注意力计算,支持灵活的块级稀疏模式,通过selectIdx指定每个Q块选择的KV块,实现高效的稀疏注意力计算。|默认确定性实现| - |
229|[aclnnRecurrentGatedDeltaRule](../../attention/recurrent_gated_delta_rule/docs/aclnnRecurrentGatedDeltaRule.md)|完成变步长的Recurrent Gated Delta Rule计算。|默认确定性实现| 默认确定性实现 |227|[aclnnRecurrentGatedDeltaRule](../../attention/recurrent_gated_delta_rule/docs/aclnnRecurrentGatedDeltaRule.md)|完成变步长的Recurrent Gated Delta Rule计算。|默认确定性实现| 默认确定性实现 |
230|[aclnnRingAttentionUpdate](../../attention/ring_attention_update/docs/aclnnRingAttentionUpdate.md)|将两次FlashAttention的输出根据其不同的softmax的max和sum更新。|默认确定性实现| 默认确定性实现 |228|[aclnnRingAttentionUpdate](../../attention/ring_attention_update/docs/aclnnRingAttentionUpdate.md)|将两次FlashAttention的输出根据其不同的softmax的max和sum更新。|默认确定性实现| 默认确定性实现 |
@@ -1,128 +1,3 @@
1-# QuantAllReduce1+# aclnnQuantAllReduce
2 2 
3-## 产品支持情况3+该算子暂无Ascend C代码实现,欢迎开发者补充贡献。
4- 
5-| 产品 | 是否支持 |
6-| :----------------------------------------------------------- | :------: |
7-| <term>Ascend 950DT</term> | √ |
8-| <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | × |
9-| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | × |
10-| <term>Atlas 200I/500 A2 推理产品</term> | × |
11-| <term>Atlas 推理系列产品</term> | × |
12-| <term>Atlas 训练系列产品</term> | × |
13- 
14-## 功能说明
15- 
16-算子功能:实现低比特数据的AllReduce通信,在通信的过程中对数据进行反量化,并输出通信结果。
17- 
18-- **计算公式**:
19- 
20- $$
21- AllGatherData = AllGather(x)
22- $$
23- 
24- $$
25- AllGatherScales = AllGather(scales)
26- $$
27- 
28- $$
29- output = Reduce(AllGatherScales * AllGatherData)
30- $$
31- 
32- 其中的Reduce计算是将来自不同rank的数据进行reduce计算。
33- 
34-## 参数说明
35- 
36-<table style="undefined;table-layout: fixed; width: 1567px"><colgroup>
37- <col style="width: 170px">
38- <col style="width: 120px">
39- <col style="width: 300px">
40- <col style="width: 330px">
41- <col style="width: 212px">
42- <col style="width: 100px">
43- <col style="width: 190px">
44- <col style="width: 145px">
45- </colgroup>
46- <thead>
47- <tr>
48- <th>参数名</th>
49- <th>输入/输出/属性</th>
50- <th>描述</th>
51- <th>使用说明</th>
52- <th>数据类型</th>
53- <th>数据格式</th>
54- <th>维度(shape)</th>
55- <th>连续Tensor</th>
56- </tr></thead>
57- <tbody>
58- <tr>
59- <td>x</td>
60- <td>输入</td>
61- <td>公式中的输入x。</td>
62- <td><li>不支持空Tensor。</li><li>支持的shape为:(bs, H)或者(b, s, H)。b为batch size,s为sequence length,H为hidden size。</li></td>
63- <td>INT8、HIFLOAT8、FLOAT8_E4M3FN、FLOAT8_E5M2</td>
64- <td>ND</td>
65- <td>2-3</td>
66- <td>√</td>
67- </tr>
68- <tr>
69- <td>scales</td>
70- <td>输入</td>
71- <td>公式中的输入scales。</td>
72- <td><li>不支持空Tensor。</li><li>当scales的数据类型为FLOAT8_E8M0时,x的数据类型必须为FLOAT8_E4M3FN、FLOAT8_E5M2,x的shape为(bs, H)或者(b, s, H),scales的shape必须对应x的shape为(bs, H/64, 2)或者(b, s, H/64, 2)。</li><li>当scales的数据类型为FLOAT时,x的数据类型必须为INT8、HIFLOAT8、FLOAT8_E4M3FN、FLOAT8_E5M2,x的shape为(bs, H)或者(b, s, H),scales的shape必须对应x的shape为(bs, H/128)或者(b, s, H/128)。</li></td>
73- <td>FLOAT、FLOAT8_E8M0</td>
74- <td>ND</td>
75- <td>2-4</td>
76- <td>√</td>
77- </tr>
78- <tr>
79- <td>group</td>
80- <td>属性</td>
81- <td>通信域标识。</td>
82- <td><li>Host侧标识列组的字符串,通信域名称。</li><li>通过Hccl提供的接口"extern HcclResult HcclGetCommName(HcclComm comm, char* commName);"获取,其中commName即为group。</li></td>
83- <td>Char*、String</td>
84- <td>-</td>
85- <td>-</td>
86- <td>-</td>
87- </tr>
88- <tr>
89- <td>reduceOp</td>
90- <td>可选属性</td>
91- <td>公式中的reduce操作类型。</td>
92- <td>当前仅支持"sum"操作。</td>
93- <td>Char*、String</td>
94- <td>-</td>
95- <td>-</td>
96- <td>-</td>
97- </tr>
98- <tr>
99- <td>output</td>
100- <td>输出</td>
101- <td>公式中的输出output。</td>
102- <td><li>不支持空Tensor。</li><li>支持的shape为(bs, H)或者(b, s, H),output的shape与x保持一致。</li></td>
103- <td>FLOAT、FLOAT16、BFLOAT16</td>
104- <td>ND</td>
105- <td>2-3</td>
106- <td>√</td>
107- </tr>
108- </tbody>
109-</table>
110- 
111-## 约束说明
112- 
113-- 当x的数据类型为FLOAT8_E4M3FN、FLOAT8_E5M2并且scales的数据类型为FLOAT8_E8M0时,输入数据的量化方式为mx量化。
114-- 当x的数据类型为INT8、HIFLOAT8、FLOAT8_E4M3FN、FLOAT8_E5M2并且scales的数据类型为FLOAT时,输入数据的量化方式为pertoken-pergroup量化(groupSize=128)。
115-- 只在Ascend950系列平台开启。
116-- 不支持空Tensor输入。
117-- 通信引擎约束:
118- - Ascend950DT: 仅支持UB-Memory通信。
119-- 通信域大小支持2, 4, 8。
120-- 通信域使用约束:同一通信域内仅允许连续执行`aclnnQuantAllReduce`和`aclnnQuantReduceScatter`算且子,该通信域中不允许有其他通信算子。
121-- `HCCL_BUFFSIZE`:调用本算子前需检查`HCCL_BUFFSIZE`环境变量取值是否合理,该环境变量表示单个通信域占用内存大小,单位MB,不配置时默认为200MB。要求满足`HCCL_BUFFSIZE`>= 2 * (`xDataSize` + `scalesDataSize + 1`)。其中`xDataSize`为输入`x`的数据大小,计算公式为:`xDataSize = b * s * H * 1 (Byte)`,`scalesDataSize`为`scales`的数据大小,当量化方式为pertoken-pergroup量化时,计算公式为:`scalesDataSize = b * s * H / 128 * 4 (Byte)`,当量化方式为mx量化时,计算公式为:`scalesDataSize = b * s * H / 32 * 1 (Byte)`。
122-- H范围仅支持[1024, 8192],要求128对齐。
123- 
124-## 调用说明
125- 
126-| 调用方式 | 样例代码 | 说明 |
127-| :--------: | :----------------------------------------: | :-------------------------------------------------------: |
128-| aclnn接口 | [test_aclnn_quant_all_reduce.cpp](./examples/test_aclnn_quant_all_reduce.cpp) | 通过[aclnnQuantAllReduce](./docs/aclnnQuantAllReduce.md)接口方式调用quant_all_reduce算子。 |
@@ -1,469 +0,0 @@
1-# aclnnQuantAllReduce
2- 
3-## 产品支持情况
4- 
5-<!-- npu="950" id1 -->
6-- <term>Ascend 950DT</term>:支持
7-<!-- end id1 -->
8-<!-- npu="A3" id2 -->
9-- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:不支持
10-<!-- end id2 -->
11-<!-- npu="910b" id3 -->
12-- <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持
13-<!-- end id3 -->
14-<!-- npu="310b" id4 -->
15-- <term>Atlas 200I/500 A2 推理产品</term>:不支持
16-<!-- end id4 -->
17-<!-- npu="310p" id5 -->
18-- <term>Atlas 推理系列产品</term>:不支持
19-<!-- end id5 -->
20-<!-- npu="910" id6 -->
21-- <term>Atlas 训练系列产品</term>:不支持
22-<!-- end id6 -->
23- 
24-**说明:** 使用该接口时,请确保驱动固件包和CANN包都为配套的8.0.RC2版本或者配套的更高版本,否则将会引发报错,比如BUS ERROR等。
25- 
26-## 功能说明
27- 
28-- **接口功能**:实现低比特数据的AllReduce通信,在通信的过程中对数据进行反量化,并输出通信结果。
29-具体实现依据数据量大小有两种情况:
30- 
31-- **计算公式**:
32- 
33- $$
34- AllGatherData = AllGather(x)
35- $$
36- 
37- $$
38- AllGatherScales = AllGather(scales)
39- $$
40- 
41- $$
42- output = Reduce(AllGatherScales * AllGatherData)
43- $$
44- 
45- 其中的Reduce计算是将来自不同rank的数据进行reduce计算。
46- 
47-## 函数原型
48- 
49-该算子分为两段式接口,必须先调用“aclnnQuantAllReduceGetWorkspaceSize”接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用“aclnnQuantAllReduce”接口执行计算。
50- 
51-```cpp
52-aclnnStatus aclnnQuantAllReduceGetWorkspaceSize(
53- const aclTensor* x,
54- const aclTensor* scales,
55- const char* group,
56- const char* reduceOp,
57- aclTensor* output,
58- uint64_t* workspaceSize,
59- aclOpExecutor** executor)
60-```
61- 
62-```cpp
63-aclnnStatus aclnnQuantAllReduce(
64- void* workspace,
65- uint64_t workspaceSize,
66- aclOpExecutor* executor,
67- const aclrtStream stream)
68-```
69- 
70-## aclnnQuantAllReduceGetWorkspaceSize
71- 
72-- **参数说明**
73- 
74- <table style="undefined;table-layout: fixed; width: 1556px"><colgroup>
75- <col style="width: 161px">
76- <col style="width: 141px">
77- <col style="width: 245px">
78- <col style="width: 408px">
79- <col style="width: 191px">
80- <col style="width: 120px">
81- <col style="width: 145px">
82- <col style="width: 145px">
83- </colgroup>
84- <thead>
85- <tr>
86- <th>参数名</th>
87- <th>输入/输出/属性</th>
88- <th>描述</th>
89- <th>使用说明</th>
90- <th>数据类型</th>
91- <th>数据格式</th>
92- <th>维度(shape)</th>
93- <th>连续Tensor</th>
94- </tr></thead>
95- <tbody>
96- <tr>
97- <td>x</td>
98- <td>输入</td>
99- <td>公式中的输入x。</td>
100- <td><ul><li>不支持空Tensor。</li><li>支持的shape为:(bs, H)或者(b, s, H)。b为batch size,s为sequence length,H为hidden size。</li></ul></td>
101- <td>INT8、HIFLOAT8、FLOAT8_E4M3FN、FLOAT8_E5M2</td>
102- <td>ND</td>
103- <td>2-3</td>
104- <td>√</td>
105- </tr>
106- <tr>
107- <td>scales</td>
108- <td>输入</td>
109- <td>公式中的输入scales。</td>
110- <td><ul><li>不支持空Tensor。</li><li>当scales的数据类型为FLOAT8_E8M0时,x的数据类型必须为FLOAT8_E4M3FN、FLOAT8_E5M2,x的shape为(bs, H)或者(b, s, H),scales的shape必须对应x的shape为(bs, H/64, 2)或者(b, s, H/64, 2)。</li><li>当scales的数据类型为FLOAT时,x的数据类型必须为INT8、HIFLOAT8、FLOAT8_E4M3FN、FLOAT8_E5M2,x的shape为(bs, H)或者(b, s, H),scales的shape必须对应x的shape为(bs, H/128)或者(b, s, H/128)。</li></ul></td>
111- <td>FLOAT、FLOAT8_E8M0</td>
112- <td>ND</td>
113- <td>2-4</td>
114- <td>√</td>
115- </tr>
116- <tr>
117- <td>group</td>
118- <td>属性</td>
119- <td>通信域标识。</td>
120- <td><ul><li>Host侧标识列组的字符串,通信域名称。</li><li>通过Hccl提供的接口"extern HcclResult HcclGetCommName(HcclComm comm, char* commName);"获取,其中commName即为group。</li></ul></td>
121- <td>Char*、String</td>
122- <td>-</td>
123- <td>-</td>
124- <td>-</td>
125- </tr>
126- <tr>
127- <td>reduceOp</td>
128- <td>可选属性</td>
129- <td>公式中的reduce操作类型。</td>
130- <td>当前仅支持"sum"操作。</td>
131- <td>Char*、String</td>
132- <td>-</td>
133- <td>-</td>
134- <td>-</td>
135- </tr>
136- <tr>
137- <td>output</td>
138- <td>输出</td>
139- <td>公式中的输出output。</td>
140- <td><ul><li>不支持空Tensor。</li><li>支持的shape为(bs, H)或者(b, s, H),output的shape与x保持一致。</li></ul></td>
141- <td>FLOAT、FLOAT16、BFLOAT16</td>
142- <td>ND</td>
143- <td>2-3</td>
144- <td>√</td>
145- </tr>
146- <tr>
147- <td>workspaceSize</td>
148- <td>输出</td>
149- <td>返回需要在device侧申请的workspace大小。</td>
150- <td>-</td>
151- <td>-</td>
152- <td>-</td>
153- <td>-</td>
154- <td>-</td>
155- </tr>
156- <tr>
157- <td>executor</td>
158- <td>输出</td>
159- <td>返回op执行器,包含了算子计算流程。</td>
160- <td>-</td>
161- <td>-</td>
162- <td>-</td>
163- <td>-</td>
164- <td>-</td>
165- </tr>
166- </tbody>
167- </table>
168- 
169-- **返回值**
170- 
171- aclnnStatus:返回状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。
172- 
173- 第一段接口完成入参校验,出现以下场景时报错:
174- 
175- <table style="undefined;table-layout: fixed; width: 1149px"><colgroup>
176- <col style="width: 282px">
177- <col style="width: 120px">
178- <col style="width: 747px">
179- </colgroup>
180- <thead>
181- <tr>
182- <th>返回值</th>
183- <th>错误码</th>
184- <th>描述</th>
185- </tr></thead>
186- <tbody>
187- <tr>
188- <td>ACLNN_ERR_PARAM_NULLPTR</td>
189- <td>161001</td>
190- <td>x、scales、output存在空指针。</td>
191- </tr>
192- <tr>
193- <td rowspan="3">ACLNN_ERR_PARAM_INVALID</td>
194- <td rowspan="3">161002</td>
195- <td>x、scales或output的数据类型不在支持的范围之内。</td>
196- </tr>
197- <tr>
198- <td>x、scales的数据类型和shape不匹配。</td>
199- </tr>
200- <tr>
201- <td>x、scales的维度不在支持的范围之内。</td>
202- </tr>
203- </tbody>
204- </table>
205- 
206-## aclnnQuantAllReduce
207- 
208-- **参数说明**
209- 
210- <table style="undefined;table-layout: fixed; width: 1150px"><colgroup>
211- <col style="width: 168px">
212- <col style="width: 128px">
213- <col style="width: 854px">
214- </colgroup>
215- <thead>
216- <tr>
217- <th>参数名</th>
218- <th>输入/输出</th>
219- <th>描述</th>
220- </tr></thead>
221- <tbody>
222- <tr>
223- <td>workspace</td>
224- <td>输入</td>
225- <td>在Device侧申请的workspace内存地址。</td>
226- </tr>
227- <tr>
228- <td>workspaceSize</td>
229- <td>输入</td>
230- <td>在Device侧申请的workspace大小,由第一段接口aclnnQuantAllReduceGetWorkspaceSize获取。</td>
231- </tr>
232- <tr>
233- <td>executor</td>
234- <td>输入</td>
235- <td>op执行器,包含了算子计算流程。</td>
236- </tr>
237- <tr>
238- <td>stream</td>
239- <td>输入</td>
240- <td>指定执行任务的Stream。</td>
241- </tr>
242- </tbody></table>
243- 
244-- **返回值**
245- 
246- 返回aclnnStatus状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。
247- 
248-## 约束说明
249- 
250-- 确定性计算:
251- - aclnnQuantAllReduce默认非确定性实现,支持通过aclrtCtxSetSysParamOpt开启确定性。
252- 
253-- 当x的数据类型为FLOAT8_E4M3FN、FLOAT8_E5M2并且scales的数据类型为FLOAT8_E8M0时,输入数据的量化方式为mx量化。
254-- 当x的数据类型为INT8、HIFLOAT8、FLOAT8_E4M3FN、FLOAT8_E5M2并且scales的数据类型为FLOAT时,输入数据的量化方式为pertoken-pergroup量化(groupSize=128)。
255-- 只在Ascend950系列平台开启。
256-- 不支持空Tensor输入。
257-- 通信引擎约束:
258- - Ascend950DT: 仅支持UB-Memory通信。
259-- 通信域大小支持2, 4, 8。
260-- 通信域使用约束:同一通信域内仅允许连续执行`aclnnQuantAllReduce`和`aclnnQuantReduceScatter`算子,且该通信域中不允许有其他通信算子。
261-- `HCCL_BUFFSIZE`:调用本算子前需检查`HCCL_BUFFSIZE`环境变量取值是否合理,该环境变量表示单个通信域占用内存大小,单位MB,不配置时默认为200MB。要求满足`HCCL_BUFFSIZE`>= 2 * (`xDataSize` + `scalesDataSize + 1`)。其中`xDataSize`为输入`x`的数据大小,计算公式为:`xDataSize = b * s * H * 1 (Byte)`,`scalesDataSize`为`scales`的数据大小,当量化方式为pertoken-pergroup量化时,计算公式为:`scalesDataSize = b * s * H / 128 * 4 (Byte)`,当量化方式为mx量化时,计算公式为:`scalesDataSize = b * s * H / 32 * 1 (Byte)`。
262-- H范围仅支持[1024, 8192],要求128对齐。
263- 
264-## 调用示例
265- 
266-示例代码如下,仅供参考,具体编译和执行过程请参考编译与运行样例。
267- 
268-说明:本示例代码调用了部分HCCL集合通信库接口:HcclCommInitClusterInfoConfig、HcclGetCommName、HcclCommDestroy,请参考[<<HCCL API (C)>>](https://www.hiascend.com/document/detail/zh/CANNCommunityEdition/850alpha002/API/hcclapiref/hcclcpp_07_0001.html)。
269- 
270-<!-- npu="950" id7 -->
271-- <term>Ascend 950DT</term>:
272- 
273- ```Cpp
274- #include <thread>
275- #include <iostream>
276- #include <vector>
277- #include <string>
278- #include <cstring>
279- #include "hccl/hccl.h"
280- #include "aclnnop/aclnn_quant_all_reduce.h"
281- using namespace std;
282- 
283- #define CHECK_RET(cond, return_expr) \
284- do { \
285- if (!(cond)) { \
286- return_expr; \
287- } \
288- } while (0)
289- 
290- #define LOG_PRINT(message, ...) \
291- do { \
292- printf(message, ##__VA_ARGS__); \
293- } while (0)
294- 
295- constexpr int DEV_NUM = 2; // 设备数量
296- 
297- int64_t GetShapeSize(const std::vector<int64_t> &shape)
298- {
299- int64_t shape_size = 1;
300- for (auto i : shape) {
301- shape_size *= i;
302- }
303- return shape_size;
304- }
305- 
306- template<typename T>
307- int CreateAclTensor(const std::vector<T> &hostData, const std::vector<int64_t> &shape, void **deviceAddr,
308- aclDataType dataType, aclTensor **tensor)
309- {
310- auto size = GetShapeSize(shape) * sizeof(T);
311- auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST);
312- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtMalloc failed. ret: %d\n", ret);
313- return ret);
314- ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE);
315- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtMemcpy failed. ret: %d\n", ret);
316- return ret);
317- std::vector<int64_t> strides(shape.size(), 1);
318- for (int64_t i = shape.size() - 2; i >= 0; i--) {
319- strides[i] = shape[i + 1] * strides[i + 1];
320- }
321- *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND,
322- shape.data(), shape.size(), *deviceAddr);
323- return 0;
324- }
325- 
326- struct Args {
327- uint32_t rankId;
328- HcclComm hcclComm;
329- aclrtStream stream;
330- aclrtContext context;
331- };
332- 
333- int LaunchOneThreadQuantAllReduce(Args &args)
334- {
335- int ret = aclrtSetCurrentContext(args.context);
336- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtSetCurrentContext failed. ret = %d\n", ret);
337- return ret);
338- char hcomName[128] = {0};
339- ret = HcclGetCommName(args.hcclComm, hcomName);
340- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] HcclGetCommName failed. ret = %d\n", ret);
341- return -1);
342- LOG_PRINT("[INFO] rank = %d, hcomName = %s, stream = %p\n", args.rankId, hcomName, args.stream);
343- std::vector<int64_t> xShape = {1024, 5120}; // (bs, H)
344- std::vector<int64_t> scalesShape = {1024, 80, 2}; // (bs, H/64, 2)
345- std::vector<int64_t> outputShape = {1024, 5120}; // (bs, H)
346- void *xDeviceAddr = nullptr;
347- void *scalesDeviceAddr = nullptr;
348- void *outputDeviceAddr = nullptr;
349- void *workspaceAddr = nullptr;
350- 
351- aclTensor *x = nullptr;
352- aclTensor *scales = nullptr;
353- aclTensor *output = nullptr;
354- uint64_t workspaceSize = 0;
355- aclOpExecutor *executor = nullptr;
356- 
357- long long xShapeSize = GetShapeSize(xShape);
358- long long scalesShapeSize = GetShapeSize(scalesShape);
359- long long outputShapeSize = GetShapeSize(outputShape);
360- 
361- std::vector<int8_t> xHostData(xShapeSize, 0);
362- std::vector<int8_t> scalesHostData(scalesShapeSize, 0);
363- std::vector<int16_t> outputHostData(outputShapeSize, 0);
364- 
365- // 创建tensor
366- ret = CreateAclTensor(xHostData, xShape, &xDeviceAddr, aclDataType::ACL_FLOAT8_E5M2, &x);
367- CHECK_RET(ret == ACL_SUCCESS, return ret);
368- ret = CreateAclTensor(scalesHostData, scalesShape, &scalesDeviceAddr, aclDataType::ACL_FLOAT8_E8M0, &scales);
369- CHECK_RET(ret == ACL_SUCCESS, return ret);
370- ret = CreateAclTensor(outputHostData, outputShape, &outputDeviceAddr, aclDataType::ACL_FLOAT16, &output);
371- CHECK_RET(ret == ACL_SUCCESS, return ret);
372- 
373- // 调用第一阶段接口
374- ret = aclnnQuantAllReduceGetWorkspaceSize(
375- x, scales, hcomName, "sum", output, &workspaceSize, &executor);
376- CHECK_RET(ret == ACL_SUCCESS,
377- LOG_PRINT("[ERROR] aclnnQuantAllReduceGetWorkspaceSize failed. ret = %d \n", ret);
378- return ret);
379- // 根据第一阶段接口计算出的workspaceSize申请device内存
380- if (workspaceSize > 0) {
381- ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST);
382- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtMalloc workspace failed. ret = %d \n", ret);
383- return ret);
384- }
385- // 调用第二阶段接口
386- ret = aclnnQuantAllReduce(workspaceAddr, workspaceSize, executor, args.stream);
387- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclnnQuantAllReduce failed. ret = %d \n", ret); return ret);
388- //(固定写法)同步等待任务执行结束
389- ret = aclrtSynchronizeStreamWithTimeout(args.stream, 10000);
390- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtSynchronizeStreamWithTimeout failed. ret = %d \n", ret);
391- return ret);
392- LOG_PRINT("[INFO] device_%d aclnnQuantAllReduce execute successfully.\n", args.rankId);
393- // 释放device资源,需要根据具体API的接口定义修改
394- if (x != nullptr) {
395- aclDestroyTensor(x);
396- }
397- if (scales != nullptr) {
398- aclDestroyTensor(scales);
399- }
400- if (output != nullptr) {
401- aclDestroyTensor(output);
402- }
403- 
404- if (xDeviceAddr != nullptr) {
405- aclrtFree(xDeviceAddr);
406- }
407- if (scalesDeviceAddr != nullptr) {
408- aclrtFree(scalesDeviceAddr);
409- }
410- if (outputDeviceAddr != nullptr) {
411- aclrtFree(outputDeviceAddr);
412- }
413- if (workspaceSize > 0) {
414- aclrtFree(workspaceAddr);
415- }
416- ret = HcclCommDestroy(args.hcclComm);
417- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] HcclCommDestroy failed. ret = %d \n", ret); return ret);
418- ret = aclrtDestroyStream(args.stream);
419- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtDestroyStream failed. ret = %d \n", ret); return ret);
420- ret = aclrtResetDevice(args.rankId);
421- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtResetDevice failed. ret = %d \n", ret); return ret);
422- ret = aclrtDestroyContext(args.context);
423- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtDestroyContext failed. ret = %d \n", ret); return ret);
424- return 0;
425- }
426- 
427- int main(int argc, char *argv[])
428- {
429- int ret = aclInit(nullptr);
430- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclInit failed. ret = %d \n", ret);
431- return ret);
432- aclrtStream stream[DEV_NUM];
433- aclrtContext context[DEV_NUM];
434- for (uint32_t rankId = 0; rankId < DEV_NUM; rankId++) {
435- ret = aclrtSetDevice(rankId);
436- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtSetDevice failed. ret = %d \n", ret); return ret);
437- ret = aclrtCreateContext(&context[rankId], rankId);
438- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtCreateContext failed. ret = %d \n", ret); return ret);
439- ret = aclrtCreateStream(&stream[rankId]);
440- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtCreateStream failed. ret = %d \n", ret); return ret);
441- }
442- int32_t devices[DEV_NUM];
443- for (int i = 0; i < DEV_NUM; i++) {
444- devices[i] = i;
445- }
446- // 初始化集合通信域
447- HcclComm comms[DEV_NUM];
448- ret = HcclCommInitAll(DEV_NUM, devices, comms);
449- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] HcclCommInitAll failed. ret = %d \n", ret); return ret);
450- 
451- Args args[DEV_NUM];
452- // 启动多线程
453- std::vector<std::unique_ptr<std::thread>> threads(DEV_NUM);
454- for (uint32_t rankId = 0; rankId < DEV_NUM; rankId++) {
455- args[rankId].rankId = rankId;
456- args[rankId].hcclComm = comms[rankId];
457- args[rankId].context = context[rankId];
458- args[rankId].stream = stream[rankId];
459- threads[rankId].reset(new(std::nothrow) std::thread(&LaunchOneThreadQuantAllReduce, std::ref(args[rankId])));
460- }
461- for (uint32_t rankId = 0; rankId < DEV_NUM; rankId++) {
462- threads[rankId]->join();
463- }
464- aclFinalize();
465- return 0;
466- }
467- ```
468- 
469-<!-- end id7 -->
@@ -1,200 +0,0 @@
1-/**
2- * Copyright (c) 2025 Huawei Technologies Co., Ltd.
3- * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4- * CANN Open Software License Agreement Version 2.0 (the "License").
5- * Please refer to the License for details. You may not use this file except in compliance with the License.
6- * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7- * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8- * See LICENSE in the root of the software repository for the full text of the License.
9- */
10- 
11-/*!
12- * \file test_aclnn_quant_all_reduce.cpp
13- * \brief aclnn测试样例
14- */
15- 
16-#include <thread>
17-#include <iostream>
18-#include <vector>
19-#include <string>
20-#include <cstring>
21-#include "hccl/hccl.h"
22-#include "aclnnop/aclnn_quant_all_reduce.h"
23-using namespace std;
24- 
25-#define CHECK_RET(cond, return_expr) \
26- do { \
27- if (!(cond)) { \
28- return_expr; \
29- } \
30- } while (0)
31- 
32-#define LOG_PRINT(message, ...) \
33- do { \
34- printf(message, ##__VA_ARGS__); \
35- } while (0)
36- 
37-constexpr int DEV_NUM = 2; // 设备数量
38- 
39-int64_t GetShapeSize(const std::vector<int64_t> &shape)
40-{
41- int64_t shape_size = 1;
42- for (auto i : shape) {
43- shape_size *= i;
44- }
45- return shape_size;
46-}
47- 
48-template <typename T>
49-int CreateAclTensor(const std::vector<T> &hostData, const std::vector<int64_t> &shape, void **deviceAddr,
50- aclDataType dataType, aclTensor **tensor)
51-{
52- auto size = GetShapeSize(shape) * sizeof(T);
53- auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST);
54- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtMalloc failed. ret: %d\n", ret); return ret);
55- ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE);
56- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtMemcpy failed. ret: %d\n", ret); return ret);
57- std::vector<int64_t> strides(shape.size(), 1);
58- for (int64_t i = shape.size() - 2; i >= 0; i--) {
59- strides[i] = shape[i + 1] * strides[i + 1];
60- }
61- *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND,
62- shape.data(), shape.size(), *deviceAddr);
63- return 0;
64-}
65- 
66-struct Args {
67- uint32_t rankId;
68- HcclComm hcclComm;
69- aclrtStream stream;
70- aclrtContext context;
71-};
72- 
73-int LaunchOneThreadQuantAllReduce(Args &args)
74-{
75- int ret = aclrtSetCurrentContext(args.context);
76- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtSetCurrentContext failed. ret = %d\n", ret); return ret);
77- char hcomName[128] = {0};
78- ret = HcclGetCommName(args.hcclComm, hcomName);
79- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] HcclGetCommName failed. ret = %d\n", ret); return -1);
80- LOG_PRINT("[INFO] rank = %d, hcomName = %s, stream = %p\n", args.rankId, hcomName, args.stream);
81- std::vector<int64_t> xShape = {1024, 5120}; // (bs, H)
82- std::vector<int64_t> scalesShape = {1024, 80, 2}; // (bs, H/64, 2)
83- std::vector<int64_t> outputShape = {1024, 5120}; // (bs, H)
84- void *xDeviceAddr = nullptr;
85- void *scalesDeviceAddr = nullptr;
86- void *outputDeviceAddr = nullptr;
87- void *workspaceAddr = nullptr;
88- 
89- aclTensor *x = nullptr;
90- aclTensor *scales = nullptr;
91- aclTensor *output = nullptr;
92- uint64_t workspaceSize = 0;
93- aclOpExecutor *executor = nullptr;
94- 
95- long long xShapeSize = GetShapeSize(xShape);
96- long long scalesShapeSize = GetShapeSize(scalesShape);
97- long long outputShapeSize = GetShapeSize(outputShape);
98- 
99- std::vector<int8_t> xHostData(xShapeSize, 0);
100- std::vector<int8_t> scalesHostData(scalesShapeSize, 0);
101- std::vector<int16_t> outputHostData(outputShapeSize, 0);
102- 
103- // 创建tensor
104- ret = CreateAclTensor(xHostData, xShape, &xDeviceAddr, aclDataType::ACL_FLOAT8_E5M2, &x);
105- CHECK_RET(ret == ACL_SUCCESS, return ret);
106- ret = CreateAclTensor(scalesHostData, scalesShape, &scalesDeviceAddr, aclDataType::ACL_FLOAT8_E8M0, &scales);
107- CHECK_RET(ret == ACL_SUCCESS, return ret);
108- ret = CreateAclTensor(outputHostData, outputShape, &outputDeviceAddr, aclDataType::ACL_FLOAT16, &output);
109- CHECK_RET(ret == ACL_SUCCESS, return ret);
110- 
111- // 调用第一阶段接口
112- ret = aclnnQuantAllReduceGetWorkspaceSize(x, scales, hcomName, "sum", output, &workspaceSize, &executor);
113- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclnnQuantAllReduceGetWorkspaceSize failed. ret = %d \n", ret);
114- return ret);
115- // 根据第一阶段接口计算出的workspaceSize申请device内存
116- if (workspaceSize > 0) {
117- ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST);
118- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtMalloc workspace failed. ret = %d \n", ret); return ret);
119- }
120- // 调用第二阶段接口
121- ret = aclnnQuantAllReduce(workspaceAddr, workspaceSize, executor, args.stream);
122- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclnnQuantAllReduce failed. ret = %d \n", ret); return ret);
123- // (固定写法)同步等待任务执行结束
124- ret = aclrtSynchronizeStreamWithTimeout(args.stream, 10000);
125- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtSynchronizeStreamWithTimeout failed. ret = %d \n", ret);
126- return ret);
127- LOG_PRINT("[INFO] device_%d aclnnQuantAllReduce execute successfully.\n", args.rankId);
128- // 释放device资源,需要根据具体API的接口定义修改
129- if (x != nullptr) {
130- aclDestroyTensor(x);
131- }
132- if (scales != nullptr) {
133- aclDestroyTensor(scales);
134- }
135- if (output != nullptr) {
136- aclDestroyTensor(output);
137- }
138- 
139- if (xDeviceAddr != nullptr) {
140- aclrtFree(xDeviceAddr);
141- }
142- if (scalesDeviceAddr != nullptr) {
143- aclrtFree(scalesDeviceAddr);
144- }
145- if (outputDeviceAddr != nullptr) {
146- aclrtFree(outputDeviceAddr);
147- }
148- if (workspaceSize > 0) {
149- aclrtFree(workspaceAddr);
150- }
151- ret = HcclCommDestroy(args.hcclComm);
152- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] HcclCommDestroy failed. ret = %d \n", ret); return ret);
153- ret = aclrtDestroyStream(args.stream);
154- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtDestroyStream failed. ret = %d \n", ret); return ret);
155- ret = aclrtResetDevice(args.rankId);
156- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtResetDevice failed. ret = %d \n", ret); return ret);
157- ret = aclrtDestroyContext(args.context);
158- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtDestroyContext failed. ret = %d \n", ret); return ret);
159- return 0;
160-}
161- 
162-int main(int argc, char *argv[])
163-{
164- int ret = aclInit(nullptr);
165- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclInit failed. ret = %d \n", ret); return ret);
166- aclrtStream stream[DEV_NUM];
167- aclrtContext context[DEV_NUM];
168- for (uint32_t rankId = 0; rankId < DEV_NUM; rankId++) {
169- ret = aclrtSetDevice(rankId);
170- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtSetDevice failed. ret = %d \n", ret); return ret);
171- ret = aclrtCreateContext(&context[rankId], rankId);
172- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtCreateContext failed. ret = %d \n", ret); return ret);
173- ret = aclrtCreateStream(&stream[rankId]);
174- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtCreateStream failed. ret = %d \n", ret); return ret);
175- }
176- int32_t devices[DEV_NUM];
177- for (int i = 0; i < DEV_NUM; i++) {
178- devices[i] = i;
179- }
180- // 初始化集合通信域
181- HcclComm comms[DEV_NUM];
182- ret = HcclCommInitAll(DEV_NUM, devices, comms);
183- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] HcclCommInitAll failed. ret = %d \n", ret); return ret);
184- 
185- Args args[DEV_NUM];
186- // 启动多线程
187- std::vector<std::unique_ptr<std::thread>> threads(DEV_NUM);
188- for (uint32_t rankId = 0; rankId < DEV_NUM; rankId++) {
189- args[rankId].rankId = rankId;
190- args[rankId].hcclComm = comms[rankId];
191- args[rankId].context = context[rankId];
192- args[rankId].stream = stream[rankId];
193- threads[rankId].reset(new (std::nothrow) std::thread(&LaunchOneThreadQuantAllReduce, std::ref(args[rankId])));
194- }
195- for (uint32_t rankId = 0; rankId < DEV_NUM; rankId++) {
196- threads[rankId]->join();
197- }
198- aclFinalize();
199- return 0;
200-}
@@ -29,195 +29,15 @@
29#include "common/utils/hccl_util.h"29#include "common/utils/hccl_util.h"
30#include "aclnn_kernels/transdata.h"30#include "aclnn_kernels/transdata.h"
31 31 
32-namespace {
33- 
34-using namespace op;
35-using namespace l0op;
36- 
37-enum class NnopbaseHcclServerType : uint32_t {
38- NNOPBASE_HCCL_SERVER_TYPE_AICPU = 0,
39- NNOPBASE_HCCL_SERVER_TYPE_MTE,
40- NNOPBASE_HCCL_SERVER_TYPE_CCU,
41- NNOPBASE_HCCL_SERVER_TYPE_END
42-};
43- 
44-static constexpr size_t HCCL_GROUP_NAME_LENGTH_MAX = 128U; // group长度小于128字符
45- 
46-// K-G量化支持的Dtype
47-static const std::initializer_list<op::DataType> X_DTYPE_KG_SUPPORT_LIST = {
48- op::DataType::DT_INT8, op::DataType::DT_HIFLOAT8, op::DataType::DT_FLOAT8_E4M3FN, op::DataType::DT_FLOAT8_E5M2};
49-static const std::initializer_list<op::DataType> SCALES_DTYPE_KG_SUPPORT_LIST = {op::DataType::DT_FLOAT};
50- 
51-// MX量化支持的Dtype
52-static const std::initializer_list<op::DataType> X_DTYPE_MX_SUPPORT_LIST = {op::DataType::DT_FLOAT8_E4M3FN,
53- op::DataType::DT_FLOAT8_E5M2};
54-static const std::initializer_list<op::DataType> SCALES_DTYPE_MX_SUPPORT_LIST = {op::DataType::DT_FLOAT8_E8M0};
55- 
56-// output支持的Dtype
57-static const std::initializer_list<op::DataType> OUTPUT_DTYPE_SUPPORT_LIST = {
58- op::DataType::DT_FLOAT16, op::DataType::DT_BF16, op::DataType::DT_FLOAT};
59- 
60-// 检查入参是否为nullptr
61-static bool QuantAllReduceCheckNotNull(const aclTensor *x, const aclTensor *scales, const aclTensor *output)
62-{
63- OP_CHECK_NULL(x, return false);
64- OP_CHECK_NULL(scales, return false);
65- OP_CHECK_NULL(output, return false);
66- return true;
67-}
68- 
69-// 检查K-G量化方案中x、scales、output的数据类型是否在算子的支持列表内
70-static bool QuantAllReduceCheckKGAllDtypesValid(const aclTensor *x, const aclTensor *scales, const aclTensor *output)
71-{
72- if (CheckType(x->GetDataType(), X_DTYPE_KG_SUPPORT_LIST) &&
73- CheckType(scales->GetDataType(), SCALES_DTYPE_KG_SUPPORT_LIST) &&
74- CheckType(output->GetDataType(), OUTPUT_DTYPE_SUPPORT_LIST)) {
75- return true;
76- } else {
77- return false;
78- }
79-}
80- 
81-// 检查MX量化方案中x、scales、output的数据类型是否在算子的支持列表内
82-static bool QuantAllReduceCheckMXAllDtypesValid(const aclTensor *x, const aclTensor *scales, const aclTensor *output)
83-{
84- if (CheckType(x->GetDataType(), X_DTYPE_MX_SUPPORT_LIST) &&
85- CheckType(scales->GetDataType(), SCALES_DTYPE_MX_SUPPORT_LIST) &&
86- CheckType(output->GetDataType(), OUTPUT_DTYPE_SUPPORT_LIST)) {
87- return true;
88- } else {
89- return false;
90- }
91-}
92- 
93-// 统一数据类型检查
94-static bool QuantAllReduceCheckAllDtypesValid(const aclTensor *x, const aclTensor *scales, const aclTensor *output)
95-{
96- bool isAllDtypesValid = false;
97- isAllDtypesValid = (QuantAllReduceCheckKGAllDtypesValid(x, scales, output) ||
98- QuantAllReduceCheckMXAllDtypesValid(x, scales, output));
99- if (!isAllDtypesValid) {
100- OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON("aclnnQuantAllReduceGetWorkspaceSize", "x/scales/output",
101- (std::string(op::ToString(x->GetDataType()).GetString()) + "/" +
102- op::ToString(scales->GetDataType()).GetString() + "/" +
103- op::ToString(output->GetDataType()).GetString())
104- .c_str(),
105- "The dtypes of x, scales and output must be the same.");
106- }
107- return isAllDtypesValid;
108-}
109- 
110-static bool QuantAllReduceCheckAllFormatValid(const aclTensor *x, const aclTensor *scales, const aclTensor *output)
111-{
112- if (IsPrivateFormat(x->GetStorageFormat())) {
113- OP_LOGE_FOR_INVALID_FORMAT("aclnnQuantAllReduceGetWorkspaceSize", "x",
114- op::ToString(x->GetStorageFormat()).GetString(), "non-Private Format");
115- return false;
116- }
117- if (IsPrivateFormat(scales->GetStorageFormat())) {
118- OP_LOGE_FOR_INVALID_FORMAT("aclnnQuantAllReduceGetWorkspaceSize", "scales",
119- op::ToString(scales->GetStorageFormat()).GetString(), "non-Private Format");
120- return false;
121- }
122- if (IsPrivateFormat(output->GetStorageFormat())) {
123- OP_LOGE_FOR_INVALID_FORMAT("aclnnQuantAllReduceGetWorkspaceSize", "output",
124- op::ToString(output->GetStorageFormat()).GetString(), "non-Private Format");
125- return false;
126- }
127- 
128- // 内部只处理ND格式,这里做reformat操作
129- if (x->GetStorageFormat() != op::Format::FORMAT_ND) {
130- OP_LOGW("x origin format is: %s.", op::ToString(x->GetStorageFormat()).GetString());
131- x = l0op::ReFormat(x, op::Format::FORMAT_ND);
132- CHECK_RET(x != nullptr, false);
133- }
134- if (scales->GetStorageFormat() != op::Format::FORMAT_ND) {
135- OP_LOGW("scales origin format is: %s.", op::ToString(scales->GetStorageFormat()).GetString());
136- scales = l0op::ReFormat(scales, op::Format::FORMAT_ND);
137- CHECK_RET(scales != nullptr, false);
138- }
139- if (output->GetStorageFormat() != op::Format::FORMAT_ND) {
140- OP_LOGW("output origin format is: %s.", op::ToString(output->GetStorageFormat()).GetString());
141- output = l0op::ReFormat(output, op::Format::FORMAT_ND);
142- CHECK_RET(output != nullptr, false);
143- }
144- 
145- return true;
146-}
147- 
148-static bool QuantAllReduceCheckGroupLength(const char *group)
149-{
150- if (group == nullptr) {
151- OP_LOGE_WITH_INVALID_INPUT("aclnnQuantAllReduceGetWorkspaceSize", "group");
152- return false;
153- }
154- 
155- size_t groupLen = strnlen(group, HCCL_GROUP_NAME_LENGTH_MAX); // group长度≥128字符, 返回HCCL_GROUP_NAME_LENGTH_MAX
156- if (groupLen >= HCCL_GROUP_NAME_LENGTH_MAX) {
157- OP_LOGE_FOR_INVALID_VALUE("aclnnQuantAllReduceGetWorkspaceSize", "group length",
158- std::to_string(groupLen).c_str(),
159- ("less than " + std::to_string(HCCL_GROUP_NAME_LENGTH_MAX)).c_str());
160- return false;
161- }
162- 
163- return true;
164-}
165- 
166-// 参数综合校验
167-static aclnnStatus QuantAllReduceCheckParams(const aclTensor *x, const aclTensor *scales, const char *group,
168- const aclTensor *output)
169-{
170- // 1. 检查参数是否为空指针
171- CHECK_RET(QuantAllReduceCheckNotNull(x, scales, output), ACLNN_ERR_PARAM_NULLPTR);
172- 
173- // 2. 检查输入的数据类型是否在API支持的数据类型范围之内,需要根据api定义校验
174- CHECK_RET(QuantAllReduceCheckAllDtypesValid(x, scales, output), ACLNN_ERR_PARAM_INVALID);
175- 
176- // 3. 检查参数数据格式是否在API支持的数据类型范围之内,需要根据api定义校验
177- CHECK_RET(QuantAllReduceCheckAllFormatValid(x, scales, output), ACLNN_ERR_PARAM_INVALID);
178- 
179- // 4. 检查group参数是否在要求范围之内
180- CHECK_RET(QuantAllReduceCheckGroupLength(group), ACLNN_ERR_PARAM_INVALID);
181- 
182- return ACLNN_SUCCESS;
183-}
184-} // namespace
185- 
186-extern "C" aclnnStatus aclnnInnerQuantAllReduceGetWorkspaceSize(const aclTensor *x, const aclTensor *scales,
187- const char *group, const char *reduceOp,
188- uint64_t yDtype, int64_t worldSize, aclTensor *output,
189- uint64_t *workspaceSize, aclOpExecutor **executor);
190- 
191-extern "C" aclnnStatus aclnnInnerQuantAllReduce(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor,
192- const aclrtStream stream);
193- 
194-extern "C" void __attribute__((weak)) NnopbaseSetHcclServerType(void *executor, NnopbaseHcclServerType sType);
195- 
196extern "C" aclnnStatus aclnnQuantAllReduceGetWorkspaceSize(const aclTensor *x, const aclTensor *scales,32extern "C" aclnnStatus aclnnQuantAllReduceGetWorkspaceSize(const aclTensor *x, const aclTensor *scales,
197 const char *group, const char *reduceOp, aclTensor *output,33 const char *group, const char *reduceOp, aclTensor *output,
198 uint64_t *workspaceSize, aclOpExecutor **executor)34 uint64_t *workspaceSize, aclOpExecutor **executor)
199{35{
200- aclnnStatus retParam = QuantAllReduceCheckParams(x, scales, group, output);36+ return ACLNN_SUCCESS;
201- CHECK_RET(retParam == ACLNN_SUCCESS, retParam);
202- uint64_t yDtype = static_cast<uint64_t>(output->GetDataType());
203- int64_t worldSize = -1;
204- aclnnStatus ret =
205- aclnnInnerQuantAllReduceGetWorkspaceSize(x, scales, const_cast<char *>(group), const_cast<char *>(reduceOp),
206- yDtype, worldSize, output, workspaceSize, executor);
207- OP_LOGD("QuantAllReduce, aclnnGetWorkspaceSize ret %d.", ret);
208- return ret;
209}37}
210 38 
211extern "C" aclnnStatus aclnnQuantAllReduce(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor,39extern "C" aclnnStatus aclnnQuantAllReduce(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor,
212 const aclrtStream stream)40 const aclrtStream stream)
213{41{
214- if (NnopbaseSetHcclServerType) {
215- NnopbaseSetHcclServerType(executor, NnopbaseHcclServerType::NNOPBASE_HCCL_SERVER_TYPE_MTE);
216- }
217- aclnnStatus ret = aclnnInnerQuantAllReduce(workspace, workspaceSize, executor, stream);
218- if (ret != ACLNN_SUCCESS) {
219- OP_LOGE_LIBOPAPI_REPORT("aclnnQuantAllReduce", "QuantAllReduce, This is an error in launch aicore");
220- return ACLNN_ERR_INNER;
221- }
222 return ACLNN_SUCCESS;42 return ACLNN_SUCCESS;
223}43}
@@ -1,150 +1,3 @@
1# aclnnQuantReduceScatter1# aclnnQuantReduceScatter
2 2 
3-## 产品支持情况3+该算子暂无Ascend C代码实现,欢迎开发者补充贡献。
4- 
5-| 产品 | 是否支持 |
6-| :------------------------------------------------------------------------------ | :------: |
7-| <term>Ascend 950DT</term> | √ |
8-| <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term> | × |
9-| <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term> | × |
10-| <term>Atlas 200I/500 A2 推理产品</term> | × |
11-| <term>Atlas 推理系列产品</term> | × |
12-| <term>Atlas 训练系列产品</term> | × |
13- 
14-**说明:** 使用该接口时,请确保驱动固件包和CANN包都为配套的8.0.RC2版本或者配套的更高版本,否则将会引发报错,比如BUS ERROR等。
15- 
16-## 功能说明
17- 
18-- **算子功能**:实现quant + reduceScatter融合计算。
19-- **计算公式**:
20- 
21- $$
22- output=Reduce(AllToAllScales * AllToAllData)
23- $$
24- 
25- $$
26- AllToAllData=AllToAll(x)
27- $$
28- 
29- $$
30- AllToAllScales=AllToAll(scales)
31- $$
32- 
33- 其中的Reduce计算是将来自不同rank的数据进行reduce计算。
34- $$
35- 
36-- **参数说明:**
37- <table style="undefined;table-layout: fixed; width: 1567px"><colgroup>
38- <col style="width: 170px">
39- <col style="width: 120px">
40- <col style="width: 300px">
41- <col style="width: 330px">
42- <col style="width: 212px">
43- <col style="width: 100px">
44- <col style="width: 190px">
45- <col style="width: 145px">
46- </colgroup>
47- <thead>
48- <tr>
49- <th>参数名</th>
50- <th>输入/输出</th>
51- <th>描述</th>
52- <th>使用说明</th>
53- <th>数据类型</th>
54- <th>数据格式</th>
55- <th>维度(shape)</th>
56- <th>连续Tensor</th>
57- </tr>
58- </thead>
59- <tbody>
60- <tr>
61- <td>x</td>
62- <td>输入</td>
63- <td>公式中的输入x</td>
64- <td><ul><li>不支持空Tensor。</li><li>支持的shape为:(BS, H)或者(B, S, H)。B为batch size,S为sequence length,H为hidden size。当前版本输入x的H支持1024~8192中任意128对齐泛化。</li></ul></td>
65- <td>INT8, HIFLOAT8, FLOAT8_E4M3FN, FLOAT8_E5M2</td>
66- <td>ND</td>
67- <td>2-3</td>
68- <td>√</td>
69- </tr>
70- <tr>
71- <td>scales</td>
72- <td>输入</td>
73- <td>公式中的输入scales</td>
74- <td><ul><li>不支持空Tensor。</li><li>当scales的数据类型为FLOAT8_E8M0时,x的数据类型必须为FLOAT8_E4M3FN、FLOAT8_E5M2,x的shape为(BS, H)或者(B, S, H),scales的shape必须对应x的shape为(BS, H/64, 2)或者(B, S, H/64, 2)。</li><li>当scales的数据类型为FLOAT时,x的数据类型必须为INT8、HIFLOAT8、FLOAT8_E4M3FN、FLOAT8_E5M2,x的shape为(BS, H)或者(B, S, H),scales的shape必须对应x的shape为(BS, H/128)或者(B, S, H/128)。</li></ul></td>
75- <td>FLOAT, FLOAT8_E8M0</td>
76- <td>ND</td>
77- <td>2-4</td>
78- <td>√</td>
79- </tr>
80- <tr>
81- <td>group</td>
82- <td>输入</td>
83- <td>通信域标识</td>
84- <td>通信域标识</td>
85- <td>String</td>
86- <td>-</td>
87- <td>-</td>
88- <td>-</td>
89- </tr>
90- <tr>
91- <td>reduceOp</td>
92- <td>输入</td>
93- <td>公式中的reduce操作类型。</td>
94- <td>当前仅支持"sum"</td>
95- <td>string</td>
96- <td>-</td>
97- <td>-</td>
98- <td>-</td>
99- </tr>
100- <tr>
101- <td>output</td>
102- <td>输出</td>
103- <td>公式中的输出output。</td>
104- <td><ul><li>不支持空Tensor。</li><li>当x的shape是(BS,H)的时候,output的shape必须为(BS/rankNum,H);当x的shape是(B,S,H)的时候,output的shape必须为(B*S/rankNum,H)。rankNum表示通信域大小。</li></ul></td>
105- <td>FLOAT、FLOAT16、BFLOAT16</td>
106- <td>ND</td>
107- <td>2</td>
108- <td>√</td>
109- </tr>
110- <tr>
111- <td>workspaceSize</td>
112- <td>输出</td>
113- <td>返回需要在Device侧申请的workspace大小。</td>
114- <td>-</td>
115- <td>-</td>
116- <td>-</td>
117- <td>-</td>
118- <td>-</td>
119- </tr>
120- <tr>
121- <td>executor</td>
122- <td>输出</td>
123- <td>返回op执行器,包含了算子计算流程。</td>
124- <td>-</td>
125- <td>-</td>
126- <td>-</td>
127- <td>-</td>
128- <td>-</td>
129- </tr>
130- </tbody>
131- </table>
132- 
133-## 约束说明
134- 
135-- 当x的数据类型为FLOAT8_E4M3FN, FLOAT8_E5M2并且scales的数据类型为FLOAT8_E8M0时,输入数据的量化方式为mx量化。
136-- 当x的数据类型为INT8、HIFLOAT8、FLOAT8_E4M3FN, FLOAT8_E5M2并且scales的数据类型为FLOAT时,输入数据的量化方式为pertoken-pergroup量化(groupSize=128)。
137-- 只在Ascend950系列平台开启。
138-- 不支持空tensor输入。
139-- 通信引擎约束:
140- - Ascend950DT: 仅支持UB-Memory通信。
141-- 通信域大小支持2、4、8。
142-- 通信域使用约束:同一通信域内仅允许连续执行`aclnnQuantAllReduce`和`aclnnQuantReduceScatter`算子,且该通信域中不允许有其他通信算子。
143-- `HCCL_BUFFSIZE`:调用本算子前需检查`HCCL_BUFFSIZE`环境变量取值是否合理,该环境变量表示单个通信域占用内存大小,单位MB,不配置时默认为200MB。要求满足`HCCL_BUFFSIZE`>= 2 * (`xDataSize` + `scalesDataSize + 1`)。其中`xDataSize`为输入`x`的数据大小,计算公式为:`xDataSize = BS * H * 1 (Byte)`,`scalesDataSize`为`scales`的数据大小,当量化方式为pertoken-pergroup量化时,计算公式为:`scalesDataSize = BS * H / 128 * 4 (Byte)`,当量化方式为mx量化时,计算公式为:`scalesDataSize = BS * H / 32 * 1 (Byte)`。
144-- H范围仅支持[1024, 8192],要求128对齐。
145- 
146-## 调用说明
147- 
148-| 调用方式 | 样例代码 | 说明 |
149-| :--------: | :----------------------------------------: | :-------------------------------------------------------: |
150-| aclnn接口 | [test_aclnn_quant_reduce_scatter.cpp](./examples/test_aclnn_quant_reduce_scatter.cpp) | 通过[aclnnQuantReduceScatter](./docs/aclnnQuantReduceScatter.md)接口方式调用quant_reduce_scatter算子。 |
@@ -1,466 +0,0 @@
1-# aclnnQuantReduceScatter
2- 
3-## 产品支持情况
4- 
5-<!-- npu="950" id1 -->
6-- <term>Ascend 950DT</term>:支持
7-<!-- end id1 -->
8-<!-- npu="A3" id2 -->
9-- <term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>:不支持
10-<!-- end id2 -->
11-<!-- npu="910b" id3 -->
12-- <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>:不支持
13-<!-- end id3 -->
14-<!-- npu="310b" id4 -->
15-- <term>Atlas 200I/500 A2 推理产品</term>:不支持
16-<!-- end id4 -->
17-<!-- npu="310p" id5 -->
18-- <term>Atlas 推理系列产品</term>:不支持
19-<!-- end id5 -->
20-<!-- npu="910" id6 -->
21-- <term>Atlas 训练系列产品</term>:不支持
22-<!-- end id6 -->
23- 
24-**说明:** 使用该接口时,请确保驱动固件包和CANN包都为配套的8.0.RC2版本或者配套的更高版本,否则将会引发报错,比如BUS ERROR等。
25- 
26-## 功能说明
27- 
28-- **接口功能**:实现quant + reduceScatter融合计算。
29-- **计算公式**:
30- 
31- $$
32- output=Reduce(AllToAllScales * AllToAllData)
33- $$
34- 
35- $$
36- AllToAllData=AllToAll(x)
37- $$
38- 
39- $$
40- AllToAllScales=AllToAll(scales)
41- $$
42- 
43- 其中的Reduce计算是将来自不同rank的数据进行reduce计算。
44- 
45-## 函数原型
46- 
47-该算子分为两段式接口,必须先调用“aclnnQuantReduceScatterGetWorkspaceSize”接口获取计算所需workspace大小以及包含了算子计算流程的执行器,再调用“aclnnQuantReduceScatter”接口执行计算。
48- 
49-```cpp
50-aclnnStatus aclnnQuantReduceScatterGetWorkspaceSize(
51- const aclTensor *x,
52- const aclTensor *scales,
53- const char *group,
54- const char *reduceOp,
55- aclTensor *output,
56- uint64_t *workspaceSize,
57- aclOpExecutor **executor)
58-```
59- 
60-```cpp
61-aclnnStatus aclnnQuantReduceScatter(
62- void *workspace,
63- uint64_t workspaceSize,
64- aclOpExecutor *executor,
65- const aclrtStream stream)
66-```
67- 
68-## aclnnQuantReduceScatterGetWorkspaceSize
69- 
70-- **参数说明**
71- 
72- <table style="undefined;table-layout: fixed; width: 1556px"><colgroup>
73- <col style="width: 161px">
74- <col style="width: 141px">
75- <col style="width: 245px">
76- <col style="width: 408px">
77- <col style="width: 191px">
78- <col style="width: 120px">
79- <col style="width: 145px">
80- <col style="width: 145px">
81- </colgroup>
82- <thead>
83- <tr>
84- <th>参数名</th>
85- <th>输入/输出</th>
86- <th>描述</th>
87- <th>使用说明</th>
88- <th>数据类型</th>
89- <th>数据格式</th>
90- <th>维度(shape)</th>
91- <th>连续Tensor</th>
92- </tr>
93- </thead>
94- <tbody>
95- <tr>
96- <td>x</td>
97- <td>输入</td>
98- <td>公式中的输入x。</td>
99- <td><ul><li>不支持空Tensor。</li><li>支持的shape为:(BS, H)或者(B, S, H)。B为batch size,S为sequence length,H为hidden size。当前版本输入x的H支持1024~8192中任意128对齐泛化。</li></ul></td>
100- <td>INT8, HIFLOAT8, FLOAT8_E4M3FN, FLOAT8_E5M2</td>
101- <td>ND</td>
102- <td>2-3</td>
103- <td>√</td>
104- </tr>
105- <tr>
106- <td>scales</td>
107- <td>输入</td>
108- <td>公式中的输入scales。</td>
109- <td><ul><li>不支持空Tensor。</li><li>当scales的数据类型为FLOAT8_E8M0时,x的数据类型必须为FLOAT8_E4M3FN、FLOAT8_E5M2,x的shape为(BS, H)或者(B, S, H),scales的shape必须对应x的shape为(BS, H/64, 2)或者(B, S, H/64, 2)。</li><li>当scales的数据类型为FLOAT时,x的数据类型必须为INT8、HIFLOAT8、FLOAT8_E4M3FN、FLOAT8_E5M2,x的shape为(BS, H)或者(B, S, H),scales的shape必须对应x的shape为(BS, H/128)或者(B, S, H/128)。</li></ul></td>
110- <td>FLOAT, FLOAT8_E8M0</td>
111- <td>ND</td>
112- <td>2-4</td>
113- <td>√</td>
114- </tr>
115- <tr>
116- <td>group</td>
117- <td>输入</td>
118- <td>通信域标识。</td>
119- <td>通信域标识</td>
120- <td>String</td>
121- <td>-</td>
122- <td>-</td>
123- <td>-</td>
124- </tr>
125- <tr>
126- <td>reduceOp</td>
127- <td>输入</td>
128- <td>公式中的reduce操作类型。</td>
129- <td>当前仅支持"sum"</td>
130- <td>string</td>
131- <td>-</td>
132- <td>-</td>
133- <td>-</td>
134- </tr>
135- <tr>
136- <td>output</td>
137- <td>输出</td>
138- <td>公式中的输出output。</td>
139- <td><ul><li>不支持空Tensor。</li><li>当x的shape是(BS,H)的时候,output的shape必须为(BS/rankNum,H);当x的shape是(B,S,H)的时候,output的shape必须为(B*S/rankNum,H)。rankNum表示通信域大小。</li></ul></td>
140- <td>FLOAT、FLOAT16、BFLOAT16</td>
141- <td>ND</td>
142- <td>2</td>
143- <td>√</td>
144- </tr>
145- <tr>
146- <td>workspaceSize</td>
147- <td>输出</td>
148- <td>返回需要在Device侧申请的workspace大小。</td>
149- <td>-</td>
150- <td>-</td>
151- <td>-</td>
152- <td>-</td>
153- <td>-</td>
154- </tr>
155- <tr>
156- <td>executor</td>
157- <td>输出</td>
158- <td>返回op执行器,包含了算子计算流程。</td>
159- <td>-</td>
160- <td>-</td>
161- <td>-</td>
162- <td>-</td>
163- <td>-</td>
164- </tr>
165- </tbody>
166- </table>
167- 
168-- **返回值**
169- 
170- aclnnStatus:返回状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。
171- 
172- 第一段接口完成入参校验,出现以下场景时报错:
173- 
174- <table style="undefined;table-layout: fixed; width: 1149px"><colgroup>
175- <col style="width: 282px">
176- <col style="width: 120px">
177- <col style="width: 747px">
178- </colgroup>
179- <thead>
180- <tr>
181- <th>返回值</th>
182- <th>错误码</th>
183- <th>描述</th>
184- </tr></thead>
185- <tbody>
186- <tr>
187- <td>ACLNN_ERR_PARAM_NULLPTR</td>
188- <td>161001</td>
189- <td>x、scales、output存在空指针。</td>
190- </tr>
191- <tr>
192- <td rowspan="3">ACLNN_ERR_PARAM_INVALID</td>
193- <td rowspan="3">161002</td>
194- <td>x、scales或output的数据类型不在支持的范围之内。</td>
195- </tr>
196- <tr>
197- <td>x、scales的数据类型和shape不匹配。</td>
198- </tr>
199- <tr>
200- <td>x、scales的维度不在支持的范围之内。</td>
201- </tr>
202- </tbody>
203- </table>
204- 
205-## aclnnQuantReduceScatter
206- 
207-- **参数说明**
208- 
209- <table style="undefined;table-layout: fixed; width: 1150px"><colgroup>
210- <col style="width: 168px">
211- <col style="width: 128px">
212- <col style="width: 854px">
213- </colgroup>
214- <thead>
215- <tr>
216- <th>参数名</th>
217- <th>输入/输出</th>
218- <th>描述</th>
219- </tr></thead>
220- <tbody>
221- <tr>
222- <td>workspace</td>
223- <td>输入</td>
224- <td>在Device侧申请的workspace内存地址。</td>
225- </tr>
226- <tr>
227- <td>workspaceSize</td>
228- <td>输入</td>
229- <td>在Device侧申请的workspace大小,由第一段接口aclnnQuantReduceScatterGetWorkspaceSize获取。</td>
230- </tr>
231- <tr>
232- <td>executor</td>
233- <td>输入</td>
234- <td>op执行器,包含了算子计算流程。</td>
235- </tr>
236- <tr>
237- <td>stream</td>
238- <td>输入</td>
239- <td>指定执行任务的stream。</td>
240- </tr>
241- </tbody></table>
242- 
243-- **返回值**
244- 
245- 返回aclnnStatus状态码,具体参见[aclnn返回码](../../../docs/zh/context/aclnn_return_code.md)。
246- 
247-## 约束说明
248- 
249-- 当x的数据类型为FLOAT8_E4M3FN, FLOAT8_E5M2并且scales的数据类型为FLOAT8_E8M0时,输入数据的量化方式为mx量化。
250-- 当x的数据类型为INT8、HIFLOAT8、FLOAT8_E4M3FN, FLOAT8_E5M2并且scales的数据类型为FLOAT时,输入数据的量化方式为pertoken-pergroup量化(groupSize=128)。
251-- 只在Ascend950系列平台开启。
252-- 不支持空tensor输入。
253-- 通信引擎约束:
254- - Ascend950DT: 仅支持UB-Memory通信。
255-- 通信域大小支持2、4、8。
256-- 通信域使用约束:同一通信域内仅允许连续执行`aclnnQuantAllReduce`和`aclnnQuantReduceScatter`算子,且该通信域中不允许有其他通信算子。
257-- `HCCL_BUFFSIZE`:调用本算子前需检查`HCCL_BUFFSIZE`环境变量取值是否合理,该环境变量表示单个通信域占用内存大小,单位MB,不配置时默认为200MB。要求满足`HCCL_BUFFSIZE`>= 2 * (`xDataSize` + `scalesDataSize + 1`)。其中`xDataSize`为输入`x`的数据大小,计算公式为:`xDataSize = BS * H * 1 (Byte)`,`scalesDataSize`为`scales`的数据大小,当量化方式为pertoken-pergroup量化时,计算公式为:`scalesDataSize = BS * H / 128 * 4 (Byte)`,当量化方式为mx量化时,计算公式为:`scalesDataSize = BS * H / 32 * 1 (Byte)`。
258-- H范围仅支持[1024, 8192],要求128对齐。
259- 
260-## 调用示例
261- 
262-示例代码如下,仅供参考,具体编译和执行过程请参考编译与运行样例。
263- 
264-- <term>Atlas A5训练系列产品/Atlas A5推理系列产品</term>:
265- 
266- ```Cpp
267- #include <thread>
268- #include <iostream>
269- #include <vector>
270- #include <string>
271- #include <cstring>
272- #include "hccl/hccl.h"
273- #include "aclnnop/aclnn_quant_reduce_scatter.h"
274- 
275- #define CHECK_RET(cond, return_expr) \
276- do { \
277- if (!(cond)) { \
278- return_expr; \
279- } \
280- } while (0)
281- 
282- #define LOG_PRINT(message, ...) \
283- do { \
284- printf(message, ##__VA_ARGS__); \
285- } while (0)
286- 
287- constexpr int DEV_NUM = 2;
288- 
289- int64_t GetShapeSize(const std::vector<int64_t> &shape)
290- {
291- int64_t shape_size = 1;
292- for (auto i : shape) {
293- shape_size *= i;
294- }
295- return shape_size;
296- }
297- 
298- template<typename T>
299- int CreateAclTensor(const std::vector<T> &hostData, const std::vector<int64_t> &shape, void **deviceAddr,
300- aclDataType dataType, aclTensor **tensor)
301- {
302- auto size = GetShapeSize(shape) * sizeof(T);
303- auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST);
304- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtMalloc failed. ret: %d\n", ret);
305- return ret);
306- ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE);
307- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtMemcpy failed. ret: %d\n", ret);
308- return ret);
309- std::vector<int64_t> strides(shape.size(), 1);
310- for (int64_t i = shape.size() - 2; i >= 0; i--) {
311- strides[i] = shape[i + 1] * strides[i + 1];
312- }
313- *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND,
314- shape.data(), shape.size(), *deviceAddr);
315- return 0;
316- }
317- 
318- struct Args {
319- int rankId;
320- HcclComm hcclComm;
321- aclrtStream stream;
322- aclrtContext context;
323- };
324- 
325- int LaunchOneThreadQtReduceScatter(Args &args)
326- {
327- int ret = aclrtSetCurrentContext(args.context);
328- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtSetCurrentContext failed. ret = %d\n", ret);
329- return ret);
330- char hcomName[128] = {0};
331- ret = HcclGetCommName(args.hcclComm, hcomName);
332- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] HcclGetCommName failed. ret = %d\n", ret);
333- return -1);
334- LOG_PRINT("[INFO] rank = %d, hcomName = %s, stream = %p\n", args.rankId, hcomName, args.stream);
335- std::vector<int64_t> xShape = {1024, 5120};
336- std::vector<int64_t> scalesShape = {1024, 40};
337- std::vector<int64_t> outputShape = {1024 / DEV_NUM, 5120};
338- void *xDeviceAddr = nullptr;
339- void *scalesDeviceAddr = nullptr;
340- void *outputDeviceAddr = nullptr;
341- void *workspaceAddr = nullptr;
342- 
343- aclTensor *x = nullptr;
344- aclTensor *scales = nullptr;
345- aclTensor *output = nullptr;
346- uint64_t workspaceSize = 0;
347- aclOpExecutor *executor = nullptr;
348- 
349- long long xShapeSize = GetShapeSize(xShape);
350- long long scalesShapeSize = GetShapeSize(scalesShape);
351- long long outputShapeSize = GetShapeSize(outputShape);
352- 
353- std::vector<int8_t> xHostData(xShapeSize, 0);
354- std::vector<int8_t> scalesHostData(scalesShapeSize, 0);
355- std::vector<int16_t> outputHostData(outputShapeSize, 0);
356- 
357- // 创建tensor
358- ret = CreateAclTensor(xHostData, xShape, &xDeviceAddr, aclDataType::ACL_FLOAT8_E5M2, &x);
359- CHECK_RET(ret == ACL_SUCCESS, return ret);
360- ret = CreateAclTensor(scalesHostData, scalesShape, &scalesDeviceAddr, aclDataType::ACL_FLOAT, &scales);
361- CHECK_RET(ret == ACL_SUCCESS, return ret);
362- ret = CreateAclTensor(outputHostData, outputShape, &outputDeviceAddr, aclDataType::ACL_FLOAT16, &output);
363- CHECK_RET(ret == ACL_SUCCESS, return ret);
364- 
365- // 调用第一阶段接口
366- ret = aclnnQuantReduceScatterGetWorkspaceSize(
367- x, scales, hcomName, "sum", output, &workspaceSize, &executor);
368- CHECK_RET(ret == ACL_SUCCESS,
369- LOG_PRINT("[ERROR] aclnnQuantReduceScatterGetWorkspaceSize failed. ret = %d \n", ret);
370- return ret);
371- // 根据第一阶段接口计算出的workspaceSize申请device内存
372- if (workspaceSize > 0) {
373- ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST);
374- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtMalloc workspace failed. ret = %d \n", ret);
375- return ret);
376- }
377- // 调用第二阶段接口
378- ret = aclnnQuantReduceScatter(workspaceAddr, workspaceSize, executor, args.stream);
379- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclnnQuantReduceScatter failed. ret = %d \n", ret);
380- return ret);
381- //(固定写法)同步等待任务执行结束
382- ret = aclrtSynchronizeStreamWithTimeout(args.stream, 10000);
383- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtSynchronizeStreamWithTimeout failed. ret = %d \n", ret);
384- return ret);
385- LOG_PRINT("[INFO] device_%d aclnnQuantReduceScatter execute successfully.\n", args.rankId);
386- // 释放device资源,需要根据具体API的接口定义修改
387- if (x != nullptr) {
388- aclDestroyTensor(x);
389- }
390- if (scales != nullptr) {
391- aclDestroyTensor(scales);
392- }
393- if (output != nullptr) {
394- aclDestroyTensor(output);
395- }
396- 
397- if (xDeviceAddr != nullptr) {
398- aclrtFree(xDeviceAddr);
399- }
400- if (scalesDeviceAddr != nullptr) {
401- aclrtFree(scalesDeviceAddr);
402- }
403- if (outputDeviceAddr != nullptr) {
404- aclrtFree(outputDeviceAddr);
405- }
406- if (workspaceSize > 0) {
407- aclrtFree(workspaceAddr);
408- }
409- ret = aclrtDestroyStream(args.stream);
410- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtDestroyStream failed. ret = %d \n", ret);
411- return ret);
412- 
413- ret = HcclCommDestroy(args.hcclComm);
414- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] HcclCommDestroy failed. ret = %d \n", ret);
415- return ret);
416- 
417- ret = aclrtDestroyContext(args.context);
418- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtDestroyContext failed. ret = %d \n", ret);
419- return ret);
420- 
421- ret = aclrtResetDevice(args.rankId);
422- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtResetDevice failed. ret = %d \n", ret);
423- return ret);
424- 
425- return 0;
426- }
427- int main(int argc, char *argv[])
428- {
429- int ret = aclInit(nullptr);
430- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclInit failed. ret = %d \n", ret); return ret);
431- aclrtStream stream[DEV_NUM];
432- aclrtContext context[DEV_NUM];
433- for (uint32_t rankId = 0; rankId < DEV_NUM; rankId++) {
434- ret = aclrtSetDevice(rankId);
435- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtSetDevice failed. ret = %d \n", ret); return ret);
436- ret = aclrtCreateContext(&context[rankId], rankId);
437- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtCreateContext failed. ret = %d \n", ret); return ret);
438- ret = aclrtCreateStream(&stream[rankId]);
439- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtCreateStream failed. ret = %d \n", ret); return ret);
440- }
441- int32_t devices[DEV_NUM];
442- for (int i = 0; i < DEV_NUM; i++) {
443- devices[i] = i;
444- }
445- // 初始化集合通信域
446- HcclComm comms[DEV_NUM];
447- ret = HcclCommInitAll(DEV_NUM, devices, comms);
448- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] HcclCommInitAll failed. ret = %d \n", ret); return ret);
449- 
450- Args args[DEV_NUM];
451- // 启动多线程
452- std::vector<std::unique_ptr<std::thread>> threads(DEV_NUM);
453- for (uint32_t rankId = 0; rankId < DEV_NUM; rankId++) {
454- args[rankId].rankId = rankId;
455- args[rankId].hcclComm = comms[rankId];
456- args[rankId].context = context[rankId];
457- args[rankId].stream = stream[rankId];
458- threads[rankId].reset(new(std::nothrow) std::thread(&LaunchOneThreadQtReduceScatter, std::ref(args[rankId])));
459- }
460- for (uint32_t rankId = 0; rankId < DEV_NUM; rankId++) {
461- threads[rankId]->join();
462- }
463- aclFinalize();
464- return 0;
465- }
466- ```
@@ -1,202 +0,0 @@
1-/**
2- * Copyright (c) 2025 Huawei Technologies Co., Ltd.
3- * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4- * CANN Open Software License Agreement Version 2.0 (the "License").
5- * Please refer to the License for details. You may not use this file except in compliance with the License.
6- * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7- * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8- * See LICENSE in the root of the software repository for the full text of the License.
9- */
10- 
11-/*!
12- * \file test_aclnn_quant_reduce_scatter.cpp
13- * \brief
14- */
15- 
16-#include <thread>
17-#include <iostream>
18-#include <vector>
19-#include <string>
20-#include <cstring>
21-#include "hccl/hccl.h"
22-#include "aclnnop/aclnn_quant_reduce_scatter.h"
23- 
24-#define CHECK_RET(cond, return_expr) \
25- do { \
26- if (!(cond)) { \
27- return_expr; \
28- } \
29- } while (0)
30- 
31-#define LOG_PRINT(message, ...) \
32- do { \
33- printf(message, ##__VA_ARGS__); \
34- } while (0)
35- 
36-constexpr int DEV_NUM = 2;
37- 
38-int64_t GetShapeSize(const std::vector<int64_t> &shape)
39-{
40- int64_t shape_size = 1;
41- for (auto i : shape) {
42- shape_size *= i;
43- }
44- return shape_size;
45-}
46- 
47-template <typename T>
48-int CreateAclTensor(const std::vector<T> &hostData, const std::vector<int64_t> &shape, void **deviceAddr,
49- aclDataType dataType, aclTensor **tensor)
50-{
51- auto size = GetShapeSize(shape) * sizeof(T);
52- auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST);
53- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtMalloc failed. ret: %d\n", ret); return ret);
54- ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE);
55- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtMemcpy failed. ret: %d\n", ret); return ret);
56- std::vector<int64_t> strides(shape.size(), 1);
57- for (int64_t i = shape.size() - 2; i >= 0; i--) {
58- strides[i] = shape[i + 1] * strides[i + 1];
59- }
60- *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND,
61- shape.data(), shape.size(), *deviceAddr);
62- return 0;
63-}
64- 
65-struct Args {
66- int rankId;
67- HcclComm hcclComm;
68- aclrtStream stream;
69- aclrtContext context;
70-};
71- 
72-int LaunchOneThreadQtReduceScatter(Args &args)
73-{
74- int ret = aclrtSetCurrentContext(args.context);
75- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtSetCurrentContext failed. ret = %d\n", ret); return ret);
76- char hcomName[128] = {0};
77- ret = HcclGetCommName(args.hcclComm, hcomName);
78- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] HcclGetCommName failed. ret = %d\n", ret); return -1);
79- LOG_PRINT("[INFO] rank = %d, hcomName = %s, stream = %p\n", args.rankId, hcomName, args.stream);
80- std::vector<int64_t> xShape = {1024, 5120};
81- std::vector<int64_t> scalesShape = {1024, 40};
82- std::vector<int64_t> outputShape = {1024 / DEV_NUM, 5120};
83- void *xDeviceAddr = nullptr;
84- void *scalesDeviceAddr = nullptr;
85- void *outputDeviceAddr = nullptr;
86- void *workspaceAddr = nullptr;
87- 
88- aclTensor *x = nullptr;
89- aclTensor *scales = nullptr;
90- aclTensor *output = nullptr;
91- uint64_t workspaceSize = 0;
92- aclOpExecutor *executor = nullptr;
93- 
94- long long xShapeSize = GetShapeSize(xShape);
95- long long scalesShapeSize = GetShapeSize(scalesShape);
96- long long outputShapeSize = GetShapeSize(outputShape);
97- 
98- std::vector<int8_t> xHostData(xShapeSize, 0);
99- std::vector<int8_t> scalesHostData(scalesShapeSize, 0);
100- std::vector<int16_t> outputHostData(outputShapeSize, 0);
101- 
102- // 创建tensor
103- ret = CreateAclTensor(xHostData, xShape, &xDeviceAddr, aclDataType::ACL_FLOAT8_E5M2, &x);
104- CHECK_RET(ret == ACL_SUCCESS, return ret);
105- ret = CreateAclTensor(scalesHostData, scalesShape, &scalesDeviceAddr, aclDataType::ACL_FLOAT, &scales);
106- CHECK_RET(ret == ACL_SUCCESS, return ret);
107- ret = CreateAclTensor(outputHostData, outputShape, &outputDeviceAddr, aclDataType::ACL_FLOAT16, &output);
108- CHECK_RET(ret == ACL_SUCCESS, return ret);
109- 
110- // 调用第一阶段接口
111- ret = aclnnQuantReduceScatterGetWorkspaceSize(x, scales, hcomName, "sum", output, &workspaceSize, &executor);
112- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclnnQuantReduceScatterGetWorkspaceSize failed. ret = %d \n", ret);
113- return ret);
114- // 根据第一阶段接口计算出的workspaceSize申请device内存
115- if (workspaceSize > 0) {
116- ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST);
117- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtMalloc workspace failed. ret = %d \n", ret); return ret);
118- }
119- // 调用第二阶段接口
120- ret = aclnnQuantReduceScatter(workspaceAddr, workspaceSize, executor, args.stream);
121- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclnnQuantReduceScatter failed. ret = %d \n", ret); return ret);
122- // (固定写法)同步等待任务执行结束
123- ret = aclrtSynchronizeStreamWithTimeout(args.stream, 10000);
124- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtSynchronizeStreamWithTimeout failed. ret = %d \n", ret);
125- return ret);
126- LOG_PRINT("[INFO] device_%d aclnnQuantReduceScatter execute successfully.\n", args.rankId);
127- // 释放device资源,需要根据具体API的接口定义修改
128- if (x != nullptr) {
129- aclDestroyTensor(x);
130- }
131- if (scales != nullptr) {
132- aclDestroyTensor(scales);
133- }
134- if (output != nullptr) {
135- aclDestroyTensor(output);
136- }
137- 
138- if (xDeviceAddr != nullptr) {
139- aclrtFree(xDeviceAddr);
140- }
141- if (scalesDeviceAddr != nullptr) {
142- aclrtFree(scalesDeviceAddr);
143- }
144- if (outputDeviceAddr != nullptr) {
145- aclrtFree(outputDeviceAddr);
146- }
147- if (workspaceSize > 0) {
148- aclrtFree(workspaceAddr);
149- }
150- ret = aclrtDestroyStream(args.stream);
151- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtDestroyStream failed. ret = %d \n", ret); return ret);
152- 
153- ret = HcclCommDestroy(args.hcclComm);
154- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] HcclCommDestroy failed. ret = %d \n", ret); return ret);
155- 
156- ret = aclrtDestroyContext(args.context);
157- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtDestroyContext failed. ret = %d \n", ret); return ret);
158- 
159- ret = aclrtResetDevice(args.rankId);
160- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtResetDevice failed. ret = %d \n", ret); return ret);
161- 
162- return 0;
163-}
164-int main(int argc, char *argv[])
165-{
166- int ret = aclInit(nullptr);
167- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclInit failed. ret = %d \n", ret); return ret);
168- aclrtStream stream[DEV_NUM];
169- aclrtContext context[DEV_NUM];
170- for (uint32_t rankId = 0; rankId < DEV_NUM; rankId++) {
171- ret = aclrtSetDevice(rankId);
172- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtSetDevice failed. ret = %d \n", ret); return ret);
173- ret = aclrtCreateContext(&context[rankId], rankId);
174- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtCreateContext failed. ret = %d \n", ret); return ret);
175- ret = aclrtCreateStream(&stream[rankId]);
176- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] aclrtCreateStream failed. ret = %d \n", ret); return ret);
177- }
178- int32_t devices[DEV_NUM];
179- for (int i = 0; i < DEV_NUM; i++) {
180- devices[i] = i;
181- }
182- // 初始化集合通信域
183- HcclComm comms[DEV_NUM];
184- ret = HcclCommInitAll(DEV_NUM, devices, comms);
185- CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("[ERROR] HcclCommInitAll failed. ret = %d \n", ret); return ret);
186- 
187- Args args[DEV_NUM];
188- // 启动多线程
189- std::vector<std::unique_ptr<std::thread>> threads(DEV_NUM);
190- for (uint32_t rankId = 0; rankId < DEV_NUM; rankId++) {
191- args[rankId].rankId = rankId;
192- args[rankId].hcclComm = comms[rankId];
193- args[rankId].context = context[rankId];
194- args[rankId].stream = stream[rankId];
195- threads[rankId].reset(new (std::nothrow) std::thread(&LaunchOneThreadQtReduceScatter, std::ref(args[rankId])));
196- }
197- for (uint32_t rankId = 0; rankId < DEV_NUM; rankId++) {
198- threads[rankId]->join();
199- }
200- aclFinalize();
201- return 0;
202-}
@@ -27,144 +27,16 @@
27#include "common/utils/hccl_util.h"27#include "common/utils/hccl_util.h"
28#include "mc2_log_compat.h"28#include "mc2_log_compat.h"
29 29 
30-using namespace op;
31- 
32-namespace {
33-enum class NnopbaseHcclServerType : uint32_t {
34- NNOPBASE_HCCL_SERVER_TYPE_AICPU = 0,
35- NNOPBASE_HCCL_SERVER_TYPE_MTE,
36- NNOPBASE_HCCL_SERVER_TYPE_CCU,
37- NNOPBASE_HCCL_SERVER_TYPE_END
38-};
39- 
40-static constexpr size_t HCCL_GROUP_NAME_LENGTH_MAX = 128U; // group长度小于128字符
41- 
42-// 根据API定义,列出K-G量化所能支持的所有dtype
43-const std::initializer_list<op::DataType> X_DTYPE_KG_SUPPORT_LIST = {
44- op::DataType::DT_INT8, op::DataType::DT_HIFLOAT8, op::DataType::DT_FLOAT8_E4M3FN, op::DataType::DT_FLOAT8_E5M2};
45-const std::initializer_list<op::DataType> SCALES_DTYPE_KG_SUPPORT_LIST = {op::DataType::DT_FLOAT};
46- 
47-// 根据API定义,列出MX量化所能支持的所有dtype
48-const std::initializer_list<op::DataType> X_DTYPE_MX_SUPPORT_LIST = {op::DataType::DT_FLOAT8_E4M3FN,
49- op::DataType::DT_FLOAT8_E5M2};
50-const std::initializer_list<op::DataType> SCALES_DTYPE_MX_SUPPORT_LIST = {op::DataType::DT_FLOAT8_E8M0};
51- 
52-const std::initializer_list<op::DataType> OUTPUT_DTYPE_SUPPORT_LIST = {op::DataType::DT_FLOAT16, op::DataType::DT_BF16,
53- op::DataType::DT_FLOAT};
54- 
55-// 检查入参是否为nullptr
56-static bool CheckNotNull(const aclTensor *x, const aclTensor *scales, const aclTensor *output)
57-{
58- OP_CHECK_NULL(x, return false);
59- OP_CHECK_NULL(scales, return false);
60- OP_CHECK_NULL(output, return false);
61- return true;
62-}
63- 
64-// 检查x、scales、output的数据类型是否在算子的支持列表之内
65-static bool CheckKGAllDtypesValid(const aclTensor *x, const aclTensor *scales, const aclTensor *output)
66-{
67- if (CheckType(x->GetDataType(), X_DTYPE_KG_SUPPORT_LIST) &&
68- CheckType(scales->GetDataType(), SCALES_DTYPE_KG_SUPPORT_LIST) &&
69- CheckType(output->GetDataType(), OUTPUT_DTYPE_SUPPORT_LIST)) {
70- return true;
71- } else {
72- return false;
73- }
74-}
75- 
76-static bool CheckMXAllDtypesValid(const aclTensor *x, const aclTensor *scales, const aclTensor *output)
77-{
78- if (CheckType(x->GetDataType(), X_DTYPE_MX_SUPPORT_LIST) &&
79- CheckType(scales->GetDataType(), SCALES_DTYPE_MX_SUPPORT_LIST) &&
80- CheckType(output->GetDataType(), OUTPUT_DTYPE_SUPPORT_LIST)) {
81- return true;
82- } else {
83- return false;
84- }
85-}
86- 
87-static bool CheckAllDtypesValid(const aclTensor *x, const aclTensor *scales, const aclTensor *output)
88-{
89- bool isAllDtypesValid = false;
90- isAllDtypesValid = CheckKGAllDtypesValid(x, scales, output) || CheckMXAllDtypesValid(x, scales, output);
91- if (!isAllDtypesValid) {
92- OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON("aclnnQuantReduceScatter", "x/scales/output",
93- (std::string(op::ToString(x->GetDataType()).GetString()) + "/" +
94- op::ToString(scales->GetDataType()).GetString() + "/" +
95- op::ToString(output->GetDataType()).GetString())
96- .c_str(),
97- "The dtypes of x, scales and output must be valid");
98- }
99- return isAllDtypesValid;
100-}
101- 
102-static bool CheckGroupLength(const char *group)
103-{
104- if (group == nullptr) {
105- OP_LOGE_WITH_INVALID_INPUT("aclnnQuantReduceScatter", "group");
106- return false;
107- }
108- 
109- size_t groupLen = strnlen(group, HCCL_GROUP_NAME_LENGTH_MAX); // group长度≥128字符, 返回HCCL_GROUP_NAME_LENGTH_MAX
110- if (groupLen >= HCCL_GROUP_NAME_LENGTH_MAX) {
111- OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(
112- "aclnnQuantReduceScatter", "group", "length exceeds " + std::to_string(HCCL_GROUP_NAME_LENGTH_MAX),
113- "The length of group must be less than " + std::to_string(HCCL_GROUP_NAME_LENGTH_MAX) + " characters");
114- return false;
115- }
116- 
117- return true;
118-}
119- 
120-static aclnnStatus CheckParams(const aclTensor *x, const aclTensor *scales, const char *group, const aclTensor *output)
121-{
122- // 1. 检查参数是否为空指针
123- CHECK_RET(CheckNotNull(x, scales, output), ACLNN_ERR_PARAM_NULLPTR);
124- // 2. 检查输入的数据类型是否在API支持的数据类型范围之内,需要根据api定义校验
125- CHECK_RET(CheckAllDtypesValid(x, scales, output), ACLNN_ERR_PARAM_INVALID);
126- // 3. 检查group参数是否在要求范围之内
127- CHECK_RET(CheckGroupLength(group), ACLNN_ERR_PARAM_INVALID);
128- 
129- return ACLNN_SUCCESS;
130-}
131-} // namespace
132- 
133-extern "C" aclnnStatus aclnnInnerQuantReduceScatterGetWorkspaceSize(const aclTensor *x, const aclTensor *scales,
134- const char *group, const char *reduceOp,
135- uint64_t yDtype, int64_t worldSize,
136- aclTensor *output, uint64_t *workspaceSize,
137- aclOpExecutor **executor);
138-extern "C" aclnnStatus aclnnInnerQuantReduceScatter(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor,
139- const aclrtStream stream);
140-extern "C" void __attribute__((weak)) NnopbaseSetHcclServerType(void *executor, NnopbaseHcclServerType sType);
141- 
142extern "C" aclnnStatus aclnnQuantReduceScatterGetWorkspaceSize(const aclTensor *x, const aclTensor *scales,30extern "C" aclnnStatus aclnnQuantReduceScatterGetWorkspaceSize(const aclTensor *x, const aclTensor *scales,
143 const char *group, const char *reduceOp,31 const char *group, const char *reduceOp,
144 aclTensor *output, uint64_t *workspaceSize,32 aclTensor *output, uint64_t *workspaceSize,
145 aclOpExecutor **executor)33 aclOpExecutor **executor)
146{34{
147- aclnnStatus retParam = CheckParams(x, scales, group, output);35+ return ACLNN_SUCCESS;
148- CHECK_RET(retParam == ACLNN_SUCCESS, retParam);
149- uint64_t yDtype = static_cast<uint64_t>(output->GetDataType());
150- int64_t worldSize = -1;
151- aclnnStatus ret =
152- aclnnInnerQuantReduceScatterGetWorkspaceSize(x, scales, const_cast<char *>(group), const_cast<char *>(reduceOp),
153- yDtype, worldSize, output, workspaceSize, executor);
154- OP_LOGD("QuantReduceScatter, aclnnnGetWorkspaceSize ret %d.", ret);
155- return ret;
156}36}
157 37 
158extern "C" aclnnStatus aclnnQuantReduceScatter(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor,38extern "C" aclnnStatus aclnnQuantReduceScatter(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor,
159 const aclrtStream stream)39 const aclrtStream stream)
160{40{
161- if (NnopbaseSetHcclServerType) {
162- NnopbaseSetHcclServerType(executor, NnopbaseHcclServerType::NNOPBASE_HCCL_SERVER_TYPE_MTE);
163- }
164- aclnnStatus ret = aclnnInnerQuantReduceScatter(workspace, workspaceSize, executor, stream);
165- if (ret != ACLNN_SUCCESS) {
166- OP_LOGE_LIBOPAPI_REPORT("aclnnQuantReduceScatter", "This is an error in launch aicore");
167- return ACLNN_ERR_INNER;
168- }
169 return ACLNN_SUCCESS;41 return ACLNN_SUCCESS;
170}42}
@@ -8,543 +8,7 @@
8 * See LICENSE in the root of the software repository for the full text of the License.8 * See LICENSE in the root of the software repository for the full text of the License.
9 */9 */
10 10 
11-#include <float.h>11+/*!
12-#include <array>12+ * \file test_aclnn_quant_reduce_scatter.cpp
13-#include <vector>13+ * \brief aclnn ut
14-#include "gtest/gtest.h"14+ */
15-#include <gmock/gmock.h>
16-#include "../../../op_api/aclnn_quant_reduce_scatter.h"
17-#include "op_api_ut_common/tensor_desc.h"
18-#include "op_api_ut_common/op_api_ut.h"
19-#include "opdev/platform.h"
20- 
21-using namespace op;
22-using namespace std;
23- 
24-class TestAclnnQuantReduceScatter : public testing::Test {
25-protected:
26- static void SetUpTestCase()
27- {
28- op::SetPlatformNpuArch(NpuArch::DAV_3510);
29- cout << "TestAclnnQuantReduceScatter SetUp" << endl;
30- }
31- 
32- static void TearDownTestCase()
33- {
34- op::SetPlatformSocVersion(op::SocVersion::ASCEND910B);
35- cout << "TestAclnnQuantReduceScatter TearDown" << endl;
36- }
37-};
38- 
39-struct QuantReduceScatterAclnnTestParam {
40- string case_name;
41- vector<int64_t> xShape; // x数据shape
42- vector<int64_t> scalesShape; // scales数据shape
43- vector<int64_t> outputShape; // output数据shape
44- char *group; // 通信域标识
45- aclDataType xDtype; // x数据dtype
46- aclDataType scalesDtype; // scales数据dtype
47- aclDataType outputDtype; // 输出数据dtype
48- aclnnStatus aclnnStatusUt;
49-};
50- 
51-static QuantReduceScatterAclnnTestParam g_casesParams[] = {
52- // 正常用例
53- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H5120_FLOAT8E4M3FN_FLOAT8E8M0_FLOAT16_true",
54- {1024, 5120},
55- {1024, 80, 2},
56- {512, 5120},
57- "quant_reduce_scatter_test_group",
58- ACL_FLOAT8_E4M3FN,
59- ACL_FLOAT8_E8M0,
60- ACL_FLOAT16,
61- ACLNN_SUCCESS},
62- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H5120_FLOAT8E5M2_FLOAT8E8M0_BF16_true",
63- {1024, 5120},
64- {1024, 80, 2},
65- {512, 5120},
66- "quant_reduce_scatter_test_group",
67- ACL_FLOAT8_E5M2,
68- ACL_FLOAT8_E8M0,
69- ACL_BF16,
70- ACLNN_SUCCESS},
71- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H7168_FLOAT8E4M3FN_FLOAT8E8M0_FLOAT16_true",
72- {1024, 7168},
73- {1024, 112, 2},
74- {512, 7168},
75- "quant_reduce_scatter_test_group",
76- ACL_FLOAT8_E4M3FN,
77- ACL_FLOAT8_E8M0,
78- ACL_FLOAT16,
79- ACLNN_SUCCESS},
80- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H7168_FLOAT8E5M2_FLOAT8E8M0_BF16_true",
81- {1024, 7168},
82- {1024, 112, 2},
83- {512, 7168},
84- "quant_reduce_scatter_test_group",
85- ACL_FLOAT8_E5M2,
86- ACL_FLOAT8_E8M0,
87- ACL_BF16,
88- ACLNN_SUCCESS},
89- 
90- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H5120_INT8_FLOAT_FLOAT16_true",
91- {1024, 5120},
92- {1024, 40},
93- {512, 5120},
94- "quant_reduce_scatter_test_group",
95- ACL_INT8,
96- ACL_FLOAT,
97- ACL_FLOAT16,
98- ACLNN_SUCCESS},
99- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H5120_HIFLOAT8_FLOAT_BF16_true",
100- {1024, 5120},
101- {1024, 40},
102- {512, 5120},
103- "quant_reduce_scatter_test_group",
104- ACL_HIFLOAT8,
105- ACL_FLOAT,
106- ACL_BF16,
107- ACLNN_SUCCESS},
108- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H5120_FLOAT8E4M3FN_FLOAT_FLOAT_true",
109- {1024, 5120},
110- {1024, 40},
111- {512, 5120},
112- "quant_reduce_scatter_test_group",
113- ACL_FLOAT8_E4M3FN,
114- ACL_FLOAT,
115- ACL_FLOAT,
116- ACLNN_SUCCESS},
117- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H5120_FLOAT8E5M2_FLOAT_FLOAT_true",
118- {1024, 5120},
119- {1024, 40},
120- {512, 5120},
121- "quant_reduce_scatter_test_group",
122- ACL_FLOAT8_E5M2,
123- ACL_FLOAT,
124- ACL_FLOAT,
125- ACLNN_SUCCESS},
126- 
127- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H7168_INT8_FLOAT_FLOAT16_true",
128- {1024, 7168},
129- {1024, 56},
130- {512, 7168},
131- "quant_reduce_scatter_test_group",
132- ACL_INT8,
133- ACL_FLOAT,
134- ACL_FLOAT16,
135- ACLNN_SUCCESS},
136- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H7168_HIFLOAT8_FLOAT_BF16_true",
137- {1024, 7168},
138- {1024, 56},
139- {512, 7168},
140- "quant_reduce_scatter_test_group",
141- ACL_HIFLOAT8,
142- ACL_FLOAT,
143- ACL_BF16,
144- ACLNN_SUCCESS},
145- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H7168_FLOAT8E4M3FN_FLOAT_FLOAT_true",
146- {1024, 7168},
147- {1024, 56},
148- {512, 7168},
149- "quant_reduce_scatter_test_group",
150- ACL_FLOAT8_E4M3FN,
151- ACL_FLOAT,
152- ACL_FLOAT,
153- ACLNN_SUCCESS},
154- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H7168_FLOAT8E5M2_FLOAT_FLOAT_true",
155- {1024, 7168},
156- {1024, 56},
157- {512, 7168},
158- "quant_reduce_scatter_test_group",
159- ACL_FLOAT8_E5M2,
160- ACL_FLOAT,
161- ACL_FLOAT,
162- ACLNN_SUCCESS},
163- {"test_aclnn_quant_reduce_scatter_kg_BS2048_H5120_INT8_FLOAT_FLOAT16_true",
164- {2048, 5120},
165- {2048, 40},
166- {1024, 5120},
167- "quant_reduce_scatter_test_group",
168- ACL_INT8,
169- ACL_FLOAT,
170- ACL_FLOAT16,
171- ACLNN_SUCCESS},
172- {"test_aclnn_quant_reduce_scatter_kg_BS2048_H5120_HIFLOAT8_FLOAT_BF16_true",
173- {2048, 5120},
174- {2048, 40},
175- {1024, 5120},
176- "quant_reduce_scatter_test_group",
177- ACL_HIFLOAT8,
178- ACL_FLOAT,
179- ACL_BF16,
180- ACLNN_SUCCESS},
181- {"test_aclnn_quant_reduce_scatter_mx_BS2048_H7168_FLOAT8E5M2_FLOAT8E8M0_BF16_true",
182- {2048, 7168},
183- {2048, 112, 2},
184- {1024, 7168},
185- "quant_reduce_scatter_test_group",
186- ACL_FLOAT8_E5M2,
187- ACL_FLOAT8_E8M0,
188- ACL_BF16,
189- ACLNN_SUCCESS},
190- {"test_aclnn_quant_reduce_scatter_kg_BS4096_H7168_INT8_FLOAT_FLOAT16_true",
191- {4096, 7168},
192- {4096, 56},
193- {2048, 7168},
194- "quant_reduce_scatter_test_group",
195- ACL_INT8,
196- ACL_FLOAT,
197- ACL_FLOAT16,
198- ACLNN_SUCCESS},
199- {"test_aclnn_quant_reduce_scatter_kg_BS8192_H5120_FLOAT8E4M3FN_FLOAT_FLOAT_true",
200- {8192, 5120},
201- {8192, 40},
202- {4096, 5120},
203- "quant_reduce_scatter_test_group",
204- ACL_FLOAT8_E4M3FN,
205- ACL_FLOAT,
206- ACL_FLOAT,
207- ACLNN_SUCCESS},
208- 
209- // x类型异常
210- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H5120_INT8_FLOAT8E8M0_FLOAT16_xDtype_false",
211- {1024, 5120},
212- {1024, 80, 2},
213- {512, 5120},
214- "quant_reduce_scatter_test_group",
215- ACL_INT8,
216- ACL_FLOAT8_E8M0,
217- ACL_FLOAT16,
218- ACLNN_ERR_PARAM_INVALID}, // mx量化时,x不应该为INT8,HIFLOAT8
219- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H7168_INT8_FLOAT8E8M0_FLOAT16_xDtype_false",
220- {1024, 7168},
221- {1024, 112, 2},
222- {512, 7168},
223- "quant_reduce_scatter_test_group",
224- ACL_INT8,
225- ACL_FLOAT8_E8M0,
226- ACL_FLOAT16,
227- ACLNN_ERR_PARAM_INVALID},
228- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H5120_HIFLOAT8_FLOAT8E8M0_FLOAT16_xDtype_false",
229- {1024, 5120},
230- {1024, 80, 2},
231- {512, 5120},
232- "quant_reduce_scatter_test_group",
233- ACL_HIFLOAT8,
234- ACL_FLOAT8_E8M0,
235- ACL_FLOAT16,
236- ACLNN_ERR_PARAM_INVALID},
237- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H7168_HIFLOAT8_FLOAT8E8M0_FLOAT16_xDtype_false",
238- {1024, 7168},
239- {1024, 112, 2},
240- {512, 7168},
241- "quant_reduce_scatter_test_group",
242- ACL_HIFLOAT8,
243- ACL_FLOAT8_E8M0,
244- ACL_FLOAT16,
245- ACLNN_ERR_PARAM_INVALID},
246- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H5120_FLOAT16_FLOAT_FLOAT_xDtype_false",
247- {1024, 5120},
248- {1024, 40},
249- {512, 5120},
250- "quant_reduce_scatter_test_group",
251- ACL_FLOAT16,
252- ACL_FLOAT,
253- ACL_FLOAT,
254- ACLNN_ERR_PARAM_INVALID}, // K-G量化时,x不应该为FLOAT16
255- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H7168_FLOAT16_FLOAT_FLOAT_xDtype_false",
256- {1024, 7168},
257- {1024, 56},
258- {512, 7168},
259- "quant_reduce_scatter_test_group",
260- ACL_FLOAT16,
261- ACL_FLOAT,
262- ACL_FLOAT,
263- ACLNN_ERR_PARAM_INVALID},
264- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H5120_FLOAT_FLOAT_FLOAT16_xDtype_false",
265- {1024, 5120},
266- {1024, 80, 2},
267- {512, 5120},
268- "quant_reduce_scatter_test_group",
269- ACL_FLOAT,
270- ACL_FLOAT,
271- ACL_FLOAT16,
272- ACLNN_ERR_PARAM_INVALID},
273- 
274- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H5120_FLOAT_FLOAT8E8M0_FLOAT16_xDtype_false",
275- {1024, 5120},
276- {1024, 80, 2},
277- {512, 5120},
278- "quant_reduce_scatter_test_group",
279- ACL_FLOAT,
280- ACL_FLOAT8_E8M0,
281- ACL_FLOAT16,
282- ACLNN_ERR_PARAM_INVALID}, // x不应该为FLOAT
283- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H7168_FLOAT_FLOAT8E8M0_FLOAT16_xDtype_false",
284- {1024, 7168},
285- {1024, 112, 2},
286- {512, 7168},
287- "quant_reduce_scatter_test_group",
288- ACL_FLOAT,
289- ACL_FLOAT8_E8M0,
290- ACL_FLOAT16,
291- ACLNN_ERR_PARAM_INVALID},
292- // scales类型异常
293- {"test_aclnn_quant_reduce_scatter_BS1024_H5120_INT8_FLOAT16_FLOAT16_scalesDtype_false",
294- {1024, 5120},
295- {1024, 40},
296- {512, 5120},
297- "quant_reduce_scatter_test_group",
298- ACL_INT8,
299- ACL_FLOAT16,
300- ACL_FLOAT16,
301- ACLNN_ERR_PARAM_INVALID}, // scales不支持FLOAT16
302- {"test_aclnn_quant_reduce_scatter_BS1024_H5120_HIFLOAT8_FLOAT16_BF16_scalesDtype_false",
303- {1024, 5120},
304- {1024, 40},
305- {512, 5120},
306- "quant_reduce_scatter_test_group",
307- ACL_HIFLOAT8,
308- ACL_FLOAT16,
309- ACL_BF16,
310- ACLNN_ERR_PARAM_INVALID},
311- {"test_aclnn_quant_reduce_scatter_BS1024_H5120_FLOAT8E4M3FN_FLOAT16_FLOAT_scalesDtype_false",
312- {1024, 5120},
313- {1024, 40},
314- {512, 5120},
315- "quant_reduce_scatter_test_group",
316- ACL_FLOAT8_E4M3FN,
317- ACL_FLOAT16,
318- ACL_FLOAT,
319- ACLNN_ERR_PARAM_INVALID},
320- {"test_aclnn_quant_reduce_scatter_BS1024_H5120_FLOAT8E5M2_FLOAT16_FLOAT_scalesDtype_false",
321- {1024, 5120},
322- {1024, 40},
323- {512, 5120},
324- "quant_reduce_scatter_test_group",
325- ACL_FLOAT8_E5M2,
326- ACL_FLOAT16,
327- ACL_FLOAT,
328- ACLNN_ERR_PARAM_INVALID},
329- {"test_aclnn_quant_reduce_scatter_BS1024_H5120_FLOAT8E5M2_FLOAT8E4M3FN_FLOAT_scalesDtype_false",
330- {1024, 5120},
331- {1024, 40},
332- {512, 5120},
333- "quant_reduce_scatter_test_group",
334- ACL_FLOAT8_E5M2,
335- ACL_FLOAT8_E4M3FN,
336- ACL_FLOAT,
337- ACLNN_ERR_PARAM_INVALID}, // scales不支持FLOAT8_E4M3FN
338- 
339- {"test_aclnn_quant_reduce_scatter_BS1024_H7168_INT8_FLOAT16_FLOAT16_scalesDtype_false",
340- {1024, 7168},
341- {1024, 56},
342- {512, 7168},
343- "quant_reduce_scatter_test_group",
344- ACL_INT8,
345- ACL_FLOAT16,
346- ACL_FLOAT16,
347- ACLNN_ERR_PARAM_INVALID},
348- {"test_aclnn_quant_reduce_scatter_BS1024_H7168_HIFLOAT8_FLOAT16_FLOAT_scalesDtype_false",
349- {1024, 7168},
350- {1024, 56},
351- {512, 7168},
352- "quant_reduce_scatter_test_group",
353- ACL_HIFLOAT8,
354- ACL_FLOAT16,
355- ACL_BF16,
356- ACLNN_ERR_PARAM_INVALID},
357- {"test_aclnn_quant_reduce_scatter_BS1024_H7168_FLOAT8E4M3FN_FLOAT16_FLOAT_scalesDtype_false",
358- {1024, 7168},
359- {1024, 56},
360- {512, 7168},
361- "quant_reduce_scatter_test_group",
362- ACL_FLOAT8_E4M3FN,
363- ACL_FLOAT16,
364- ACL_FLOAT,
365- ACLNN_ERR_PARAM_INVALID},
366- {"test_aclnn_quant_reduce_scatter_BS1024_H7168_FLOAT8E5M2_FLOAT16_FLOAT_scalesDtype_false",
367- {1024, 7168},
368- {1024, 56},
369- {512, 7168},
370- "quant_reduce_scatter_test_group",
371- ACL_FLOAT8_E5M2,
372- ACL_FLOAT16,
373- ACL_FLOAT,
374- ACLNN_ERR_PARAM_INVALID},
375- // output类型异常
376- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H5120_INT8_FLOAT_FLOAT8E5M2_outputDtype_false",
377- {1024, 5120},
378- {1024, 40},
379- {512, 5120},
380- "quant_reduce_scatter_test_group",
381- ACL_INT8,
382- ACL_FLOAT,
383- ACL_FLOAT8_E5M2,
384- ACLNN_ERR_PARAM_INVALID}, // output不支持ACL_FLOAT8_E5M2
385- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H5120_HIFLOAT8_FLOAT_FLOAT8E5M2_outputDtype_false",
386- {1024, 5120},
387- {1024, 40},
388- {512, 5120},
389- "quant_reduce_scatter_test_group",
390- ACL_HIFLOAT8,
391- ACL_FLOAT,
392- ACL_FLOAT8_E5M2,
393- ACLNN_ERR_PARAM_INVALID},
394- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H5120_FLOAT8E4M3FN_FLOAT_FLOAT8E5M2_outputDtype_false",
395- {1024, 5120},
396- {1024, 40},
397- {512, 5120},
398- "quant_reduce_scatter_test_group",
399- ACL_FLOAT8_E4M3FN,
400- ACL_FLOAT,
401- ACL_FLOAT8_E5M2,
402- ACLNN_ERR_PARAM_INVALID},
403- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H5120_FLOAT8E5M2_FLOAT_FLOAT8E5M2_outputDtype_false",
404- {1024, 5120},
405- {1024, 40},
406- {512, 5120},
407- "quant_reduce_scatter_test_group",
408- ACL_FLOAT8_E5M2,
409- ACL_FLOAT,
410- ACL_FLOAT8_E5M2,
411- ACLNN_ERR_PARAM_INVALID},
412- 
413- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H7168_INT8_FLOAT_FLOAT8E5M2_outputDtype_false",
414- {1024, 7168},
415- {1024, 56},
416- {512, 7168},
417- "quant_reduce_scatter_test_group",
418- ACL_INT8,
419- ACL_FLOAT,
420- ACL_FLOAT8_E5M2,
421- ACLNN_ERR_PARAM_INVALID}, // output不支持ACL_FLOAT8_E5M2
422- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H7168_HIFLOAT8_FLOAT_FLOAT8E5M2_outputDtype_false",
423- {1024, 7168},
424- {1024, 56},
425- {512, 7168},
426- "quant_reduce_scatter_test_group",
427- ACL_HIFLOAT8,
428- ACL_FLOAT,
429- ACL_FLOAT8_E5M2,
430- ACLNN_ERR_PARAM_INVALID},
431- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H7168_FLOAT8E4M3FN_FLOAT_FLOAT8E5M2_outputDtype_false",
432- {1024, 7168},
433- {1024, 56},
434- {512, 7168},
435- "quant_reduce_scatter_test_group",
436- ACL_FLOAT8_E4M3FN,
437- ACL_FLOAT,
438- ACL_FLOAT8_E5M2,
439- ACLNN_ERR_PARAM_INVALID},
440- {"test_aclnn_quant_reduce_scatter_kg_BS1024_H7168_FLOAT8E5M2_FLOAT_FLOAT8E5M2_outputDtype_false",
441- {1024, 7168},
442- {1024, 56},
443- {512, 7168},
444- "quant_reduce_scatter_test_group",
445- ACL_FLOAT8_E5M2,
446- ACL_FLOAT,
447- ACL_FLOAT8_E5M2,
448- ACLNN_ERR_PARAM_INVALID},
449- {"test_aclnn_quant_reduce_scatter__BS1024_H5120_FLOAT8E4M3FN_FLOAT8E8M0_FLOAT8E5M2_outputDtype_false",
450- {1024, 5120},
451- {1024, 80, 2},
452- {512, 5120},
453- "quant_reduce_scatter_test_group",
454- ACL_FLOAT8_E4M3FN,
455- ACL_FLOAT8_E8M0,
456- ACL_FLOAT8_E5M2,
457- ACLNN_ERR_PARAM_INVALID}};
458- 
459-static QuantReduceScatterAclnnTestParam g_groupCasesParams[] = {
460- // group长度校验用例, this_is_a_very_long_groupname_长度为30字符
461- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H5120_FLOAT8E4M3FN_FLOAT8E8M0_FLOAT16_false_group_length_127",
462- {1024, 5120},
463- {1024, 80, 2},
464- {512, 5120},
465- "this_is_a_very_long_groupname_"
466- "this_is_a_very_long_groupname_"
467- "this_is_a_very_long_groupname_"
468- "this_is_a_very_long_groupname_"
469- "this_is",
470- ACL_FLOAT8_E4M3FN,
471- ACL_FLOAT8_E8M0,
472- ACL_FLOAT16,
473- ACLNN_SUCCESS}, // group长度为127
474- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H5120_FLOAT8E4M3FN_FLOAT8E8M0_FLOAT16_false_group_length_128",
475- {1024, 5120},
476- {1024, 80, 2},
477- {512, 5120},
478- "this_is_a_very_long_groupname_"
479- "this_is_a_very_long_groupname_"
480- "this_is_a_very_long_groupname_"
481- "this_is_a_very_long_groupname_"
482- "this_is_",
483- ACL_FLOAT8_E4M3FN,
484- ACL_FLOAT8_E8M0,
485- ACL_FLOAT16,
486- ACLNN_ERR_PARAM_INVALID}, // group长度为128-越界
487- {"test_aclnn_quant_reduce_scatter_mx_BS1024_H5120_FLOAT8E4M3FN_FLOAT8E8M0_FLOAT16_false_group_length_129",
488- {1024, 5120},
489- {1024, 80, 2},
490- {512, 5120},
491- "this_is_a_very_long_groupname_"
492- "this_is_a_very_long_groupname_"
493- "this_is_a_very_long_groupname_"
494- "this_is_a_very_long_groupname_"
495- "this_is_a",
496- ACL_FLOAT8_E4M3FN,
497- ACL_FLOAT8_E8M0,
498- ACL_FLOAT16,
499- ACLNN_ERR_PARAM_INVALID}, // group长度为129-越界
500-};
501- 
502-static void TestOneParamCase(const QuantReduceScatterAclnnTestParam &param)
503-{
504- std::cout << "run case " << param.case_name << std::endl;
505- if (param.group == nullptr) {
506- std::cerr << "[ERROR]: group is null" << std::endl;
507- return;
508- }
509- vector<int64_t> xShape = param.xShape;
510- vector<int64_t> scalesShape = param.scalesShape;
511- vector<int64_t> outputShape = param.outputShape;
512- char *group = param.group;
513- aclDataType xDtype = param.xDtype;
514- aclDataType scalesDtype = param.scalesDtype;
515- aclDataType outputDtype = param.outputDtype;
516- aclnnStatus retStatus = param.aclnnStatusUt;
517- TensorDesc x = TensorDesc(xShape, xDtype, ACL_FORMAT_ND);
518- TensorDesc scales = TensorDesc(scalesShape, scalesDtype, ACL_FORMAT_ND);
519- TensorDesc output = TensorDesc(outputShape, outputDtype, ACL_FORMAT_ND);
520- const char *reduceOp = "sum";
521- auto ut = OP_API_UT(aclnnQuantReduceScatter, INPUT(x, scales, group, reduceOp), OUTPUT(output));
522- uint64_t workspaceSize = 0;
523- aclOpExecutor *executor = nullptr;
524- aclnnStatus aclRet = ut.TestGetWorkspaceSizeWithNNopbaseInner(&workspaceSize, executor);
525- if (retStatus == ACLNN_SUCCESS) {
526- EXPECT_NE(aclRet, ACLNN_ERR_PARAM_INVALID);
527- } else {
528- EXPECT_EQ(aclRet, retStatus);
529- }
530-}
531- 
532-TEST_F(TestAclnnQuantReduceScatter, CasesParamsTest)
533-{
534- if (std::size(g_casesParams) != 0) {
535- uint64_t numCases = sizeof(g_casesParams) / sizeof(g_casesParams[0]);
536- for (size_t idx = 0; idx < numCases; idx += 1) {
537- TestOneParamCase(g_casesParams[idx]);
538- }
539- }
540-}
541- 
542-TEST_F(TestAclnnQuantReduceScatter, GroupCasesParamsTest)
543-{
544- if (std::size(g_groupCasesParams) != 0) {
545- uint64_t numCases = sizeof(g_groupCasesParams) / sizeof(g_groupCasesParams[0]);
546- for (size_t idx = 0; idx < numCases; idx += 1) {
547- TestOneParamCase(g_groupCasesParams[idx]);
548- }
549- }
550-}