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); \