提交记录 86352


用户 题目 状态 得分 用时 内存 语言 代码长度
pdoom mmms4k. 测测你的短整数矩阵乘法-4k Accepted 100 610.219 ms 152924 KB C++17 21.85 KB
提交时间 评测时间
2026-09-24 23:25:51 2026-09-24 23:25:56
#pragma GCC optimize("O3,unroll-loops")
#pragma GCC target("avx2")
#include <immintrin.h>
#include <stdint.h>
#include <string.h>

#ifndef BASE
#define BASE 64
#endif
using U = uint16_t;
static constexpr int MAXN = 4096;
alignas(4096) static U pa[MAXN * MAXN], pb[MAXN * MAXN];
alignas(4096) static U pc[MAXN * MAXN], ws[MAXN * MAXN];

static inline __m256i ld(const U *p) {
    return _mm256_load_si256((const __m256i *)p);
}
static inline void st(U *p, __m256i x) {
    _mm256_store_si256((__m256i *)p, x);
}

// The recursive quadrant layout makes every Strassen addition contiguous.
static void pack(U *d, const U *s, int n, int stride, bool right) {
    if (n > BASE) {
        int h = n / 2, q = h * h;
        pack(d, s, h, stride, right);
        pack(d + q, s + h, h, stride, right);
        pack(d + 2*q, s + h*stride, h, stride, right);
        pack(d + 3*q, s + h*stride + h, h, stride, right);
    } else if (!right) {
        for(int i=0;i<n;i+=4) for(int k=0;k<n;k+=2)
            for(int r=0;r<4;++r) {
                memcpy(d,s+(i+r)*stride+k,4); d+=2;
            }
    } else {
        for(int j=0;j<n;) {
            int w=n-j>=24?24:n-j;
            for(int k=0;k<n;k+=2) {
                for(int t=0;t<w;t+=8) {
                    __m128i x=_mm_load_si128((const __m128i*)(s+k*stride+j+t));
                    __m128i y=_mm_load_si128((const __m128i*)(s+(k+1)*stride+j+t));
                    __m128i lo=_mm_unpacklo_epi16(x,y),hi=_mm_unpackhi_epi16(x,y);
                    _mm_store_si128((__m128i*)d,lo);
                    _mm_store_si128((__m128i*)(d+8),hi);
                    d+=16;
                }
            }
            j+=w;
        }
    }
}

static void unpack(U *d, const U *s, int n, int stride) {
    if (n > BASE) {
        int h = n/2, q = h*h;
        unpack(d,s,h,stride);
        unpack(d+h,s+q,h,stride);
        unpack(d+h*stride,s+2*q,h,stride);
        unpack(d+h*stride+h,s+3*q,h,stride);
    } else {
        for (int i=0;i<n;++i) memcpy(d+i*stride,s+i*n,n*sizeof(U));
    }
}

