已合并
support mxfp8 #74
linyixin创建于 3月2日
support mxfp8 #74
已合并
linyixin创建于 3月2日
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 
634std::map<int, ReduceCheckBufInitFunc> functionReduceMap = {635std::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()
121void HcclOpBaseTest::is_initdata_overflow()122void HcclOpBaseTest::is_initdata_overflow()
122{123{
123 if((dtype == HCCL_DATA_TYPE_INT8 || dtype == HCCL_DATA_TYPE_UINT8 || dtype == HCCL_DATA_TYPE_HIF8124 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#include <arpa/inet.h>27#include <arpa/inet.h>
28 28 
29constexpr s32 HCCL_TEST_REDUCE_RESERVED = 4;29constexpr s32 HCCL_TEST_REDUCE_RESERVED = 4;
30-constexpr s32 HCCL_TEST_DATA_TYPE_RESERVED = 17;30+constexpr s32 HCCL_TEST_DATA_TYPE_RESERVED = 18;
31constexpr s32 HCCL_TEST_ACCELERATOR_CONFIG_RESERVED = 8;31constexpr s32 HCCL_TEST_ACCELERATOR_CONFIG_RESERVED = 8;
32constexpr u32 HCCL_TEST_DATATYPE_BF16_SAT = 13;32constexpr u32 HCCL_TEST_DATATYPE_BF16_SAT = 13;
33HcclReduceOp test_ops[HCCL_TEST_REDUCE_RESERVED] = {33HcclReduceOp 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};
57const char *test_typenames[HCCL_TEST_DATA_TYPE_RESERVED] = {58const 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 
61int get_hccl_op_from_str(char *str)62int 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: