DGEMM与DGEMV结果不一致:C语言实现三维张量求和运算遇问题
实现张量运算 C[l,q,m] = Σₖ A[m,q,k] * B[k,l] 的三种方案
我们需要实现的是一个三维张量与二维矩阵的缩并运算,其中重复索引k为求和维度。下面提供三种可运行的实现方式:朴素嵌套循环、利用BLAS的DGEMV(矩阵-向量乘法)、利用BLAS的DGEMM(矩阵-矩阵乘法),所有代码都附带详细注释和结果验证逻辑。
1. 朴素循环实现
这种方式逻辑最直观,通过多层嵌套循环遍历所有索引逐个计算值,适合理解运算本质,但大尺寸数据下性能较差。
#include <stdio.h> #include <stdint.h> #include <stdlib.h> #include <string.h> #include <cblas.h> #include <math.h> #define M 2 // 维度m的大小 #define Q 3 // 维度q的大小 #define K 4 // 维度k的大小 #define L 5 // 维度l的大小 // 初始化张量A void init_A(double *A) { for (int m = 0; m < M; m++) { for (int q = 0; q < Q; q++) { for (int k = 0; k < K; k++) { A[m*Q*K + q*K + k] = m + q + k + 1.0; // 简单赋值便于验证 } } } } // 初始化矩阵B void init_B(double *B) { for (int k = 0; k < K; k++) { for (int l = 0; l < L; l++) { B[k*L + l] = k + l + 1.0; } } } // 朴素循环实现目标运算 void naive_implementation(double *C, const double *A, const double *B) { memset(C, 0, sizeof(double)*L*Q*M); for (int m = 0; m < M; m++) { for (int q = 0; q < Q; q++) { for (int l = 0; l < L; l++) { double sum = 0.0; for (int k = 0; k < K; k++) { // 索引映射:A[m,q,k] = A[m*Q*K + q*K + k] // B[k,l] = B[k*L + l] sum += A[m*Q*K + q*K + k] * B[k*L + l]; } // 索引映射:C[l,q,m] = C[l*Q*M + q*M + m] C[l*Q*M + q*M + m] = sum; } } } }
2. BLAS DGEMV实现
DGEMV是BLAS优化的矩阵-向量乘法接口,我们可以把运算拆解为:对每个(m,q)对,取A的对应行向量A[m,q,:],计算它与B的每一列的点积。通过转置B,我们可以用DGEMV批量完成这些点积计算,性能比朴素循环更优。
// DGEMV实现目标运算 void dgemv_implementation(double *C, const double *A, const double *B) { memset(C, 0, sizeof(double)*L*Q*M); const double alpha = 1.0; const double beta = 0.0; for (int m = 0; m < M; m++) { for (int q = 0; q < Q; q++) { // 取出A中对应(m,q)的行向量,长度为K const double *A_vec = &A[m*Q*K + q*K]; // 计算 C[:,q,m] = A_vec * B,等价于B^T * A_vec^T(转置后用DGEMV计算) cblas_dgemv(CblasRowMajor, CblasTrans, L, K, alpha, B, L, A_vec, 1, beta, &C[q*M + m], M); // 结果步长设为M,保证l递增时地址偏移正确对应C[l,q,m]的索引 } } }
3. BLAS DGEMM实现
DGEMM是BLAS中高度优化的矩阵-矩阵乘法接口,我们可以将三维张量A重构为二维矩阵:把(m,q)作为行索引,k作为列索引,得到(M*Q)×K的矩阵;再与K×L的矩阵B相乘,得到(M*Q)×L的结果矩阵,最后映射回三维张量C。这种方式性能最优,尤其适合大尺寸数据。
// DGEMM实现目标运算 void dgemm_implementation(double *C, const double *A, const double *B) { memset(C, 0, sizeof(double)*L*Q*M); const double alpha = 1.0; const double beta = 0.0; // 将A重构为(M*Q)×K的二维矩阵(行优先存储) double *A_mat = (double*)malloc(sizeof(double)*M*Q*K); memcpy(A_mat, A, sizeof(double)*M*Q*K); // 计算矩阵乘法:C_mat = A_mat * B,C_mat为(M*Q)×L的矩阵 double *C_mat = (double*)malloc(sizeof(double)*M*Q*L); cblas_dgemm(CblasRowMajor, CblasNoTrans, CblasNoTrans, M*Q, L, K, alpha, A_mat, K, B, L, beta, C_mat, L); // 将二维结果矩阵映射回三维张量C for (int l = 0; l < L; l++) { for (int q = 0; q < Q; q++) { for (int m = 0; m < M; m++) { C[l*Q*M + q*M + m] = C_mat[(q*M + m)*L + l]; } } } free(A_mat); free(C_mat); }
结果验证与主函数
下面的主函数用于初始化数据、调用三种实现,并验证结果的一致性(考虑浮点数精度误差):
// 验证两个结果是否一致 int verify_result(const double *C1, const double *C2) { const double eps = 1e-9; for (int i = 0; i < L*Q*M; i++) { if (fabs(C1[i] - C2[i]) > eps) { printf("Result mismatch at index %d: %lf vs %lf\n", i, C1[i], C2[i]); return -1; } } printf("All results match!\n"); return 0; } int main() { // 分配内存 double *A = (double*)malloc(sizeof(double)*M*Q*K); double *B = (double*)malloc(sizeof(double)*K*L); double *C_naive = (double*)malloc(sizeof(double)*L*Q*M); double *C_dgemv = (double*)malloc(sizeof(double)*L*Q*M); double *C_dgemm = (double*)malloc(sizeof(double)*L*Q*M); // 初始化数据 init_A(A); init_B(B); // 调用三种实现 naive_implementation(C_naive, A, B); dgemv_implementation(C_dgemv, A, B); dgemm_implementation(C_dgemm, A, B); // 验证结果一致性 printf("Verifying DGEMV vs Naive implementation...\n"); verify_result(C_naive, C_dgemv); printf("Verifying DGEMM vs Naive implementation...\n"); verify_result(C_naive, C_dgemm); // 释放内存 free(A); free(B); free(C_naive); free(C_dgemv); free(C_dgemm); return 0; }
编译与运行说明
编译时需要链接BLAS库,例如使用gcc:
gcc -o tensor_operation tensor_operation.c -lcblas
运行生成的可执行文件即可看到验证结果。
内容的提问来源于stack exchange,提问作者Max
相关产品推荐
相关产品推荐