static inline void micro24(const U *a,const U *b,U *c,int n) {
    const U *end=a+4*n; U *c3=c+3*n; intptr_t stride=2*n;
    asm volatile(
        "vpxor %%ymm0, %%ymm0, %%ymm0\n\t"
        "vpxor %%ymm1, %%ymm1, %%ymm1\n\t"
        "vpxor %%ymm2, %%ymm2, %%ymm2\n\t"
        "vpxor %%ymm3, %%ymm3, %%ymm3\n\t"
        "vpxor %%ymm4, %%ymm4, %%ymm4\n\t"
        "vpxor %%ymm5, %%ymm5, %%ymm5\n\t"
        "vpxor %%ymm6, %%ymm6, %%ymm6\n\t"
        "vpxor %%ymm7, %%ymm7, %%ymm7\n\t"
        "vpxor %%ymm8, %%ymm8, %%ymm8\n\t"
        "vpxor %%ymm9, %%ymm9, %%ymm9\n\t"
        "vpxor %%ymm10, %%ymm10, %%ymm10\n\t"
        "vpxor %%ymm11, %%ymm11, %%ymm11\n\t"
        "1:\n\t"
        "vmovdqa 0(%[b]), %%ymm12\n\t"
        "vmovdqa 32(%[b]), %%ymm13\n\t"
        "vpbroadcastd 0(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpmaddwd 64(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpbroadcastd 4(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
        "vpmaddwd 64(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
        "vpbroadcastd 8(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
        "vpmaddwd 64(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm8, %%ymm8\n\t"
        "vpbroadcastd 12(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm9, %%ymm9\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm10, %%ymm10\n\t"
        "vpmaddwd 64(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm11, %%ymm11\n\t"
        "vmovdqa 96(%[b]), %%ymm12\n\t"
        "vmovdqa 128(%[b]), %%ymm13\n\t"
        "vpbroadcastd 16(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpmaddwd 160(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpbroadcastd 20(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
        "vpmaddwd 160(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
        "vpbroadcastd 24(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
        "vpmaddwd 160(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm8, %%ymm8\n\t"
        "vpbroadcastd 28(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm9, %%ymm9\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm10, %%ymm10\n\t"
        "vpmaddwd 160(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm11, %%ymm11\n\t"
        "vmovdqa 192(%[b]), %%ymm12\n\t"
        "vmovdqa 224(%[b]), %%ymm13\n\t"
        "vpbroadcastd 32(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpmaddwd 256(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpbroadcastd 36(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
        "vpmaddwd 256(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
        "vpbroadcastd 40(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
        "vpmaddwd 256(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm8, %%ymm8\n\t"
        "vpbroadcastd 44(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm9, %%ymm9\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm10, %%ymm10\n\t"
        "vpmaddwd 256(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm11, %%ymm11\n\t"
        "vmovdqa 288(%[b]), %%ymm12\n\t"
        "vmovdqa 320(%[b]), %%ymm13\n\t"
        "vpbroadcastd 48(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpmaddwd 352(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpbroadcastd 52(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
        "vpmaddwd 352(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
        "vpbroadcastd 56(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
        "vpmaddwd 352(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm8, %%ymm8\n\t"
        "vpbroadcastd 60(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm9, %%ymm9\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm10, %%ymm10\n\t"
        "vpmaddwd 352(%[b]), %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm11, %%ymm11\n\t"
        "add $64, %[a]\n\t"
        "add $384, %[b]\n\t"
        "cmp %[end], %[a]\n\t"
        "jne 1b\n\t"
        "vpcmpeqd %%ymm15, %%ymm15, %%ymm15\n\t"
        "vpsrld $16, %%ymm15, %%ymm15\n\t"
        "vpand %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpand %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpand %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpackusdw %%ymm1, %%ymm0, %%ymm0\n\t"
        "vpermq $216, %%ymm0, %%ymm0\n\t"
        "vmovdqu %%ymm0, 0(%[c])\n\t"
        "vpackusdw %%ymm2, %%ymm2, %%ymm2\n\t"
        "vpermq $216, %%ymm2, %%ymm2\n\t"
        "vmovdqu %%xmm2, 32(%[c])\n\t"
        "vpand %%ymm15, %%ymm3, %%ymm3\n\t"
        "vpand %%ymm15, %%ymm4, %%ymm4\n\t"
        "vpand %%ymm15, %%ymm5, %%ymm5\n\t"
        "vpackusdw %%ymm4, %%ymm3, %%ymm3\n\t"
        "vpermq $216, %%ymm3, %%ymm3\n\t"
        "vmovdqu %%ymm3, 0(%[c],%[s],1)\n\t"
        "vpackusdw %%ymm5, %%ymm5, %%ymm5\n\t"
        "vpermq $216, %%ymm5, %%ymm5\n\t"
        "vmovdqu %%xmm5, 32(%[c],%[s],1)\n\t"
        "vpand %%ymm15, %%ymm6, %%ymm6\n\t"
        "vpand %%ymm15, %%ymm7, %%ymm7\n\t"
        "vpand %%ymm15, %%ymm8, %%ymm8\n\t"
        "vpackusdw %%ymm7, %%ymm6, %%ymm6\n\t"
        "vpermq $216, %%ymm6, %%ymm6\n\t"
        "vmovdqu %%ymm6, 0(%[c],%[s],2)\n\t"
        "vpackusdw %%ymm8, %%ymm8, %%ymm8\n\t"
        "vpermq $216, %%ymm8, %%ymm8\n\t"
        "vmovdqu %%xmm8, 32(%[c],%[s],2)\n\t"
        "vpand %%ymm15, %%ymm9, %%ymm9\n\t"
        "vpand %%ymm15, %%ymm10, %%ymm10\n\t"
        "vpand %%ymm15, %%ymm11, %%ymm11\n\t"
        "vpackusdw %%ymm10, %%ymm9, %%ymm9\n\t"
        "vpermq $216, %%ymm9, %%ymm9\n\t"
        "vmovdqu %%ymm9, 0(%[c3])\n\t"
        "vpackusdw %%ymm11, %%ymm11, %%ymm11\n\t"
        "vpermq $216, %%ymm11, %%ymm11\n\t"
        "vmovdqu %%xmm11, 32(%[c3])\n\t"
        : [a] "+&r"(a), [b] "+&r"(b)
        : [c] "r"(c), [c3] "r"(c3), [s] "r"(stride), [end] "r"(end)
        : "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5",
          "ymm6", "ymm7", "ymm8", "ymm9", "ymm10", "ymm11", "ymm12",
          "ymm13", "ymm14", "ymm15");
}

