提交记录 55815


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_dsh_v41_0919 mmmd1k. 测测你的双精度矩阵乘法-1k Accepted 100 48.014 ms 10492 KB C++17 7.22 KB
提交时间 评测时间
2026-09-19 19:21:25 2026-09-19 19:44:09
// mmmd1k: C = A*B, n = 1024, row-major doubles.
// Blocked GEMM, AVX2+FMA 6x8 register microkernel (4x8 for tail rows).
#pragma GCC target("avx2,fma")
#include <immintrin.h>
#include <cstring>

#define MR 4
#define MR6 6
#define NR 8
#define MC 120
#define KC 256
#define NC 1024

static double Apack[(size_t)MC * KC + 64] __attribute__((aligned(64)));
static double Bpack[(size_t)(KC + 4) * (NC + NR) + 64] __attribute__((aligned(64)));

// Pack 6 rows of an A block, k-major: Apack[(ig/6)*kc*6 + k*6 + i]
static void packA6(const double *A, int n, int ic, int pc, int mc, int kc) {
    double *dst = Apack;
    for (int ig = 0; ig < mc; ig += MR6) {
        const double *r0 = A + (size_t)(ic + ig) * n + pc;
        const double *r1 = r0 + n, *r2 = r1 + n, *r3 = r2 + n, *r4 = r3 + n, *r5 = r4 + n;
        for (int k = 0; k < kc; k++) {
            dst[0] = r0[k]; dst[1] = r1[k]; dst[2] = r2[k];
            dst[3] = r3[k]; dst[4] = r4[k]; dst[5] = r5[k];
            dst += MR6;
        }
    }
}

// Pack B into NR-wide panels: panel p holds kc rows of NR contiguous doubles.
static void packB(const double *B, int n, int pc, int jc, int kc, int nc) {
    double *dst = Bpack;
    for (int jr = 0; jr < nc; jr += NR) {
        const double *src = B + (size_t)pc * n + jc + jr;
        for (int k = 0; k < kc; k++) {
            _mm256_store_pd(dst, _mm256_loadu_pd(src + (size_t)k * n));
            _mm256_store_pd(dst + 4, _mm256_loadu_pd(src + (size_t)k * n + 4));
            dst += NR;
        }
    }
}

#define STROW(R, A, B) t = _mm256_loadu_pd(cp); _mm256_storeu_pd(cp, _mm256_add_pd(t, A)); \
    t = _mm256_loadu_pd(cp + 4); _mm256_storeu_pd(cp + 4, _mm256_add_pd(t, B)); cp += ldc;

static void kernel6x8(const double *Ap, const double *Bp, int kc, double *C, int ldc) {
    __m256d c0=_mm256_setzero_pd(),c1=_mm256_setzero_pd(),c2=_mm256_setzero_pd();
    __m256d c3=_mm256_setzero_pd(),c4=_mm256_setzero_pd(),c5=_mm256_setzero_pd();
    __m256d c6=_mm256_setzero_pd(),c7=_mm256_setzero_pd(),c8=_mm256_setzero_pd();
    __m256d c9=_mm256_setzero_pd(),c10=_mm256_setzero_pd(),c11=_mm256_setzero_pd();
    const double *ap = Ap, *bp = Bp;
#define K6STEP() {                                                              \
        __m256d b0 = _mm256_load_pd(bp);                                        \
        __m256d b1 = _mm256_load_pd(bp + 4);                                    \
        __m256d a;                                                              \
        a = _mm256_broadcast_sd(ap + 0); c0  = _mm256_fmadd_pd(a, b0, c0);  c6  = _mm256_fmadd_pd(a, b1, c6); \
        a = _mm256_broadcast_sd(ap + 1); c1  = _mm256_fmadd_pd(a, b0, c1);  c7  = _mm256_fmadd_pd(a, b1, c7); \
        a = _mm256_broadcast_sd(ap + 2); c2  = _mm256_fmadd_pd(a, b0, c2);  c8  = _mm256_fmadd_pd(a, b1, c8); \
        a = _mm256_broadcast_sd(ap + 3); c3  = _mm256_fmadd_pd(a, b0, c3);  c9  = _mm256_fmadd_pd(a, b1, c9); \
        a = _mm256_broadcast_sd(ap + 4); c4  = _mm256_fmadd_pd(a, b0, c4);  c10 = _mm256_fmadd_pd(a, b1, c10); \
        a = _mm256_broadcast_sd(ap + 5); c5  = _mm256_fmadd_pd(a, b0, c5);  c11 = _mm256_fmadd_pd(a, b1, c11); \
        ap += MR6; bp += NR;                                                    \
    }
    int k = 0;
    for (; k + 4 <= kc; k += 4) {
        _mm_prefetch((const char *)(ap + 384), _MM_HINT_T0);
        _mm_prefetch((const char *)(bp + 96), _MM_HINT_T0);
        K6STEP()
        K6STEP()
        K6STEP()
        K6STEP()
    }
    for (; k < kc; k++) K6STEP()
    double *cp = C; __m256d t;
    STROW(0, c0, c6) STROW(1, c1, c7) STROW(2, c2, c8)
    STROW(3, c3, c9) STROW(4, c4, c10) STROW(5, c5, c11)
}

