使用cublasGemmBatchedEx做矩阵乘法结果全零,求代码问题排查
使用cublasGemmBatchedEx执行批量矩阵乘法结果全零的问题
我尝试用cublasGemmBatchedEx函数执行批量矩阵乘法,但输出的矩阵C全为零,以下是我的代码和运行结果,请帮忙排查问题:
原代码
#include <iostream> #include <cublas_v2.h> #define M 4 #define N 4 #define K 4 // nvcc -lcublas -o matmul_gemmBatchedEx matmul_gemmBatchedEx.cu void print_matrix(float **A, int rows, int cols, int batch_size) { for (int i = 0; i < batch_size; i++){ for (int j = 0; j < rows; j++){ for(int k = 0; k < cols; k++){ std::cout << A[i][k * rows + j] << " "; } std::cout << std::endl; } std::cout << std::endl; } } int main(int argc, char* argv[]) { // Linear dimension of matrices int batch_size = 2; float *h_A[batch_size], *h_B[batch_size], *h_C[batch_size]; for (int i = 0; i < batch_size; i++){ h_A[i] = (float*)malloc(M * K * sizeof(float)); h_B[i] = (float*)malloc(K * N * sizeof(float)); h_C[i] = (float*)malloc(M * N * sizeof(float)); } for (int i = 0; i < batch_size; i++){ for (int j = 0; j < M * K; j++) h_A[i][j] = j%4; for (int j = 0; j < K * N; j++) h_B[i][j] = j%4 + 4; for (int j = 0; j < M * N; j++) h_C[i][j] = 0; } std::cout << "A =" << std::endl; print_matrix(h_A, M, K, batch_size); std::cout << "B =" << std::endl; print_matrix(h_B, K, N, batch_size); float *d_A[batch_size], *d_B[batch_size], *d_C[batch_size]; for (int i = 0; i < batch_size; i++){ cudaMalloc(&d_A[i], sizeof(float)* M * K); cudaMalloc(&d_B[i], sizeof(float)* K * N); cudaMalloc(&d_C[i], sizeof(float)* M * N); } cudaMemcpy(d_A, h_A, sizeof(float)* M * K * batch_size, cudaMemcpyHostToDevice); cudaMemcpy(d_B, h_B, sizeof(float)* K * N * batch_size, cudaMemcpyHostToDevice); cublasHandle_t handle; cublasCreate(&handle); // Set up the matrix dimensions and batch size int lda = M; int ldb = K; int ldc = M; // Set the alpha and beta parameters for the gemm operation float alpha = 1.0f; float beta = 0.0f; cublasStatus_t status = cublasGemmBatchedEx(handle, CUBLAS_OP_N, CUBLAS_OP_N, M, N, K, &alpha, (const void**)d_A, CUDA_R_32F, lda, (const void**)d_B, CUDA_R_32F, ldb, &beta, (void**)d_C, CUDA_R_32F, ldc, batch_size, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT); cudaMemcpy(h_C,d_C,sizeof(float) * M * N * batch_size, cudaMemcpyDeviceToHost); if (status == CUBLAS_STATUS_SUCCESS) { std::cout << "C =" << std::endl; print_matrix(h_C, M, N, batch_size); } else { std::cout << status << std::endl; } // Destroy the handle cublasDestroy(handle); cudaFree(d_A); cudaFree(d_B); cudaFree(d_C); cudaFreeHost(h_A); cudaFreeHost(h_B); cudaFreeHost(h_C); }
运行结果
A = 0 0 0 0 1 1 1 1 2 2 2 2 3 3 3 3 0 0 0 0 1 1 1 1 2 2 2 2 3 3 3 3 B = 4 4 4 4 5 5 5 5 6 6 6 6 7 7 7 7 4 4 4 4 5 5 5 5 6 6 6 6 7 7 7 7 C = 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0
问题排查与修复
你的代码存在三个核心问题,导致矩阵乘法结果全零:
1. 设备指针数组未传入设备内存
cublasGemmBatchedEx要求输入的指针数组(如d_A、d_B)必须存储在设备内存中,但你代码里的d_A是主机栈上的数组,设备端的cublas无法直接访问主机栈内存,导致函数无法定位到矩阵数据的真实地址。
2. 内存拷贝逻辑错误
- 原代码中
cudaMemcpy(d_A, h_A, sizeof(float)* M * K * batch_size, ...)完全错误:d_A是指针数组,拷贝大小应为sizeof(float*) * batch_size,但更关键的是,你需要先把每个主机矩阵的数据拷贝到对应的设备内存,再将指针数组本身拷贝到设备。 - 正确流程是:先循环拷贝
h_A[i]到d_A[i],再创建设备端的指针数组,将主机上的d_A数组拷贝到设备。
3. 内存释放逻辑错误
cudaFree(d_A)、cudaFree(d_B)、cudaFree(d_C)是错误操作,因为d_A是主机栈数组,应该循环释放每个d_A[i];cudaFreeHost(h_A)同样错误,h_A是栈数组,里面的元素是malloc分配的,应该循环用free释放每个h_A[i]。
修复后的完整代码
#include <iostream> #include <cublas_v2.h> #define M 4 #define N 4 #define K 4 // nvcc -lcublas -o matmul_gemmBatchedEx matmul_gemmBatchedEx.cu void print_matrix(float **A, int rows, int cols, int batch_size) { for (int i = 0; i < batch_size; i++){ for (int j = 0; j < rows; j++){ for(int k = 0; k < cols; k++){ std::cout << A[i][k * rows + j] << " "; } std::cout << std::endl; } std::cout << std::endl; } } int main(int argc, char* argv[]) { int batch_size = 2; float *h_A[batch_size], *h_B[batch_size], *h_C[batch_size]; for (int i = 0; i < batch_size; i++){ h_A[i] = (float*)malloc(M * K * sizeof(float)); h_B[i] = (float*)malloc(K * N * sizeof(float)); h_C[i] = (float*)malloc(M * N * sizeof(float)); } for (int i = 0; i < batch_size; i++){ for (int j = 0; j < M * K; j++) h_A[i][j] = j%4; for (int j = 0; j < K * N; j++) h_B[i][j] = j%4 + 4; for (int j = 0; j < M * N; j++) h_C[i][j] = 0; } std::cout << "A =" << std::endl; print_matrix(h_A, M, K, batch_size); std::cout << "B =" << std::endl; print_matrix(h_B, K, N, batch_size); // 主机端存储设备指针的数组 float *d_A[batch_size], *d_B[batch_size], *d_C[batch_size]; // 设备端的指针数组,供cublas调用 float **d_A_ptr, **d_B_ptr, **d_C_ptr; cudaMalloc(&d_A_ptr, sizeof(float*) * batch_size); cudaMalloc(&d_B_ptr, sizeof(float*) * batch_size); cudaMalloc(&d_C_ptr, sizeof(float*) * batch_size); // 拷贝每个矩阵的数据到设备内存 for (int i = 0; i < batch_size; i++){ cudaMalloc(&d_A[i], sizeof(float)* M * K); cudaMalloc(&d_B[i], sizeof(float)* K * N); cudaMalloc(&d_C[i], sizeof(float)* M * N); cudaMemcpy(d_A[i], h_A[i], sizeof(float)* M * K, cudaMemcpyHostToDevice); cudaMemcpy(d_B[i], h_B[i], sizeof(float)* K * N, cudaMemcpyHostToDevice); } // 将主机端的设备指针数组拷贝到设备内存 cudaMemcpy(d_A_ptr, d_A, sizeof(float*) * batch_size, cudaMemcpyHostToDevice); cudaMemcpy(d_B_ptr, d_B, sizeof(float*) * batch_size, cudaMemcpyHostToDevice); cudaMemcpy(d_C_ptr, d_C, sizeof(float*) * batch_size, cudaMemcpyHostToDevice); cublasHandle_t handle; cublasCreate(&handle); int lda = M; int ldb = K; int ldc = M; float alpha = 1.0f; float beta = 0.0f; cublasStatus_t status = cublasGemmBatchedEx(handle, CUBLAS_OP_N, CUBLAS_OP_N, M, N, K, &alpha, (const void**)d_A_ptr, CUDA_R_32F, lda, (const void**)d_B_ptr, CUDA_R_32F, ldb, &beta, (void**)d_C_ptr, CUDA_R_32F, ldc, batch_size, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT); // 拷贝结果回主机 for (int i = 0; i < batch_size; i++){ cudaMemcpy(h_C[i], d_C[i], sizeof(float)* M * N, cudaMemcpyDeviceToHost); } if (status == CUBLAS_STATUS_SUCCESS) { std::cout << "C =" << std::endl; print_matrix(h_C, M, N, batch_size); } else { std::cout << "CUBLAS error code: " << status << std::endl; } cublasDestroy(handle); // 正确释放内存 for (int i = 0; i < batch_size; i++){ cudaFree(d_A[i]); cudaFree(d_B[i]); cudaFree(d_C[i]); free(h_A[i]); free(h_B[i]); free(h_C[i]); } cudaFree(d_A_ptr); cudaFree(d_B_ptr); cudaFree(d_C_ptr); return 0; }
修复后的预期输出
A = 0 0 0 0 1 1 1 1 2 2 2 2 3 3 3 3 0 0 0 0 1 1 1 1 2 2 2 2 3 3 3 3 B = 4 4 4 4 5 5 5 5 6 6 6 6 7 7 7 7 4 4 4 4 5 5 5 5 6 6 6 6 7 7 7 7 C = 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60 60
内容的提问来源于stack exchange,提问作者Bokyeong
相关产品推荐
相关产品推荐