static inline void micro16(const U *a,const U *b,U *c,int n) {
    const U *end=a+4*n; U *c3=c+3*n; intptr_t stride=2*n;
    asm volatile(
        "vpxor %%ymm0, %%ymm0, %%ymm0\n\t"
        "vpxor %%ymm1, %%ymm1, %%ymm1\n\t"
        "vpxor %%ymm2, %%ymm2, %%ymm2\n\t"
        "vpxor %%ymm3, %%ymm3, %%ymm3\n\t"
        "vpxor %%ymm4, %%ymm4, %%ymm4\n\t"
        "vpxor %%ymm5, %%ymm5, %%ymm5\n\t"
        "vpxor %%ymm6, %%ymm6, %%ymm6\n\t"
        "vpxor %%ymm7, %%ymm7, %%ymm7\n\t"
        "1:\n\t"
        "vmovdqa 0(%[b]), %%ymm12\n\t"
        "vmovdqa 32(%[b]), %%ymm13\n\t"
        "vpbroadcastd 0(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpbroadcastd 4(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
        "vpbroadcastd 8(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
        "vpbroadcastd 12(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
        "vmovdqa 64(%[b]), %%ymm12\n\t"
        "vmovdqa 96(%[b]), %%ymm13\n\t"
        "vpbroadcastd 16(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpbroadcastd 20(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
        "vpbroadcastd 24(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
        "vpbroadcastd 28(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
        "vmovdqa 128(%[b]), %%ymm12\n\t"
        "vmovdqa 160(%[b]), %%ymm13\n\t"
        "vpbroadcastd 32(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpbroadcastd 36(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
        "vpbroadcastd 40(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
        "vpbroadcastd 44(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
        "vmovdqa 192(%[b]), %%ymm12\n\t"
        "vmovdqa 224(%[b]), %%ymm13\n\t"
        "vpbroadcastd 48(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpbroadcastd 52(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
        "vpbroadcastd 56(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm4, %%ymm4\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
        "vpbroadcastd 60(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm6, %%ymm6\n\t"
        "vpmaddwd %%ymm13, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
        "add $64, %[a]\n\t"
        "add $256, %[b]\n\t"
        "cmp %[end], %[a]\n\t"
        "jne 1b\n\t"
        "vpcmpeqd %%ymm15, %%ymm15, %%ymm15\n\t"
        "vpsrld $16, %%ymm15, %%ymm15\n\t"
        "vpand %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpand %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpackusdw %%ymm1, %%ymm0, %%ymm0\n\t"
        "vpermq $216, %%ymm0, %%ymm0\n\t"
        "vmovdqu %%ymm0, 0(%[c])\n\t"
        "vpand %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpand %%ymm15, %%ymm3, %%ymm3\n\t"
        "vpackusdw %%ymm3, %%ymm2, %%ymm2\n\t"
        "vpermq $216, %%ymm2, %%ymm2\n\t"
        "vmovdqu %%ymm2, 0(%[c],%[s],1)\n\t"
        "vpand %%ymm15, %%ymm4, %%ymm4\n\t"
        "vpand %%ymm15, %%ymm5, %%ymm5\n\t"
        "vpackusdw %%ymm5, %%ymm4, %%ymm4\n\t"
        "vpermq $216, %%ymm4, %%ymm4\n\t"
        "vmovdqu %%ymm4, 0(%[c],%[s],2)\n\t"
        "vpand %%ymm15, %%ymm6, %%ymm6\n\t"
        "vpand %%ymm15, %%ymm7, %%ymm7\n\t"
        "vpackusdw %%ymm7, %%ymm6, %%ymm6\n\t"
        "vpermq $216, %%ymm6, %%ymm6\n\t"
        "vmovdqu %%ymm6, 0(%[c3])\n\t"
        : [a] "+&r"(a), [b] "+&r"(b)
        : [c] "r"(c), [c3] "r"(c3), [s] "r"(stride), [end] "r"(end)
        : "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5",
          "ymm6", "ymm7", "ymm8", "ymm9", "ymm10", "ymm11", "ymm12",
          "ymm13", "ymm14", "ymm15");
}

