提交记录 31420


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_dsh_260814 mmmd1k. 测测你的双精度矩阵乘法-1k Accepted 100 44.793 ms 8844 KB C++ 4.45 KB
提交时间 评测时间
2026-08-14 01:52:49 2026-08-14 01:52:51
#include <immintrin.h>
#include <stddef.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>

#pragma GCC target("avx2,fma")

// Tunable via -DMC=... etc. Defaults.
#ifndef MC
#define MC 512
#endif
#ifndef NC
#define NC 128
#endif
#ifndef KC
#define KC 128
#endif

#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;
#pragma GCC unroll 4
    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) {
    // Loop order: ic (m), pc (k), jc (n) -- A_pack reused across jc.
    for (int ic = 0; ic < n; ic += MC) {
        for (int pc = 0; pc < n; pc += KC) {
            // 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 jc = 0; jc < n; jc += NC) {
                // 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);
                    }
                }
                // compute
                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);
                    }
                }
            }
        }
    }
}

// ---- bench harness (for customtest) ----
#ifdef BENCH
static double Ab[1024*1024] __attribute__((aligned(4096)));
static double Bb[1024*1024] __attribute__((aligned(4096)));
static double Cb[1024*1024] __attribute__((aligned(4096)));

int main() {
    const int n = 1024;
    // Fast fill (cheap, no division) so the measurement is dominated by GEMM.
    for (int i = 0; i < n*n; i++) {
        Ab[i] = (double)(i & 1023) * (1.0/1024.0) - 0.5;
        Bb[i] = (double)((i * 7) & 1023) * (1.0/1024.0) - 0.5;
    }
    matrix_multiply(n, Ab, Bb, Cb);
    volatile double acc = Cb[0] + Cb[n*n-1] + Cb[12345];
    printf("%g\n", (double)acc);
    return 0;
}
#endif

CompilationN/AN/ACompile OKScore: N/A

Testcase #144.793 ms8 MB + 652 KBAcceptedScore: 100


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