已合并
delete quant_all_reduce & quant_reduce_scatter aclnn #12003
yifuxiong创建于 22 天前
delete quant_all_reduce & quant_reduce_scatter aclnn #12003
已合并
从已删除 :pr_920_beta2合入到cann/ops-transformer9.2.0-beta.2
共 12 个文件变更+9-3181
| @@ -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 | -# QuantAllReduce | 1 | +# 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 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - | ||
| 23 | -using namespace std; | ||
| 24 | - | ||
| 25 | - | ||
| 26 | - do { \ | ||
| 27 | - if (!(cond)) { \ | ||
| 28 | - return_expr; \ | ||
| 29 | - } \ | ||
| 30 | - } while (0) | ||
| 31 | - | ||
| 32 | - | ||
| 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 | 29 | ||
| 30 | 30 | ||
| 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 | - | ||
| 196 | extern "C" aclnnStatus aclnnQuantAllReduceGetWorkspaceSize(const aclTensor *x, const aclTensor *scales, | 32 | extern "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 | ||
| 211 | extern "C" aclnnStatus aclnnQuantAllReduce(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, | 39 | extern "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 | # aclnnQuantReduceScatter | 1 | # 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 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - | ||
| 23 | - | ||
| 24 | - | ||
| 25 | - do { \ | ||
| 26 | - if (!(cond)) { \ | ||
| 27 | - return_expr; \ | ||
| 28 | - } \ | ||
| 29 | - } while (0) | ||
| 30 | - | ||
| 31 | - | ||
| 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 | 27 | ||
| 28 | 28 | ||
| 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 | - | ||
| 142 | extern "C" aclnnStatus aclnnQuantReduceScatterGetWorkspaceSize(const aclTensor *x, const aclTensor *scales, | 30 | extern "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 | ||
| 158 | extern "C" aclnnStatus aclnnQuantReduceScatter(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, | 38 | extern "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 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 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 ¶m) | ||
| 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 | -} | ||