static inline void micro8(const U *a,const U *b,U *c,int n) {
    const U *end=a+4*n; U *c3=c+3*n; intptr_t stride=2*n;
    asm volatile(
        "vpxor %%ymm0, %%ymm0, %%ymm0\n\t"
        "vpxor %%ymm1, %%ymm1, %%ymm1\n\t"
        "vpxor %%ymm2, %%ymm2, %%ymm2\n\t"
        "vpxor %%ymm3, %%ymm3, %%ymm3\n\t"
        "1:\n\t"
        "vmovdqa 0(%[b]), %%ymm12\n\t"
        "vpbroadcastd 0(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpbroadcastd 4(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpbroadcastd 8(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpbroadcastd 12(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
        "vmovdqa 32(%[b]), %%ymm12\n\t"
        "vpbroadcastd 16(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpbroadcastd 20(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpbroadcastd 24(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpbroadcastd 28(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
        "vmovdqa 64(%[b]), %%ymm12\n\t"
        "vpbroadcastd 32(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpbroadcastd 36(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpbroadcastd 40(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpbroadcastd 44(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
        "vmovdqa 96(%[b]), %%ymm12\n\t"
        "vpbroadcastd 48(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpbroadcastd 52(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpbroadcastd 56(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpbroadcastd 60(%[a]), %%ymm14\n\t"
        "vpmaddwd %%ymm12, %%ymm14, %%ymm15\n\t"
        "vpaddd %%ymm15, %%ymm3, %%ymm3\n\t"
        "add $64, %[a]\n\t"
        "add $128, %[b]\n\t"
        "cmp %[end], %[a]\n\t"
        "jne 1b\n\t"
        "vpcmpeqd %%ymm15, %%ymm15, %%ymm15\n\t"
        "vpsrld $16, %%ymm15, %%ymm15\n\t"
        "vpand %%ymm15, %%ymm0, %%ymm0\n\t"
        "vpackusdw %%ymm0, %%ymm0, %%ymm0\n\t"
        "vpermq $216, %%ymm0, %%ymm0\n\t"
        "vmovdqu %%xmm0, 0(%[c])\n\t"
        "vpand %%ymm15, %%ymm1, %%ymm1\n\t"
        "vpackusdw %%ymm1, %%ymm1, %%ymm1\n\t"
        "vpermq $216, %%ymm1, %%ymm1\n\t"
        "vmovdqu %%xmm1, 0(%[c],%[s],1)\n\t"
        "vpand %%ymm15, %%ymm2, %%ymm2\n\t"
        "vpackusdw %%ymm2, %%ymm2, %%ymm2\n\t"
        "vpermq $216, %%ymm2, %%ymm2\n\t"
        "vmovdqu %%xmm2, 0(%[c],%[s],2)\n\t"
        "vpand %%ymm15, %%ymm3, %%ymm3\n\t"
        "vpackusdw %%ymm3, %%ymm3, %%ymm3\n\t"
        "vpermq $216, %%ymm3, %%ymm3\n\t"
        "vmovdqu %%xmm3, 0(%[c3])\n\t"
        : [a] "+&r"(a), [b] "+&r"(b)
        : [c] "r"(c), [c3] "r"(c3), [s] "r"(stride), [end] "r"(end)
        : "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5",
          "ymm6", "ymm7", "ymm8", "ymm9", "ymm10", "ymm11", "ymm12",
          "ymm13", "ymm14", "ymm15");
}

