提交记录 47276


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_codex_260812 mmmd1k. 测测你的双精度矩阵乘法-1k Accepted 100 50.591 ms 16400 KB C 7.09 KB
提交时间 评测时间
2026-08-23 14:48:59 2026-08-23 14:49:02
#include <pthread.h>
#define WORKER_ONLY
#include <immintrin.h>

enum { N = 1024, MAIN_COLS = 1008, PANELS = 84 };
#ifndef PACK_PANELS
#define PACK_PANELS 4
#endif
#ifndef PACK_ROWS
#define PACK_ROWS 32
#endif
static double packed_b[N * N] __attribute__((aligned(4096)));

__attribute__((target("avx2,fma"), always_inline))
static inline void kernel4x12(const double *A, const double *bp, double *C,
                              int i, int j) {
    __m256d c00 = _mm256_setzero_pd(), c01 = c00, c02 = c00;
    __m256d c10 = c00, c11 = c00, c12 = c00;
    __m256d c20 = c00, c21 = c00, c22 = c00;
    __m256d c30 = c00, c31 = c00, c32 = c00;
    const double *a0 = A + (i + 0) * N;
    const double *a1 = A + (i + 1) * N;
    const double *a2 = A + (i + 2) * N;
    const double *a3 = A + (i + 3) * N;
    for (int k = 0; k < N; ++k, bp += 12) {
        __m256d b0 = _mm256_load_pd(bp);
        __m256d b1 = _mm256_load_pd(bp + 4);
        __m256d b2 = _mm256_load_pd(bp + 8);
        __m256d av = _mm256_broadcast_sd(a0 + k);
        c00 = _mm256_fmadd_pd(av, b0, c00);
        c01 = _mm256_fmadd_pd(av, b1, c01);
        c02 = _mm256_fmadd_pd(av, b2, c02);
        av = _mm256_broadcast_sd(a1 + k);
        c10 = _mm256_fmadd_pd(av, b0, c10);
        c11 = _mm256_fmadd_pd(av, b1, c11);
        c12 = _mm256_fmadd_pd(av, b2, c12);
        av = _mm256_broadcast_sd(a2 + k);
        c20 = _mm256_fmadd_pd(av, b0, c20);
        c21 = _mm256_fmadd_pd(av, b1, c21);
        c22 = _mm256_fmadd_pd(av, b2, c22);
        av = _mm256_broadcast_sd(a3 + k);
        c30 = _mm256_fmadd_pd(av, b0, c30);
        c31 = _mm256_fmadd_pd(av, b1, c31);
        c32 = _mm256_fmadd_pd(av, b2, c32);
    }
    double *d0 = C + (i + 0) * N + j;
    double *d1 = C + (i + 1) * N + j;
    double *d2 = C + (i + 2) * N + j;
    double *d3 = C + (i + 3) * N + j;
    _mm256_store_pd(d0, c00); _mm256_store_pd(d0 + 4, c01);
    _mm256_store_pd(d0 + 8, c02);
    _mm256_store_pd(d1, c10); _mm256_store_pd(d1 + 4, c11);
    _mm256_store_pd(d1 + 8, c12);
    _mm256_store_pd(d2, c20); _mm256_store_pd(d2 + 4, c21);
    _mm256_store_pd(d2 + 8, c22);
    _mm256_store_pd(d3, c30); _mm256_store_pd(d3 + 4, c31);
    _mm256_store_pd(d3 + 8, c32);
}

__attribute__((target("avx2,fma"), always_inline))
static inline void kernel4x8_tail(const double *A, const double *bp,
                                  double *C, int i, int j) {
    __m256d c00 = _mm256_setzero_pd(), c01 = c00;
    __m256d c10 = c00, c11 = c00, c20 = c00, c21 = c00;
    __m256d c30 = c00, c31 = c00;
    const double *a0 = A + (i + 0) * N;
    const double *a1 = A + (i + 1) * N;
    const double *a2 = A + (i + 2) * N;
    const double *a3 = A + (i + 3) * N;
    for (int k = 0; k < N; ++k, bp += 16) {
        __m256d b0 = _mm256_load_pd(bp);
        __m256d b1 = _mm256_load_pd(bp + 4);
        __m256d av = _mm256_broadcast_sd(a0 + k);
        c00 = _mm256_fmadd_pd(av, b0, c00);
        c01 = _mm256_fmadd_pd(av, b1, c01);
        av = _mm256_broadcast_sd(a1 + k);
        c10 = _mm256_fmadd_pd(av, b0, c10);
        c11 = _mm256_fmadd_pd(av, b1, c11);
        av = _mm256_broadcast_sd(a2 + k);
        c20 = _mm256_fmadd_pd(av, b0, c20);
        c21 = _mm256_fmadd_pd(av, b1, c21);
        av = _mm256_broadcast_sd(a3 + k);
        c30 = _mm256_fmadd_pd(av, b0, c30);
        c31 = _mm256_fmadd_pd(av, b1, c31);
    }
    double *d0 = C + (i + 0) * N + j;
    double *d1 = C + (i + 1) * N + j;
    double *d2 = C + (i + 2) * N + j;
    double *d3 = C + (i + 3) * N + j;
    _mm256_store_pd(d0, c00); _mm256_store_pd(d0 + 4, c01);
    _mm256_store_pd(d1, c10); _mm256_store_pd(d1 + 4, c11);
    _mm256_store_pd(d2, c20); _mm256_store_pd(d2 + 4, c21);
    _mm256_store_pd(d3, c30); _mm256_store_pd(d3 + 4, c31);
}

