已合并
scatter撤销float8_e8m0数据类型 #3497
zhang-wenbo-beat创建于 4月2日
scatter撤销float8_e8m0数据类型 #3497
已合并
共 6 个文件变更+22-144
| @@ -99,7 +99,7 @@ aclnnStatus aclnnInplaceScatterUpdate( | |||
| 99 | <td>输入/输出</td> | 99 | <td>输入/输出</td> |
| 100 | <td>待更新的tensor。</td> | 100 | <td>待更新的tensor。</td> |
| 101 | <td>维度数需要与updates一致,不支持空Tensor。</td> | 101 | <td>维度数需要与updates一致,不支持空Tensor。</td> |
| 102 | - <td>INT8、UINT8、FLOAT16、FLOAT32、INT32、BFLOAT16、FLOAT8_E4M3FN、FLOAT8_E5M2、FLOAT8_E8M0、HIFLOAT8</td> | 102 | + <td>INT8、UINT8、FLOAT16、FLOAT32、INT32、BFLOAT16、FLOAT8_E4M3FN、FLOAT8_E5M2、HIFLOAT8</td> |
| 103 | <td>ND</td> | 103 | <td>ND</td> |
| 104 | <td>2-8</td> | 104 | <td>2-8</td> |
| 105 | <td>√</td> | 105 | <td>√</td> |
| @@ -119,7 +119,7 @@ aclnnStatus aclnnInplaceScatterUpdate( | |||
| 119 | <td>输入</td> | 119 | <td>输入</td> |
| 120 | <td>存储更新数据的tensor。</td> | 120 | <td>存储更新数据的tensor。</td> |
| 121 | <td>数据类型需要与data相同,shape的维度数需要与data shape的维度数相同。不支持空Tensor。</td> | 121 | <td>数据类型需要与data相同,shape的维度数需要与data shape的维度数相同。不支持空Tensor。</td> |
| 122 | - <td>INT8、UINT8、FLOAT16、FLOAT32、INT32、BFLOAT16、FLOAT8_E4M3FN、FLOAT8_E5M2、FLOAT8_E8M0、HIFLOAT8</td> | 122 | + <td>INT8、UINT8、FLOAT16、FLOAT32、INT32、BFLOAT16、FLOAT8_E4M3FN、FLOAT8_E5M2、HIFLOAT8</td> |
| 123 | <td>ND</td> | 123 | <td>ND</td> |
| 124 | <td>2-8</td> | 124 | <td>2-8</td> |
| 125 | <td>√</td> | 125 | <td>√</td> |
| @@ -157,7 +157,7 @@ aclnnStatus aclnnInplaceScatterUpdate( | |||
| 157 | </tbody></table> | 157 | </tbody></table> |
| 158 | 158 | ||
| 159 | - <term>Atlas 训练系列产品</term>:数据类型不支持UINT8、BFLOAT16。 | 159 | - <term>Atlas 训练系列产品</term>:数据类型不支持UINT8、BFLOAT16。 |
| 160 | - - <term>Ascend 950PR/Ascend 950DT</term>:FLOAT8_E4M3FN、FLOAT8_E5M2、FLOAT8_E8M0、HIFLOAT8等数据类型仅在该型号支持。 | 160 | + - <term>Ascend 950PR/Ascend 950DT</term>:FLOAT8_E4M3FN、FLOAT8_E5M2、HIFLOAT8等数据类型仅在该型号支持。 |
| 161 | 161 | ||
| 162 | - **返回值** | 162 | - **返回值** |
| 163 | 163 | ||
| @@ -37,10 +37,10 @@ static const std::initializer_list<op::DataType> DTYPE_SUPPORT_LIST_DATA = { | |||
| 37 | op::DataType::DT_INT32}; | 37 | op::DataType::DT_INT32}; |
| 38 | 38 | ||
| 39 | static const std::initializer_list<op::DataType> ASCEND950_DTYPE_SUPPORT_LIST_DATA = { | 39 | static const std::initializer_list<op::DataType> ASCEND950_DTYPE_SUPPORT_LIST_DATA = { |
| 40 | - op::DataType::DT_INT8, op::DataType::DT_FLOAT16, op::DataType::DT_FLOAT, | 40 | + op::DataType::DT_INT8, op::DataType::DT_FLOAT16, op::DataType::DT_FLOAT, |
| 41 | - op::DataType::DT_BF16, op::DataType::DT_UINT8, op::DataType::DT_INT32, | 41 | + op::DataType::DT_BF16, op::DataType::DT_UINT8, op::DataType::DT_INT32, |
| 42 | - op::DataType::DT_FLOAT8_E4M3FN, op::DataType::DT_FLOAT8_E5M2, | 42 | + op::DataType::DT_FLOAT8_E4M3FN, op::DataType::DT_FLOAT8_E5M2, |
| 43 | - op::DataType::DT_FLOAT8_E8M0, op::DataType::DT_HIFLOAT8}; | 43 | + op::DataType::DT_HIFLOAT8}; |
| 44 | 44 | ||
| 45 | static const std::initializer_list<op::DataType> DTYPE_SUPPORT_LIST_INDICES = {op::DataType::DT_INT64, | 45 | static const std::initializer_list<op::DataType> DTYPE_SUPPORT_LIST_INDICES = {op::DataType::DT_INT64, |
| 46 | op::DataType::DT_INT32}; | 46 | op::DataType::DT_INT32}; |
| @@ -25,7 +25,7 @@ namespace ge { | |||
| 25 | 25 | ||
| 26 | * @par Inputs: | 26 | * @par Inputs: |
| 27 | * @li var: The rewritten tensor. Format is ND. Support 2D ~ 8D, when axis is -1 the last dim of var should be 32B align. | 27 | * @li var: The rewritten tensor. Format is ND. Support 2D ~ 8D, when axis is -1 the last dim of var should be 32B align. |
| 28 | -* Must be one of the following types:float16, float32, int32, int8, uint8, bfloat16, float8_e4m3fn, float8_e5m2, float8_e8m0, hifloat8. | 28 | +* Must be one of the following types:float16, float32, int32, int8, uint8, bfloat16, float8_e4m3fn, float8_e5m2, hifloat8. |
| 29 | * @li indices: The index tensor. Format is ND. Support 1D ~ 2D, when discrete, 1-dim of indices should be 2. | 29 | * @li indices: The index tensor. Format is ND. Support 1D ~ 2D, when discrete, 1-dim of indices should be 2. |
| 30 | * Must be one of the following types: int32, int64. | 30 | * Must be one of the following types: int32, int64. |
| 31 | * Index out of bounds is not supported. | 31 | * Index out of bounds is not supported. |
| @@ -44,10 +44,10 @@ namespace ge { | |||
| 44 | * Compatible with the Mindspore operator Scatter. | 44 | * Compatible with the Mindspore operator Scatter. |
| 45 | */ | 45 | */ |
| 46 | REG_OP(Scatter) | 46 | REG_OP(Scatter) |
| 47 | - .INPUT(var, TensorType({DT_FLOAT16, DT_FLOAT, DT_INT32, DT_INT8, DT_UINT8, DT_BF16, DT_FLOAT8_E4M3FN, DT_FLOAT8_E5M2, DT_FLOAT8_E8M0, DT_HIFLOAT8})) | 47 | + .INPUT(var, TensorType({DT_FLOAT16, DT_FLOAT, DT_INT32, DT_INT8, DT_UINT8, DT_BF16, DT_FLOAT8_E4M3FN, DT_FLOAT8_E5M2, DT_HIFLOAT8})) |
| 48 | .INPUT(indices, TensorType::IndexNumberType()) | 48 | .INPUT(indices, TensorType::IndexNumberType()) |
| 49 | - .INPUT(updates, TensorType({DT_FLOAT16, DT_FLOAT, DT_INT32, DT_INT8, DT_UINT8, DT_BF16, DT_FLOAT8_E4M3FN, DT_FLOAT8_E5M2, DT_FLOAT8_E8M0, DT_HIFLOAT8})) | 49 | + .INPUT(updates, TensorType({DT_FLOAT16, DT_FLOAT, DT_INT32, DT_INT8, DT_UINT8, DT_BF16, DT_FLOAT8_E4M3FN, DT_FLOAT8_E5M2, DT_HIFLOAT8})) |
| 50 | - .OUTPUT(var, TensorType({DT_FLOAT16, DT_FLOAT, DT_INT32, DT_INT8, DT_UINT8, DT_BF16, DT_FLOAT8_E4M3FN, DT_FLOAT8_E5M2, DT_FLOAT8_E8M0, DT_HIFLOAT8})) | 50 | + .OUTPUT(var, TensorType({DT_FLOAT16, DT_FLOAT, DT_INT32, DT_INT8, DT_UINT8, DT_BF16, DT_FLOAT8_E4M3FN, DT_FLOAT8_E5M2, DT_HIFLOAT8})) |
| 51 | .REQUIRED_ATTR(reduce, String) | 51 | .REQUIRED_ATTR(reduce, String) |
| 52 | .ATTR(axis, Int, -1) | 52 | .ATTR(axis, Int, -1) |
| 53 | .OP_END_FACTORY_REG(Scatter) | 53 | .OP_END_FACTORY_REG(Scatter) |
| @@ -89,7 +89,7 @@ constexpr int64_t IN_DTYPE_B64 = 8; | |||
| 89 | 89 | ||
| 90 | static map<const ge::DataType, const int32_t> g_dtypeLen = {{ge::DT_INT8, 1}, {ge::DT_UINT8, 1}, {ge::DT_FLOAT16, 2}, | 90 | static map<const ge::DataType, const int32_t> g_dtypeLen = {{ge::DT_INT8, 1}, {ge::DT_UINT8, 1}, {ge::DT_FLOAT16, 2}, |
| 91 | {ge::DT_FLOAT, 4}, {ge::DT_INT32, 4}, {ge::DT_BF16, 2}, | 91 | {ge::DT_FLOAT, 4}, {ge::DT_INT32, 4}, {ge::DT_BF16, 2}, |
| 92 | - {ge::DT_FLOAT8_E4M3FN, 1}, {ge::DT_FLOAT8_E5M2, 1}, {ge::DT_FLOAT8_E8M0, 1}, | 92 | + {ge::DT_FLOAT8_E4M3FN, 1}, {ge::DT_FLOAT8_E5M2, 1}, |
| 93 | {ge::DT_HIFLOAT8, 1}}; | 93 | {ge::DT_HIFLOAT8, 1}}; |
| 94 | 94 | ||
| 95 | std::map<std::tuple<bool, ge::DataType, ge::DataType>, int32_t> tilingKeyMap; | 95 | std::map<std::tuple<bool, ge::DataType, ge::DataType>, int32_t> tilingKeyMap; |
| @@ -121,7 +121,7 @@ ge::graphStatus ScatterTiling::GetShapeAttrsInfo() { | |||
| 121 | inputDtype != ge::DataType::DT_BF16 && inputDtype != ge::DataType::DT_INT8 && | 121 | inputDtype != ge::DataType::DT_BF16 && inputDtype != ge::DataType::DT_INT8 && |
| 122 | inputDtype != ge::DataType::DT_UINT8 && inputDtype != ge::DataType::DT_INT32 && | 122 | inputDtype != ge::DataType::DT_UINT8 && inputDtype != ge::DataType::DT_INT32 && |
| 123 | inputDtype != ge::DataType::DT_FLOAT8_E4M3FN && inputDtype != ge::DataType::DT_FLOAT8_E5M2 && | 123 | inputDtype != ge::DataType::DT_FLOAT8_E4M3FN && inputDtype != ge::DataType::DT_FLOAT8_E5M2 && |
| 124 | - inputDtype != ge::DataType::DT_FLOAT8_E8M0 && inputDtype != ge::DataType::DT_HIFLOAT8) { | 124 | + inputDtype != ge::DataType::DT_HIFLOAT8) { |
| 125 | OP_LOGE("Scatter", "invalid input dtype."); | 125 | OP_LOGE("Scatter", "invalid input dtype."); |
| 126 | return ge::GRAPH_FAILED; | 126 | return ge::GRAPH_FAILED; |
| 127 | } | 127 | } |
| @@ -141,7 +141,7 @@ ge::graphStatus ScatterTiling::GetShapeAttrsInfo() { | |||
| 141 | updatesDtype != ge::DataType::DT_BF16 && updatesDtype != ge::DataType::DT_INT8 && | 141 | updatesDtype != ge::DataType::DT_BF16 && updatesDtype != ge::DataType::DT_INT8 && |
| 142 | updatesDtype != ge::DataType::DT_UINT8 && updatesDtype != ge::DataType::DT_INT32 && | 142 | updatesDtype != ge::DataType::DT_UINT8 && updatesDtype != ge::DataType::DT_INT32 && |
| 143 | updatesDtype != ge::DataType::DT_FLOAT8_E4M3FN && updatesDtype != ge::DataType::DT_FLOAT8_E5M2 && | 143 | updatesDtype != ge::DataType::DT_FLOAT8_E4M3FN && updatesDtype != ge::DataType::DT_FLOAT8_E5M2 && |
| 144 | - updatesDtype != ge::DataType::DT_FLOAT8_E8M0 && updatesDtype != ge::DataType::DT_HIFLOAT8) { | 144 | + updatesDtype != ge::DataType::DT_HIFLOAT8) { |
| 145 | OP_LOGE("Scatter", "invalid updates dtype."); | 145 | OP_LOGE("Scatter", "invalid updates dtype."); |
| 146 | return ge::GRAPH_FAILED; | 146 | return ge::GRAPH_FAILED; |
| 147 | } | 147 | } |
| @@ -468,7 +468,6 @@ void ScatterTiling::InitTilingKeyMap() { | |||
| 468 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT32, ge::DataType::DT_INT8)] = TILING_KEY_INT32_INT8; | 468 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT32, ge::DataType::DT_INT8)] = TILING_KEY_INT32_INT8; |
| 469 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT8_E5M2)] = TILING_KEY_INT32_INT8; | 469 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT8_E5M2)] = TILING_KEY_INT32_INT8; |
| 470 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT8_E4M3FN)] = TILING_KEY_INT32_INT8; | 470 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT8_E4M3FN)] = TILING_KEY_INT32_INT8; |
| 471 | - tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT8_E8M0)] = TILING_KEY_INT32_INT8; | ||
| 472 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT32, ge::DataType::DT_HIFLOAT8)] = TILING_KEY_INT32_INT8; | 471 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT32, ge::DataType::DT_HIFLOAT8)] = TILING_KEY_INT32_INT8; |
| 473 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT32, ge::DataType::DT_UINT8)] = TILING_KEY_INT32_UINT8; | 472 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT32, ge::DataType::DT_UINT8)] = TILING_KEY_INT32_UINT8; |
| 474 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT16)] = TILING_KEY_INT32_FLOAT16; | 473 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT16)] = TILING_KEY_INT32_FLOAT16; |
| @@ -479,7 +478,6 @@ void ScatterTiling::InitTilingKeyMap() { | |||
| 479 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT64, ge::DataType::DT_INT8)] = TILING_KEY_INT64_INT8; | 478 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT64, ge::DataType::DT_INT8)] = TILING_KEY_INT64_INT8; |
| 480 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT8_E5M2)] = TILING_KEY_INT64_INT8; | 479 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT8_E5M2)] = TILING_KEY_INT64_INT8; |
| 481 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT8_E4M3FN)] = TILING_KEY_INT64_INT8; | 480 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT8_E4M3FN)] = TILING_KEY_INT64_INT8; |
| 482 | - tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT8_E8M0)] = TILING_KEY_INT64_INT8; | ||
| 483 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT64, ge::DataType::DT_HIFLOAT8)] = TILING_KEY_INT64_INT8; | 481 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT64, ge::DataType::DT_HIFLOAT8)] = TILING_KEY_INT64_INT8; |
| 484 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT64, ge::DataType::DT_UINT8)] = TILING_KEY_INT64_UINT8; | 482 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT64, ge::DataType::DT_UINT8)] = TILING_KEY_INT64_UINT8; |
| 485 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT16)] = TILING_KEY_INT64_FLOAT16; | 483 | tilingKeyMap[std::make_tuple(false, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT16)] = TILING_KEY_INT64_FLOAT16; |
| @@ -490,7 +488,6 @@ void ScatterTiling::InitTilingKeyMap() { | |||
| 490 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT32, ge::DataType::DT_INT8)] = TILING_KEY_UINT64_INT32_INT8; | 488 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT32, ge::DataType::DT_INT8)] = TILING_KEY_UINT64_INT32_INT8; |
| 491 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT8_E5M2)] = TILING_KEY_UINT64_INT32_INT8; | 489 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT8_E5M2)] = TILING_KEY_UINT64_INT32_INT8; |
| 492 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT8_E4M3FN)] = TILING_KEY_UINT64_INT32_INT8; | 490 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT8_E4M3FN)] = TILING_KEY_UINT64_INT32_INT8; |
| 493 | - tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT8_E8M0)] = TILING_KEY_UINT64_INT32_INT8; | ||
| 494 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT32, ge::DataType::DT_HIFLOAT8)] = TILING_KEY_UINT64_INT32_INT8; | 491 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT32, ge::DataType::DT_HIFLOAT8)] = TILING_KEY_UINT64_INT32_INT8; |
| 495 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT32, ge::DataType::DT_UINT8)] = TILING_KEY_UINT64_INT32_UINT8; | 492 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT32, ge::DataType::DT_UINT8)] = TILING_KEY_UINT64_INT32_UINT8; |
| 496 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT16)] = TILING_KEY_UINT64_INT32_FLOAT16; | 493 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT32, ge::DataType::DT_FLOAT16)] = TILING_KEY_UINT64_INT32_FLOAT16; |
| @@ -501,7 +498,6 @@ void ScatterTiling::InitTilingKeyMap() { | |||
| 501 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT64, ge::DataType::DT_INT8)] = TILING_KEY_UINT64_INT64_INT8; | 498 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT64, ge::DataType::DT_INT8)] = TILING_KEY_UINT64_INT64_INT8; |
| 502 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT8_E5M2)] = TILING_KEY_UINT64_INT64_INT8; | 499 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT8_E5M2)] = TILING_KEY_UINT64_INT64_INT8; |
| 503 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT8_E4M3FN)] = TILING_KEY_UINT64_INT64_INT8; | 500 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT8_E4M3FN)] = TILING_KEY_UINT64_INT64_INT8; |
| 504 | - tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT8_E8M0)] = TILING_KEY_UINT64_INT64_INT8; | ||
| 505 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT64, ge::DataType::DT_HIFLOAT8)] = TILING_KEY_UINT64_INT64_INT8; | 501 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT64, ge::DataType::DT_HIFLOAT8)] = TILING_KEY_UINT64_INT64_INT8; |
| 506 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT64, ge::DataType::DT_UINT8)] = TILING_KEY_UINT64_INT64_UINT8; | 502 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT64, ge::DataType::DT_UINT8)] = TILING_KEY_UINT64_INT64_UINT8; |
| 507 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT16)] = TILING_KEY_UINT64_INT64_FLOAT16; | 503 | tilingKeyMap[std::make_tuple(true, ge::DataType::DT_INT64, ge::DataType::DT_FLOAT16)] = TILING_KEY_UINT64_INT64_FLOAT16; |
| @@ -827,65 +827,6 @@ | |||
| 827 | } | 827 | } |
| 828 | ] | 828 | ] |
| 829 | }, | 829 | }, |
| 830 | - { | ||
| 831 | - "bin_filename": "Scatter_1f3c60ecb6aac159qsmae360606f7e01", | ||
| 832 | - "inputs": [ | ||
| 833 | - { | ||
| 834 | - "name": "var", | ||
| 835 | - "index": 0, | ||
| 836 | - "dtype": "float8_e8m0", | ||
| 837 | - "format": "ND", | ||
| 838 | - "paramType": "required", | ||
| 839 | - "shape": [ | ||
| 840 | - -2 | ||
| 841 | - ] | ||
| 842 | - }, | ||
| 843 | - { | ||
| 844 | - "name": "indices", | ||
| 845 | - "index": 1, | ||
| 846 | - "dtype": "int64", | ||
| 847 | - "format": "ND", | ||
| 848 | - "paramType": "required", | ||
| 849 | - "shape": [ | ||
| 850 | - -2 | ||
| 851 | - ] | ||
| 852 | - }, | ||
| 853 | - { | ||
| 854 | - "name": "updates", | ||
| 855 | - "index": 2, | ||
| 856 | - "dtype": "float8_e8m0", | ||
| 857 | - "format": "ND", | ||
| 858 | - "paramType": "required", | ||
| 859 | - "shape": [ | ||
| 860 | - -2 | ||
| 861 | - ] | ||
| 862 | - } | ||
| 863 | - ], | ||
| 864 | - "outputs": [ | ||
| 865 | - { | ||
| 866 | - "name": "var", | ||
| 867 | - "index": 0, | ||
| 868 | - "dtype": "float8_e8m0", | ||
| 869 | - "format": "ND", | ||
| 870 | - "paramType": "required", | ||
| 871 | - "shape": [ | ||
| 872 | - -2 | ||
| 873 | - ] | ||
| 874 | - } | ||
| 875 | - ], | ||
| 876 | - "attrs": [ | ||
| 877 | - { | ||
| 878 | - "name": "reduce", | ||
| 879 | - "dtype": "string", | ||
| 880 | - "value": "update" | ||
| 881 | - }, | ||
| 882 | - { | ||
| 883 | - "name": "axis", | ||
| 884 | - "dtype": "int", | ||
| 885 | - "value": null | ||
| 886 | - } | ||
| 887 | - ] | ||
| 888 | - }, | ||
| 889 | { | 830 | { |
| 890 | "bin_filename": "Scatter_1f7a6ceb90cc559qsmaef65a606f7e01", | 831 | "bin_filename": "Scatter_1f7a6ceb90cc559qsmaef65a606f7e01", |
| 891 | "inputs": [ | 832 | "inputs": [ |
| @@ -1004,65 +945,6 @@ | |||
| 1004 | } | 945 | } |
| 1005 | ] | 946 | ] |
| 1006 | }, | 947 | }, |
| 1007 | - { | ||
| 1008 | - "bin_filename": "Scatter_1f3c60eb90aacac3qsmae360b606f7e01", | ||
| 1009 | - "inputs": [ | ||
| 1010 | - { | ||
| 1011 | - "name": "var", | ||
| 1012 | - "index": 0, | ||
| 1013 | - "dtype": "float8_e8m0", | ||
| 1014 | - "format": "ND", | ||
| 1015 | - "paramType": "required", | ||
| 1016 | - "shape": [ | ||
| 1017 | - -2 | ||
| 1018 | - ] | ||
| 1019 | - }, | ||
| 1020 | - { | ||
| 1021 | - "name": "indices", | ||
| 1022 | - "index": 1, | ||
| 1023 | - "dtype": "int32", | ||
| 1024 | - "format": "ND", | ||
| 1025 | - "paramType": "required", | ||
| 1026 | - "shape": [ | ||
| 1027 | - -2 | ||
| 1028 | - ] | ||
| 1029 | - }, | ||
| 1030 | - { | ||
| 1031 | - "name": "updates", | ||
| 1032 | - "index": 2, | ||
| 1033 | - "dtype": "float8_e8m0", | ||
| 1034 | - "format": "ND", | ||
| 1035 | - "paramType": "required", | ||
| 1036 | - "shape": [ | ||
| 1037 | - -2 | ||
| 1038 | - ] | ||
| 1039 | - } | ||
| 1040 | - ], | ||
| 1041 | - "outputs": [ | ||
| 1042 | - { | ||
| 1043 | - "name": "var", | ||
| 1044 | - "index": 0, | ||
| 1045 | - "dtype": "float8_e8m0", | ||
| 1046 | - "format": "ND", | ||
| 1047 | - "paramType": "required", | ||
| 1048 | - "shape": [ | ||
| 1049 | - -2 | ||
| 1050 | - ] | ||
| 1051 | - } | ||
| 1052 | - ], | ||
| 1053 | - "attrs": [ | ||
| 1054 | - { | ||
| 1055 | - "name": "reduce", | ||
| 1056 | - "dtype": "string", | ||
| 1057 | - "value": "update" | ||
| 1058 | - }, | ||
| 1059 | - { | ||
| 1060 | - "name": "axis", | ||
| 1061 | - "dtype": "int", | ||
| 1062 | - "value": null | ||
| 1063 | - } | ||
| 1064 | - ] | ||
| 1065 | - }, | ||
| 1066 | { | 948 | { |
| 1067 | "bin_filename": "Scatter_a93c60eb90aacac3qsmae360b606f7e7b", | 949 | "bin_filename": "Scatter_a93c60eb90aacac3qsmae360b606f7e7b", |
| 1068 | "inputs": [ | 950 | "inputs": [ |
| @@ -22,29 +22,29 @@ static const int64_t AXIS_DEFAULT = 0; | |||
| 22 | static const std::vector<ge::DataType> varDataType = { | 22 | static const std::vector<ge::DataType> varDataType = { |
| 23 | ge::DT_INT8, ge::DT_UINT8, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_INT32, | 23 | ge::DT_INT8, ge::DT_UINT8, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_INT32, |
| 24 | ge::DT_INT8, ge::DT_UINT8, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_INT32, | 24 | ge::DT_INT8, ge::DT_UINT8, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_INT32, |
| 25 | - ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E8M0, ge::DT_HIFLOAT8, | 25 | + ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_HIFLOAT8, |
| 26 | - ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E8M0, ge::DT_HIFLOAT8 | 26 | + ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_HIFLOAT8 |
| 27 | }; | 27 | }; |
| 28 | 28 | ||
| 29 | static const std::vector<ge::DataType> indicesDataType = { | 29 | static const std::vector<ge::DataType> indicesDataType = { |
| 30 | ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, | 30 | ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, |
| 31 | ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, | 31 | ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, |
| 32 | - ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, | 32 | + ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, |
| 33 | - ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64 | 33 | + ge::DT_INT64, ge::DT_INT64, ge::DT_INT64 |
| 34 | }; | 34 | }; |
| 35 | 35 | ||
| 36 | static const std::vector<ge::DataType> updatesDataType = { | 36 | static const std::vector<ge::DataType> updatesDataType = { |
| 37 | ge::DT_INT8, ge::DT_UINT8, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_INT32, | 37 | ge::DT_INT8, ge::DT_UINT8, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_INT32, |
| 38 | ge::DT_INT8, ge::DT_UINT8, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_INT32, | 38 | ge::DT_INT8, ge::DT_UINT8, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_INT32, |
| 39 | - ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E8M0, ge::DT_HIFLOAT8, | 39 | + ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_HIFLOAT8, |
| 40 | - ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E8M0, ge::DT_HIFLOAT8 | 40 | + ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_HIFLOAT8 |
| 41 | }; | 41 | }; |
| 42 | 42 | ||
| 43 | static const std::vector<ge::Format> format = { | 43 | static const std::vector<ge::Format> format = { |
| 44 | ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, | 44 | ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, |
| 45 | ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, | 45 | ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, |
| 46 | - ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, | 46 | + ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, |
| 47 | - ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND | 47 | + ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND |
| 48 | }; | 48 | }; |
| 49 | 49 | ||
| 50 | class Scatter : public OpDef { | 50 | class Scatter : public OpDef { |