You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.25 04:20:41