已合并
support mxfp8 #74
linyixin创建于 3月2日
support mxfp8 #74
已合并
共 7 个文件变更+14-7
| @@ -628,7 +628,8 @@ std::map<int,HostBufInitFunc> functionMap = { | |||
| 628 | std::pair<int,HostBufInitFunc>(HCCL_DATA_TYPE_HIF8, host_buf_init_int8), | 628 | std::pair<int,HostBufInitFunc>(HCCL_DATA_TYPE_HIF8, host_buf_init_int8), |
| 629 | std::pair<int,HostBufInitFunc>(HCCL_DATA_TYPE_FP8E4M3, host_buf_init_int8), | 629 | std::pair<int,HostBufInitFunc>(HCCL_DATA_TYPE_FP8E4M3, host_buf_init_int8), |
| 630 | std::pair<int,HostBufInitFunc>(HCCL_DATA_TYPE_FP8E5M2, host_buf_init_int8), | 630 | std::pair<int,HostBufInitFunc>(HCCL_DATA_TYPE_FP8E5M2, host_buf_init_int8), |
| 631 | - std::pair<int,HostBufInitFunc>(HCCL_DATA_TYPE_FP8E8M0, host_buf_init_int8) | 631 | + std::pair<int,HostBufInitFunc>(HCCL_DATA_TYPE_FP8E8M0, host_buf_init_int8), |
| 632 | + std::pair<int,HostBufInitFunc>(HCCL_DATA_TYPE_MXFP8, host_buf_init_int8) | ||
| 632 | }; | 633 | }; |
| 633 | 634 | ||
| 634 | std::map<int, ReduceCheckBufInitFunc> functionReduceMap = { | 635 | std::map<int, ReduceCheckBufInitFunc> functionReduceMap = { |
| @@ -659,4 +660,5 @@ std::map<int, AllToAllCheckResult> functionAllToAllMap = { | |||
| 659 | std::pair<int, AllToAllCheckResult>(HCCL_DATA_TYPE_HIF8, alltoall_check_result_uint8), | 660 | std::pair<int, AllToAllCheckResult>(HCCL_DATA_TYPE_HIF8, alltoall_check_result_uint8), |
| 660 | std::pair<int, AllToAllCheckResult>(HCCL_DATA_TYPE_FP8E4M3, alltoall_check_result_uint8), | 661 | std::pair<int, AllToAllCheckResult>(HCCL_DATA_TYPE_FP8E4M3, alltoall_check_result_uint8), |
| 661 | std::pair<int, AllToAllCheckResult>(HCCL_DATA_TYPE_FP8E5M2, alltoall_check_result_uint8), | 662 | std::pair<int, AllToAllCheckResult>(HCCL_DATA_TYPE_FP8E5M2, alltoall_check_result_uint8), |
| 662 | - std::pair<int, AllToAllCheckResult>(HCCL_DATA_TYPE_FP8E8M0, alltoall_check_result_uint8)}; | 663 | + std::pair<int, AllToAllCheckResult>(HCCL_DATA_TYPE_FP8E8M0, alltoall_check_result_uint8), |
| 664 | + std::pair<int, AllToAllCheckResult>(HCCL_DATA_TYPE_MXFP8, alltoall_check_result_uint8)}; | ||
| @@ -79,6 +79,7 @@ void HcclOpBaseTest::init_data_count() | |||
| 79 | case HCCL_DATA_TYPE_FP8E4M3: | 79 | case HCCL_DATA_TYPE_FP8E4M3: |
| 80 | case HCCL_DATA_TYPE_FP8E5M2: | 80 | case HCCL_DATA_TYPE_FP8E5M2: |
| 81 | case HCCL_DATA_TYPE_FP8E8M0: | 81 | case HCCL_DATA_TYPE_FP8E8M0: |
| 82 | + case HCCL_DATA_TYPE_MXFP8: | ||
| 82 | data->count = (data->data_size + sizeof(unsigned char) - 1)/sizeof(unsigned char); | 83 | data->count = (data->data_size + sizeof(unsigned char) - 1)/sizeof(unsigned char); |
| 83 | data->type_size = sizeof(unsigned char); | 84 | data->type_size = sizeof(unsigned char); |
| 84 | break; | 85 | break; |
| @@ -121,7 +122,7 @@ void HcclOpBaseTest::no_verification() | |||
| 121 | void HcclOpBaseTest::is_initdata_overflow() | 122 | void HcclOpBaseTest::is_initdata_overflow() |
| 122 | { | 123 | { |
| 123 | if((dtype == HCCL_DATA_TYPE_INT8 || dtype == HCCL_DATA_TYPE_UINT8 || dtype == HCCL_DATA_TYPE_HIF8 | 124 | if((dtype == HCCL_DATA_TYPE_INT8 || dtype == HCCL_DATA_TYPE_UINT8 || dtype == HCCL_DATA_TYPE_HIF8 |
| 124 | - || dtype == HCCL_DATA_TYPE_FP8E4M3 || dtype == HCCL_DATA_TYPE_FP8E5M2 || dtype == HCCL_DATA_TYPE_FP8E8M0) | 125 | + || dtype == HCCL_DATA_TYPE_FP8E4M3 || dtype == HCCL_DATA_TYPE_FP8E5M2 || dtype == HCCL_DATA_TYPE_FP8E8M0 || dtype == HCCL_DATA_TYPE_MXFP8) |
| 125 | && rank_size >= RANKSIZE_TH_FP32) { | 126 | && rank_size >= RANKSIZE_TH_FP32) { |
| 126 | check = 0; //不进行校验 | 127 | check = 0; //不进行校验 |
| 127 | if (rank_id == root_rank && print_dump) { | 128 | if (rank_id == root_rank && print_dump) { |
| @@ -27,7 +27,7 @@ | |||
| 27 | 27 | ||
| 28 | 28 | ||
| 29 | constexpr s32 HCCL_TEST_REDUCE_RESERVED = 4; | 29 | constexpr s32 HCCL_TEST_REDUCE_RESERVED = 4; |
| 30 | -constexpr s32 HCCL_TEST_DATA_TYPE_RESERVED = 17; | 30 | +constexpr s32 HCCL_TEST_DATA_TYPE_RESERVED = 18; |
| 31 | constexpr s32 HCCL_TEST_ACCELERATOR_CONFIG_RESERVED = 8; | 31 | constexpr s32 HCCL_TEST_ACCELERATOR_CONFIG_RESERVED = 8; |
| 32 | constexpr u32 HCCL_TEST_DATATYPE_BF16_SAT = 13; | 32 | constexpr u32 HCCL_TEST_DATATYPE_BF16_SAT = 13; |
| 33 | HcclReduceOp test_ops[HCCL_TEST_REDUCE_RESERVED] = { | 33 | HcclReduceOp test_ops[HCCL_TEST_REDUCE_RESERVED] = { |
| @@ -52,11 +52,12 @@ HcclDataType test_types[HCCL_TEST_DATA_TYPE_RESERVED] = {HCCL_DATA_TYPE_INT8, /* | |||
| 52 | HCCL_DATA_TYPE_HIF8, /**< hif8 */ | 52 | HCCL_DATA_TYPE_HIF8, /**< hif8 */ |
| 53 | HCCL_DATA_TYPE_FP8E4M3, /**< fp8e4m3 */ | 53 | HCCL_DATA_TYPE_FP8E4M3, /**< fp8e4m3 */ |
| 54 | HCCL_DATA_TYPE_FP8E5M2, /**< fp8e5m2 */ | 54 | HCCL_DATA_TYPE_FP8E5M2, /**< fp8e5m2 */ |
| 55 | - HCCL_DATA_TYPE_FP8E8M0 /**< fp8e8m0 */ | 55 | + HCCL_DATA_TYPE_FP8E8M0, /**< fp8e8m0 */ |
| 56 | + HCCL_DATA_TYPE_MXFP8 /**< mxfp8 */ | ||
| 56 | }; | 57 | }; |
| 57 | const char *test_typenames[HCCL_TEST_DATA_TYPE_RESERVED] = { | 58 | const char *test_typenames[HCCL_TEST_DATA_TYPE_RESERVED] = { |
| 58 | "int8", "int16", "int32", "fp16", "fp32", "int64", "uint64", "uint8", "uint16", "uint32", "fp64", "bfp16", | 59 | "int8", "int16", "int32", "fp16", "fp32", "int64", "uint64", "uint8", "uint16", "uint32", "fp64", "bfp16", |
| 59 | - "int128", "hif8", "fp8e4m3", "fp8e5m2", "fp8e8m0"}; | 60 | + "int128", "hif8", "fp8e4m3", "fp8e5m2", "fp8e8m0", "mxfp8"}; |
| 60 | 61 | ||
| 61 | int get_hccl_op_from_str(char *str) | 62 | int get_hccl_op_from_str(char *str) |
| 62 | { | 63 | { |
| @@ -88,6 +88,7 @@ int HcclOpBaseAllgatherTest::check_buf_result() | |||
| 88 | case HCCL_DATA_TYPE_FP8E4M3: | 88 | case HCCL_DATA_TYPE_FP8E4M3: |
| 89 | case HCCL_DATA_TYPE_FP8E5M2: | 89 | case HCCL_DATA_TYPE_FP8E5M2: |
| 90 | case HCCL_DATA_TYPE_FP8E8M0: | 90 | case HCCL_DATA_TYPE_FP8E8M0: |
| 91 | + case HCCL_DATA_TYPE_MXFP8: | ||
| 91 | ret = check_buf_result_int8((char*)recv_buff_temp, (char*)check_buf, data->count * rank_size); | 92 | ret = check_buf_result_int8((char*)recv_buff_temp, (char*)check_buf, data->count * rank_size); |
| 92 | break; | 93 | break; |
| 93 | case HCCL_DATA_TYPE_INT32: | 94 | case HCCL_DATA_TYPE_INT32: |
| @@ -113,6 +113,7 @@ int HcclOpBaseAllgatherVTest::check_buf_result() | |||
| 113 | case HCCL_DATA_TYPE_FP8E4M3: | 113 | case HCCL_DATA_TYPE_FP8E4M3: |
| 114 | case HCCL_DATA_TYPE_FP8E5M2: | 114 | case HCCL_DATA_TYPE_FP8E5M2: |
| 115 | case HCCL_DATA_TYPE_FP8E8M0: | 115 | case HCCL_DATA_TYPE_FP8E8M0: |
| 116 | + case HCCL_DATA_TYPE_MXFP8: | ||
| 116 | ret = check_buf_result_int8((char*)recv_buff_temp, (char*)check_buf, data->count * rank_size); | 117 | ret = check_buf_result_int8((char*)recv_buff_temp, (char*)check_buf, data->count * rank_size); |
| 117 | break; | 118 | break; |
| 118 | case HCCL_DATA_TYPE_INT32: | 119 | case HCCL_DATA_TYPE_INT32: |
| @@ -82,6 +82,7 @@ int HcclOpBaseBrocastTest::check_buf_result() | |||
| 82 | case HCCL_DATA_TYPE_FP8E4M3: | 82 | case HCCL_DATA_TYPE_FP8E4M3: |
| 83 | case HCCL_DATA_TYPE_FP8E5M2: | 83 | case HCCL_DATA_TYPE_FP8E5M2: |
| 84 | case HCCL_DATA_TYPE_FP8E8M0: | 84 | case HCCL_DATA_TYPE_FP8E8M0: |
| 85 | + case HCCL_DATA_TYPE_MXFP8: | ||
| 85 | ret = check_buf_result_int8((char*)recv_buff_temp, (char*)check_buf, data->count); | 86 | ret = check_buf_result_int8((char*)recv_buff_temp, (char*)check_buf, data->count); |
| 86 | break; | 87 | break; |
| 87 | case HCCL_DATA_TYPE_INT32: | 88 | case HCCL_DATA_TYPE_INT32: |
| @@ -82,7 +82,7 @@ int HcclOpBaseScatterTest::check_buf_result() | |||
| 82 | case HCCL_DATA_TYPE_FP8E4M3: | 82 | case HCCL_DATA_TYPE_FP8E4M3: |
| 83 | case HCCL_DATA_TYPE_FP8E5M2: | 83 | case HCCL_DATA_TYPE_FP8E5M2: |
| 84 | case HCCL_DATA_TYPE_FP8E8M0: | 84 | case HCCL_DATA_TYPE_FP8E8M0: |
| 85 | - | 85 | + case HCCL_DATA_TYPE_MXFP8: |
| 86 | ret = check_buf_result_int8((char*)recv_buff_temp, (char*)check_buf, data->count); | 86 | ret = check_buf_result_int8((char*)recv_buff_temp, (char*)check_buf, data->count); |
| 87 | break; | 87 | break; |
| 88 | case HCCL_DATA_TYPE_INT32: | 88 | case HCCL_DATA_TYPE_INT32: |