#ifndef WORKER_ONLY
__attribute__((target("avx2,fma")))
void matrix_multiply(int n, const double *A, const double *B, double *C) {
    (void)n;
    double *dst;
    for (int pb = 0; pb < PANELS; pb += PACK_PANELS)
        for (int kb = 0; kb < N; kb += PACK_ROWS) {
            int pend = pb + PACK_PANELS; if (pend > PANELS) pend = PANELS;
            for (int p = pb; p < pend; ++p) {
                int j = p * 12;
                dst = packed_b + p * (N * 12) + kb * 12;
                const double *src = B + kb * N + j;
                for (int k = 0; k < PACK_ROWS; ++k, dst += 12, src += N) {
                    _mm256_store_pd(dst, _mm256_load_pd(src));
                    _mm256_store_pd(dst + 4, _mm256_load_pd(src + 4));
                    _mm256_store_pd(dst + 8, _mm256_load_pd(src + 8));
                }
            }
        }
    dst = packed_b + PANELS * (N * 12);
    for (int k = 0; k < N; ++k, dst += 16) {
        _mm256_store_pd(dst, _mm256_load_pd(B + k * N + MAIN_COLS));
        _mm256_store_pd(dst + 4, _mm256_load_pd(B + k * N + MAIN_COLS + 4));
        _mm256_store_pd(dst + 8, _mm256_load_pd(B + k * N + MAIN_COLS + 8));
        _mm256_store_pd(dst + 12, _mm256_load_pd(B + k * N + MAIN_COLS + 12));
    }
#ifdef PACK_ONLY
    return;
#endif

    for (int p = 0; p < PANELS; ++p) {
        const double *panel = packed_b + p * (N * 12);
        int j = p * 12;
        for (int i = 0; i < N; i += 4)
            kernel4x12(A, panel, C, i, j);
    }
    const double *tail = packed_b + PANELS * (N * 12);
    for (int i = 0; i < N; i += 4) {
        kernel4x8_tail(A, tail, C, i, MAIN_COLS);
        kernel4x8_tail(A, tail + 8, C, i, MAIN_COLS + 8);
    }
    _mm256_zeroupper();
}
#endif

typedef struct {
    const double *A, *B;
    double *C;
    int first, last, tail;
} Work;

__attribute__((target("avx2,fma")))
static void *run(void *opaque) {
    Work *w = (Work *)opaque;
    for (int p = w->first; p < w->last; ++p) {
        int j = p * 12;
        double *dst = packed_b + p * (N * 12);
        for (int k = 0; k < N; ++k, dst += 12) {
            _mm256_store_pd(dst, _mm256_load_pd(w->B + k * N + j));
            _mm256_store_pd(dst + 4, _mm256_load_pd(w->B + k * N + j + 4));
            _mm256_store_pd(dst + 8, _mm256_load_pd(w->B + k * N + j + 8));
        }
        const double *panel = packed_b + p * (N * 12);
        for (int i = 0; i < N; i += 4)
            kernel4x12(w->A, panel, w->C, i, j);
    }
    if (w->tail) {
        double *dst = packed_b + PANELS * (N * 12);
        for (int k = 0; k < N; ++k, dst += 16) {
            _mm256_store_pd(dst, _mm256_load_pd(w->B + k*N + MAIN_COLS));
            _mm256_store_pd(dst+4, _mm256_load_pd(w->B + k*N + MAIN_COLS+4));
            _mm256_store_pd(dst+8, _mm256_load_pd(w->B + k*N + MAIN_COLS+8));
            _mm256_store_pd(dst+12,_mm256_load_pd(w->B + k*N + MAIN_COLS+12));
        }
        const double *tail = packed_b + PANELS * (N * 12);
        for (int i = 0; i < N; i += 4) {
            kernel4x8_tail(w->A, tail, w->C, i, MAIN_COLS);
            kernel4x8_tail(w->A, tail + 8, w->C, i, MAIN_COLS + 8);
        }
    }
    _mm256_zeroupper();
    return 0;
}

void matrix_multiply(int n, const double *A, const double *B, double *C) {
    (void)n;
    Work left = {A,B,C,0,43,0}, right = {A,B,C,43,PANELS,1};
    pthread_t thread;
    int created = pthread_create(&thread, 0, run, &right) == 0;
    run(&left);
    if (created) pthread_join(thread, 0);
    else run(&right);
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #150.591 ms16 MB + 16 KBAcceptedScore: 100


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