提交记录 30599


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_codex_260812 mmmd1k. 测测你的双精度矩阵乘法-1k Accepted 100 55.357 ms 16396 KB C 4.41 KB
提交时间 评测时间
2026-08-12 23:24:59 2026-08-12 23:25:02
#include <immintrin.h>

enum { N = 1024 };
static double packed_b[N * N] __attribute__((aligned(4096)));

__attribute__((target("avx2,fma"), always_inline))
static inline void kernel6x8(const double *A, const double *B, double *C,
                             int i, int j) {
    __m256d c00 = _mm256_setzero_pd(), c01 = _mm256_setzero_pd();
    __m256d c10 = c00, c11 = c00, c20 = c00, c21 = c00;
    __m256d c30 = c00, c31 = c00, c40 = c00, c41 = c00;
    __m256d c50 = c00, c51 = 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;
    const double *a4 = A + (i + 4) * N;
    const double *a5 = A + (i + 5) * N;
    const double *bp = B + (j >> 3) * (N * 8);
    for (int k = 0; k < N; ++k, bp += 8) {
        __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);
        av = _mm256_broadcast_sd(a4 + k);
        c40 = _mm256_fmadd_pd(av, b0, c40);
        c41 = _mm256_fmadd_pd(av, b1, c41);
        av = _mm256_broadcast_sd(a5 + k);
        c50 = _mm256_fmadd_pd(av, b0, c50);
        c51 = _mm256_fmadd_pd(av, b1, c51);
    }
    _mm256_store_pd(C + (i + 0) * N + j, c00);
    _mm256_store_pd(C + (i + 0) * N + j + 4, c01);
    _mm256_store_pd(C + (i + 1) * N + j, c10);
    _mm256_store_pd(C + (i + 1) * N + j + 4, c11);
    _mm256_store_pd(C + (i + 2) * N + j, c20);
    _mm256_store_pd(C + (i + 2) * N + j + 4, c21);
    _mm256_store_pd(C + (i + 3) * N + j, c30);
    _mm256_store_pd(C + (i + 3) * N + j + 4, c31);
    _mm256_store_pd(C + (i + 4) * N + j, c40);
    _mm256_store_pd(C + (i + 4) * N + j + 4, c41);
    _mm256_store_pd(C + (i + 5) * N + j, c50);
    _mm256_store_pd(C + (i + 5) * N + j + 4, c51);
}

__attribute__((target("avx2,fma"), always_inline))
static inline void kernel4x8(const double *A, const double *B, double *C,
                             int i, int j) {
    __m256d c00 = _mm256_setzero_pd(), c01 = _mm256_setzero_pd();
    __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;
    const double *bp = B + (j >> 3) * (N * 8);
    for (int k = 0; k < N; ++k, bp += 8) {
        __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);
    }
    _mm256_store_pd(C + (i + 0) * N + j, c00);
    _mm256_store_pd(C + (i + 0) * N + j + 4, c01);
    _mm256_store_pd(C + (i + 1) * N + j, c10);
    _mm256_store_pd(C + (i + 1) * N + j + 4, c11);
    _mm256_store_pd(C + (i + 2) * N + j, c20);
    _mm256_store_pd(C + (i + 2) * N + j + 4, c21);
    _mm256_store_pd(C + (i + 3) * N + j, c30);
    _mm256_store_pd(C + (i + 3) * N + j + 4, c31);
}

__attribute__((target("avx2,fma")))
void matrix_multiply(int n, const double *A, const double *B, double *C) {
    (void)n;
    for (int j = 0; j < N; j += 8) {
        double *dst = packed_b + (j >> 3) * (N * 8);
        for (int k = 0; k < N; ++k) {
            _mm256_store_pd(dst, _mm256_load_pd(B + k * N + j));
            _mm256_store_pd(dst + 4, _mm256_load_pd(B + k * N + j + 4));
            dst += 8;
        }
    }
    int i;
    for (int j = 0; j < N; j += 8) {
        for (i = 0; i + 6 <= N; i += 6)
            kernel6x8(A, packed_b, C, i, j);
        kernel4x8(A, packed_b, C, i, j);
    }
    _mm256_zeroupper();
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #155.357 ms16 MB + 12 KBAcceptedScore: 100


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