From 1d5cc226f4dd200250778634b202bb092e1a5ea7 Mon Sep 17 00:00:00 2001 From: Eric Lunderberg Date: Tue, 15 Feb 2022 13:44:00 -0600 Subject: [PATCH] [Cuda] Updated bfloat16 math defs. Required to pass `test_cuda_bf16_vectorize_add` in `tests/python/unittest/test_target_codegen_cuda.py`. --- src/target/source/literal/cuda_half_t.h | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/target/source/literal/cuda_half_t.h b/src/target/source/literal/cuda_half_t.h index 0ecefc3178f6..c3a1c66e5874 100644 --- a/src/target/source/literal/cuda_half_t.h +++ b/src/target/source/literal/cuda_half_t.h @@ -338,7 +338,7 @@ __pack_nv_bfloat162(const nv_bfloat16 x, const nv_bfloat16 y) { // so we define them here to make sure the generated CUDA code // is valid. #define CUDA_UNSUPPORTED_HALF_MATH_BINARY(HALF_MATH_NAME, FP32_MATH_NAME) \ -static inline __device__ __host__ half HALF_MATH_NAME(half x, half y) { \ +static inline __device__ __host__ nv_bfloat16 HALF_MATH_NAME(nv_bfloat16 x, nv_bfloat16 y) { \ float tmp_x = __bfloat162float(x); \ float tmp_y = __bfloat162float(y); \ float result = FP32_MATH_NAME(tmp_x, tmp_y); \ @@ -346,7 +346,7 @@ static inline __device__ __host__ half HALF_MATH_NAME(half x, half y) { \ } #define CUDA_UNSUPPORTED_HALF_MATH_UNARY(HALF_MATH_NAME, FP32_MATH_NAME) \ -static inline __device__ __host__ half HALF_MATH_NAME(half x) { \ +static inline __device__ __host__ nv_bfloat16 HALF_MATH_NAME(nv_bfloat16 x) { \ float tmp_x = __bfloat162float(x); \ float result = FP32_MATH_NAME(tmp_x); \ return __float2bfloat16(result); \