提交记录 30603


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_codex_260812 mmmd1k. 测测你的双精度矩阵乘法-1k Runtime Error 0 26.055 ms 12332 KB C 6.55 KB
提交时间 评测时间
2026-08-12 23:29:12 2026-08-12 23:29:13
#include <pthread.h>
#define WORKER_ONLY
#include <immintrin.h>

enum { N = 1024, MAIN_COLS = 1008, PANELS = 84 };
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 = packed_b;
    for (int j = 0; j < MAIN_COLS; j += 12)
        for (int k = 0; k < N; ++k, dst += 12) {
            _mm256_store_pd(dst, _mm256_load_pd(B + k * N + j));
            _mm256_store_pd(dst + 4, _mm256_load_pd(B + k * N + j + 4));
            _mm256_store_pd(dst + 8, _mm256_load_pd(B + k * N + j + 8));
        }
    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;
    pthread_create(&thread, 0, run, &right);
    run(&left);
    pthread_join(thread, 0);
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #126.055 ms12 MB + 44 KBRuntime ErrorScore: 0


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