static void kernel4x8(const double *Ap, const double *Bp, int kc, int nc,
                      double *C, int ldc) {
    __m256d c0=_mm256_setzero_pd(),c1=_mm256_setzero_pd(),c2=_mm256_setzero_pd(),c3=_mm256_setzero_pd();
    __m256d c4=_mm256_setzero_pd(),c5=_mm256_setzero_pd(),c6=_mm256_setzero_pd(),c7=_mm256_setzero_pd();
    const double *ap = Ap, *bp = Bp;
    for (int k = 0; k < kc; k++) {
        __m256d b0 = _mm256_load_pd(bp), b1 = _mm256_load_pd(bp + 4), a;
        a = _mm256_broadcast_sd(ap + 0); c0 = _mm256_fmadd_pd(a, b0, c0); c4 = _mm256_fmadd_pd(a, b1, c4);
        a = _mm256_broadcast_sd(ap + 1); c1 = _mm256_fmadd_pd(a, b0, c1); c5 = _mm256_fmadd_pd(a, b1, c5);
        a = _mm256_broadcast_sd(ap + 2); c2 = _mm256_fmadd_pd(a, b0, c2); c6 = _mm256_fmadd_pd(a, b1, c6);
        a = _mm256_broadcast_sd(ap + 3); c3 = _mm256_fmadd_pd(a, b0, c3); c7 = _mm256_fmadd_pd(a, b1, c7);
        ap += MR; bp += NR;
    }
    double *cp = C; __m256d t;
    STROW(0, c0, c4) STROW(1, c1, c5) STROW(2, c2, c6) STROW(3, c3, c7)
}

static void packA4(const double *A, int n, int ic, int pc, int mc, int kc) {
    double *dst = Apack;
    for (int ig = 0; ig < mc; ig += MR) {
        const double *r0 = A + (size_t)(ic + ig) * n + pc;
        const double *r1 = r0 + n, *r2 = r1 + n, *r3 = r2 + n;
        for (int k = 0; k < kc; k++) {
            dst[0] = r0[k]; dst[1] = r1[k]; dst[2] = r2[k]; dst[3] = r3[k];
            dst += MR;
        }
    }
}

static void edge_rows(const double *A, const double *B, double *C, int n, int i0, int i1, int j0, int j1) {
    for (int i = i0; i < i1; i++) {
        for (int j = j0; j < j1; j++) {
            const double *ar = A + (size_t)i * n;
            double s = 0;
            for (int k = 0; k < n; k++) s += ar[k] * B[(size_t)k * n + j];
            C[(size_t)i * n + j] = s;
        }
    }
}

void matrix_multiply(int n, const double *A, const double *B, double *C) {
    memset(C, 0, sizeof(double) * (size_t)n * n);
    int n6 = n - (n % MR6);   // rows handled by the 6-row kernel
    int n4 = n - (n % NR);    // cols handled by the vector kernels
    int ntail = n - n6;
    int n4r = n6 + (ntail - (ntail % MR));   // rows handled by the vector kernels
    for (int jc = 0; jc < n4; jc += NC) {
        int nc = n4 - jc; if (nc > NC) nc = NC;
        for (int pc = 0; pc < n; pc += KC) {
            int kc = n - pc; if (kc > KC) kc = KC;
            packB(B, n, pc, jc, kc, nc);
            for (int ic = 0; ic < n6; ic += MC) {
                int mc = n6 - ic; if (mc > MC) mc = MC;
                packA6(A, n, ic, pc, mc, kc);
                for (int jr = 0; jr < nc; jr += NR) {
                    const double *bp = Bpack + (size_t)(jr / NR) * kc * NR;
                    double *cb = C + (size_t)ic * n + jc + jr;
                    for (int ir = 0; ir < mc; ir += MR6)
                        kernel6x8(Apack + (size_t)(ir / MR6) * kc * MR6, bp, kc,
                                  cb + (size_t)ir * n, n);
                }
            }
            if (n6 < n4r) {   // rows [n6, n4r) with the 4-row kernel
                int mc = n4r - n6;
                packA4(A, n, n6, pc, mc, kc);
                for (int jr = 0; jr < nc; jr += NR) {
                    const double *bp = Bpack + (size_t)(jr / NR) * kc * NR;
                    double *cb = C + (size_t)n6 * n + jc + jr;
                    for (int ir = 0; ir < mc; ir += MR)
                        kernel4x8(Apack + (size_t)(ir / MR) * kc * MR, bp, kc, nc,
                                  cb + (size_t)ir * n, n);
                }
            }
        }
    }
    if (n4 < n) edge_rows(A, B, C, n, 0, n, n4, n);
    if (n4r < n) edge_rows(A, B, C, n, n4r, n, 0, n4);
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #148.014 ms10 MB + 252 KBAcceptedScore: 100


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