static void base(const U *a,const U *b,U *c,int n) {
    int j=0;
    for(;j+24<=n;j+=24)
        for(int i=0;i<n;i+=4) micro24(a+i*n,b+j*n,c+i*n+j,n);
    if(n-j==16)
        for(int i=0;i<n;i+=4) micro16(a+i*n,b+j*n,c+i*n+j,n);
    if(n-j==8)
        for(int i=0;i<n;i+=4) micro8(a+i*n,b+j*n,c+i*n+j,n);
}

template<bool Sub>
static inline void combine(U *d,const U *a,const U *b,int len) {
    for(int i=0;i<len;i+=16) {
        __m256i x=ld(a+i),y=ld(b+i);
        st(d+i,Sub?_mm256_sub_epi16(x,y):_mm256_add_epi16(x,y));
    }
}

// Strassen-Winograd: seven products, fifteen additions, two temporaries.
// All operations are in Z/(2^16); truncation never changes the answer.
static void mul(const U *a,const U *b,U *c,int n,U *work) {
    if(n<=BASE) { base(a,b,c,n); return; }
    int h=n/2,q=h*h;
    const U *a11=a,*a12=a+q,*a21=a+2*q,*a22=a+3*q;
    const U *b11=b,*b12=b+q,*b21=b+2*q,*b22=b+3*q;
    U *c11=c,*c12=c+q,*c21=c+2*q,*c22=c+3*q;
    U *x=work,*y=work+q,*next=work+2*q;
    combine<true>(x,a11,a21,q);
    combine<true>(y,b22,b12,q);
    mul(x,y,c21,h,next);                       // P7
    combine<false>(x,a21,a22,q);
    combine<true>(y,b12,b11,q);
    mul(x,y,c22,h,next);                       // P5
    combine<true>(x,x,a11,q);
    combine<true>(y,b22,y,q);
    mul(x,y,c12,h,next);                       // P6
    combine<true>(x,a12,x,q);
    mul(x,b22,c11,h,next);                     // P3
    combine<true>(y,y,b21,q);
    mul(a22,y,x,h,next);                       // P4
    mul(a11,b11,y,h,next);                     // P1
    for(int i=0;i<q;i+=16) {
        __m256i u2=_mm256_add_epi16(ld(y+i),ld(c12+i));
        __m256i u3=_mm256_add_epi16(u2,ld(c21+i));
        __m256i p5=ld(c22+i);
        st(c12+i,_mm256_add_epi16(_mm256_add_epi16(u2,p5),ld(c11+i)));
        st(c21+i,_mm256_sub_epi16(u3,ld(x+i)));
        st(c22+i,_mm256_add_epi16(u3,p5));
    }
    mul(a12,b21,c11,h,next);                   // P2
    combine<false>(c11,c11,y,q);
}

void matrix_multiply(int n,const short *A,const short *B,short *C) {
    pack(pa,(const U *)A,n,n,false);
    pack(pb,(const U *)B,n,n,true);
    mul(pa,pb,pc,n,ws);
    unpack((U *)C,pc,n,n);
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #1610.219 ms149 MB + 348 KBAcceptedScore: 100


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