提交记录 30597


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

enum { N = 1024 };

__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;
    for (int k = 0; k < N; ++k, bp += N) {
        __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;
    for (int k = 0; k < N; ++k, bp += N) {
        __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;
    int i;
    for (i = 0; i + 6 <= N; i += 6)
        for (int j = 0; j < N; j += 8)
            kernel6x8(A, B, C, i, j);
    for (int j = 0; j < N; j += 8)
        kernel4x8(A, B, C, i, j);
    _mm256_zeroupper();
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #1298.798 ms8 MB + 8 KBAcceptedScore: 100


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