提交记录 31303


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_dsh_260814 mmmd1k. 测测你的双精度矩阵乘法-1k Compile Error 0 0 ns 0 KB C++ 3.74 KB
提交时间 评测时间
2026-08-14 01:26:25 2026-08-14 01:26:27
#include <immintrin.h>
#include <stddef.h>

// mmmd1k: C = A*B, n=1024, row-major, double precision.
// Blocked GEMM: MC=NC=KC=256 (divide 1024), microkernel 4x8 AVX2+FMA.

#define MC 256
#define KC 256
#define NC 256
#define MR 4
#define NR 8

static double A_pack[MC * KC] __attribute__((aligned(4096)));
static double B_pack[KC * NC] __attribute__((aligned(4096)));

// A_pack: [(p*KC + k)*MR + r] = A[(ic + p*MR + r)*n + (pc + k)]
// B_pack: [(q*KC + k)*NR + c] = B[(pc + k)*n + (jc + q*NR + c)]

static inline void microkernel(const double *a, const double *b, double *c, int lda) {
    __m256d c00 = _mm256_loadu_pd(c + 0*lda);
    __m256d c01 = _mm256_loadu_pd(c + 0*lda + 4);
    __m256d c10 = _mm256_loadu_pd(c + 1*lda);
    __m256d c11 = _mm256_loadu_pd(c + 1*lda + 4);
    __m256d c20 = _mm256_loadu_pd(c + 2*lda);
    __m256d c21 = _mm256_loadu_pd(c + 2*lda + 4);
    __m256d c30 = _mm256_loadu_pd(c + 3*lda);
    __m256d c31 = _mm256_loadu_pd(c + 3*lda + 4);

    const double *ap = a;
    const double *bp = b;
    for (int k = 0; k < KC; k++) {
        __m256d b0 = _mm256_loadu_pd(bp);
        __m256d b1 = _mm256_loadu_pd(bp + 4);
        __m256d a0 = _mm256_broadcast_sd(ap + 0);
        c00 = _mm256_fmadd_pd(a0, b0, c00);
        c01 = _mm256_fmadd_pd(a0, b1, c01);
        __m256d a1 = _mm256_broadcast_sd(ap + 1);
        c10 = _mm256_fmadd_pd(a1, b0, c10);
        c11 = _mm256_fmadd_pd(a1, b1, c11);
        __m256d a2 = _mm256_broadcast_sd(ap + 2);
        c20 = _mm256_fmadd_pd(a2, b0, c20);
        c21 = _mm256_fmadd_pd(a2, b1, c21);
        __m256d a3 = _mm256_broadcast_sd(ap + 3);
        c30 = _mm256_fmadd_pd(a3, b0, c30);
        c31 = _mm256_fmadd_pd(a3, b1, c31);
        ap += MR;
        bp += NR;
    }

    _mm256_storeu_pd(c + 0*lda, c00);
    _mm256_storeu_pd(c + 0*lda + 4, c01);
    _mm256_storeu_pd(c + 1*lda, c10);
    _mm256_storeu_pd(c + 1*lda + 4, c11);
    _mm256_storeu_pd(c + 2*lda, c20);
    _mm256_storeu_pd(c + 2*lda + 4, c21);
    _mm256_storeu_pd(c + 3*lda, c30);
    _mm256_storeu_pd(c + 3*lda + 4, c31);
}

void matrix_multiply(int n, const double *A, const double *B, double *C) {
    const int M = n, N = n, K = n;
    const int npM = M / MR;   // 256
    const int npN = N / NR;   // 128

    for (int jc = 0; jc < N; jc += NC) {
        for (int pc = 0; pc < K; pc += KC) {
            // pack B: KC x NC
            for (int k = 0; k < KC; k++) {
                const double *Brow = B + (pc + k) * n + jc;
                double *bp = B_pack + k * NR;
                for (int q = 0; q < NC / NR; q++) {
                    __m256d v0 = _mm256_loadu_pd(Brow + q * NR);
                    __m256d v1 = _mm256_loadu_pd(Brow + q * NR + 4);
                    _mm256_storeu_pd(bp + q * KC * NR, v0);
                    _mm256_storeu_pd(bp + q * KC * NR + 4, v1);
                }
            }
            for (int ic = 0; ic < M; ic += MC) {
                // pack A: MC x KC
                for (int p = 0; p < MC / MR; p++) {
                    for (int k = 0; k < KC; k++) {
                        const double *Acol = A + (ic + p * MR) * n + (pc + k);
                        double *ap = A_pack + (p * KC + k) * MR;
                        ap[0] = Acol[0];
                        ap[1] = Acol[n];
                        ap[2] = Acol[2*n];
                        ap[3] = Acol[3*n];
                    }
                }
                for (int p = 0; p < MC / MR; p++) {
                    double *cp = C + (ic + p * MR) * n + jc;
                    const double *ap = A_pack + p * KC * MR;
                    for (int q = 0; q < NC / NR; q++) {
                        const double *bp = B_pack + q * KC * NR;
                        microkernel(ap, bp, cp + q * NR, n);
                    }
                }
            }
        }
    }
}

CompilationN/AN/ACompile ErrorScore: N/A


Judge Duck Online | 评测鸭在线
Server Time: 2026-09-12 09:25:08 | Loaded in 1 ms | Server Status
个人娱乐项目,仅供学习交流使用 | 捐赠