CUDA核函数中不转换格式直接打印FP16元素的正确方法
CUDA中half类型的printf正确用法
问题原因
你遇到的问题根源在于:CUDA中的half是非平凡可拷贝的结构体类型(定义为struct __half),无法直接通过printf的可变参数列表传递,这会触发编译警告;同时%f格式符对应float类型,参数不匹配导致打印出错误的0值。
解决方案
1. 打印浮点数值(需转换为float)
如果要输出half的浮点数值,无法绕过转换步骤,但可以使用CUDA设备端内置的__half2float()函数快速转换,这是最直接的方式:
修改后的代码:
#include <stdio.h> #include <cuda_fp16.h> #define DTYPE half #define PRINT_FMT "output %f\n" __global__ void __launch_bounds__(1024) test_print_kernel(DTYPE *__restrict__ O) { if (((int)blockIdx.x == 0) && ((int)threadIdx.x == 0)) { O[0] = ((DTYPE)(2)); __syncthreads(); // 用CUDA内置函数转换为float后打印 printf(PRINT_FMT, __half2float(O[0])); } } int main(int argc, char **argv) { DTYPE *h_O; cudaStream_t stream; cudaStreamCreateWithFlags(&stream, cudaStreamNonBlocking); cudaMallocHost(&h_O, 1 * sizeof(DTYPE)); test_print_kernel<<<dim3(1, 1, 1), dim3(1, 1, 1), 0, stream>>>(h_O); cudaDeviceSynchronize(); }
编译运行后会输出正确的output 2.000000。
2. 打印原始十六进制表示(不转换类型)
如果严格要求不转换为float/double,可以打印half类型的原始16位二进制对应的十六进制值(比如half类型的2.0对应十六进制0x4000):
修改后的代码:
#include <stdio.h> #include <cuda_fp16.h> #define DTYPE half #define PRINT_FMT "output 0x%hx\n" __global__ void __launch_bounds__(1024) test_print_kernel(DTYPE *__restrict__ O) { if (((int)blockIdx.x == 0) && ((int)threadIdx.x == 0)) { O[0] = ((DTYPE)(2)); __syncthreads(); // 强制转换为unsigned short打印十六进制原始值 printf(PRINT_FMT, *(unsigned short*)&O[0]); } } int main(int argc, char **argv) { DTYPE *h_O; cudaStream_t stream; cudaStreamCreateWithFlags(&stream, cudaStreamNonBlocking); cudaMallocHost(&h_O, 1 * sizeof(DTYPE)); test_print_kernel<<<dim3(1, 1, 1), dim3(1, 1, 1), 0, stream>>>(h_O); cudaDeviceSynchronize(); }
运行后会输出output 0x4000,对应half类型的2.0。
内容的提问来源于stack exchange,提问作者Chun Ting Li
相关产品推荐
相关产品推荐

