提交记录 119671


用户 题目 状态 得分 用时 内存 语言 代码长度
saffah_cc_v41_260924 mmmi4k. 测测你的整数矩阵乘法-4k Accepted 100 1.778 s 262016 KB C++17 34.41 KB
提交时间 评测时间
2026-10-02 02:52:32 2026-10-02 02:52:38
/* FILE-SCOPE FLAGS: judge compiles with g++ -O2 -static -nostdlib -U_FORTIFY_SOURCE and NO
 * -march, so the baseline is x86-64 = SSE2 ONLY. Widen the target for EVERY function in
 * the file (including all packing / C-traffic code) and lift -O2 to -O3 + unroll. */
#pragma GCC optimize("O3")
#pragma GCC target("avx2,bmi,bmi2,popcnt,lzcnt")
/* int32 GEMM exact mod 2^32, AVX2.  Three-madd vpmaddwd mikro (KC=128 baked in),
 * wrapped in EXACT Strassen over Z/2^32 (7/8 of the multiply work per level).
 * GENERATED by gen_strassen.py -- do not edit by hand.
 */
#include <immintrin.h>
#include <string.h>
#include <stdint.h>
#include <stdlib.h>

#define MR 4
#define NCG 1
#define KC 128
#define MC 4
#define NC 128
#define NPAIR (KC / 2)
#define APSTR  (NPAIR * 16 + 16)
#define NR (8 * NCG)
#define BASE 128
typedef int32_t T;

#define AVX2 __attribute__((target("avx2")))
#define LD(p) _mm256_load_si256((const __m256i *)(p))
#define LDA(p) _mm256_load_si256((const __m256i *)(p))
#define STA(p, v) _mm256_store_si256((__m256i *)(p), (v))
#define ST(p, v) _mm256_store_si256((__m256i *)(p), (v))

static int32_t Apk[(MC + 4) * APSTR + 64] __attribute__((aligned(64)));
static int32_t Bpk[16896] __attribute__((aligned(64)));  /* [jg][m][B0|B1] */
#define NTSEL 0
/* MIKRO-BEGIN */
#define AVX2 __attribute__((target("avx2")))
static __m256i LA[128] __attribute__((aligned(32)));
AVX2 __attribute__((always_inline)) static inline void mikro(int npair, const int32_t *ap, const int32_t *bp,
                              int32_t *c, int ldc, int first) {
  (void)npair;
  __m256i q0, q1, q2, q3;
  __asm__ volatile(
    "vpxor %%xmm15, %%xmm15, %%xmm15\n\t"
    "vmovdqa %%ymm15, %%ymm0\n\t"
    "vmovdqa %%ymm15, %%ymm4\n\t"
    "vmovdqa %%ymm15, %%ymm1\n\t"
    "vmovdqa %%ymm15, %%ymm5\n\t"
    "vmovdqa %%ymm15, %%ymm2\n\t"
    "vmovdqa %%ymm15, %%ymm6\n\t"
    "vmovdqa %%ymm15, %%ymm3\n\t"
    "vmovdqa %%ymm15, %%ymm7\n\t"
    "test %[cnt], %[cnt]\n\t"
    "jle 2f\n\t"
    "1:\n\t"
    "prefetcht0 1024(%[bp])\n\t"
    "vmovdqu 0(%[bp]), %%ymm8\n\t"
    "vmovdqu 32(%[bp]), %%ymm9\n\t"
    "vpmaddwd 0(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpmaddwd 4160+0(%[ap]), %%ymm8, %%ymm11\n\t"
    "vpmaddwd 8320+0(%[ap]), %%ymm8, %%ymm12\n\t"
    "vpmaddwd 12480+0(%[ap]), %%ymm8, %%ymm13\n\t"
    "vpmaddwd 0(%[ap]), %%ymm9, %%ymm14\n\t"
    "vpmaddwd 4160+0(%[ap]), %%ymm9, %%ymm15\n\t"
    "vpaddd %%ymm10, %%ymm0, %%ymm0\n\t"
    "vpmaddwd 8320+0(%[ap]), %%ymm9, %%ymm10\n\t"
    "vpaddd %%ymm11, %%ymm1, %%ymm1\n\t"
    "vpmaddwd 12480+0(%[ap]), %%ymm9, %%ymm11\n\t"
    "vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
    "vpmaddwd 32+0(%[ap]), %%ymm8, %%ymm12\n\t"
    "vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
    "vpmaddwd 4160+32+0(%[ap]), %%ymm8, %%ymm13\n\t"
    "vpaddd %%ymm14, %%ymm4, %%ymm4\n\t"
    "vpmaddwd 8320+32+0(%[ap]), %%ymm8, %%ymm14\n\t"
    "vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
    "vpmaddwd 12480+32+0(%[ap]), %%ymm8, %%ymm15\n\t"
    "vmovdqu 64(%[bp]), %%ymm8\n\t"
    "vmovdqu 96(%[bp]), %%ymm9\n\t"
    "vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
    "vpmaddwd 64+0(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm11, %%ymm7, %%ymm7\n\t"
    "vpmaddwd 4160+64+0(%[ap]), %%ymm8, %%ymm11\n\t"
    "vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
    "vpmaddwd 8320+64+0(%[ap]), %%ymm8, %%ymm12\n\t"
    "vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
    "vpmaddwd 12480+64+0(%[ap]), %%ymm8, %%ymm13\n\t"
    "vpaddd %%ymm14, %%ymm6, %%ymm6\n\t"
    "vpmaddwd 64+0(%[ap]), %%ymm9, %%ymm14\n\t"
    "vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
    "vpmaddwd 4160+64+0(%[ap]), %%ymm9, %%ymm15\n\t"
    "vpaddd %%ymm10, %%ymm0, %%ymm0\n\t"
    "vpmaddwd 8320+64+0(%[ap]), %%ymm9, %%ymm10\n\t"
    "vpaddd %%ymm11, %%ymm1, %%ymm1\n\t"
    "vpmaddwd 12480+64+0(%[ap]), %%ymm9, %%ymm11\n\t"
    "vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
    "vpmaddwd 96+0(%[ap]), %%ymm8, %%ymm12\n\t"
    "vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
    "vpmaddwd 4160+96+0(%[ap]), %%ymm8, %%ymm13\n\t"
    "vpaddd %%ymm14, %%ymm4, %%ymm4\n\t"
    "vpmaddwd 8320+96+0(%[ap]), %%ymm8, %%ymm14\n\t"
    "vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
    "vpmaddwd 12480+96+0(%[ap]), %%ymm8, %%ymm15\n\t"
    "vmovdqu 128(%[bp]), %%ymm8\n\t"
    "vmovdqu 160(%[bp]), %%ymm9\n\t"
    "vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
    "vpmaddwd 128+0(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm11, %%ymm7, %%ymm7\n\t"
    "vpmaddwd 4160+128+0(%[ap]), %%ymm8, %%ymm11\n\t"
    "vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
    "vpmaddwd 8320+128+0(%[ap]), %%ymm8, %%ymm12\n\t"
    "vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
    "vpmaddwd 12480+128+0(%[ap]), %%ymm8, %%ymm13\n\t"
    "vpaddd %%ymm14, %%ymm6, %%ymm6\n\t"
    "vpmaddwd 128+0(%[ap]), %%ymm9, %%ymm14\n\t"
    "vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
    "vpmaddwd 4160+128+0(%[ap]), %%ymm9, %%ymm15\n\t"
    "vpaddd %%ymm10, %%ymm0, %%ymm0\n\t"
    "vpmaddwd 8320+128+0(%[ap]), %%ymm9, %%ymm10\n\t"
    "vpaddd %%ymm11, %%ymm1, %%ymm1\n\t"
    "vpmaddwd 12480+128+0(%[ap]), %%ymm9, %%ymm11\n\t"
    "vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
    "vpmaddwd 160+0(%[ap]), %%ymm8, %%ymm12\n\t"
    "vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
    "vpmaddwd 4160+160+0(%[ap]), %%ymm8, %%ymm13\n\t"
    "vpaddd %%ymm14, %%ymm4, %%ymm4\n\t"
    "vpmaddwd 8320+160+0(%[ap]), %%ymm8, %%ymm14\n\t"
    "vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
    "vpmaddwd 12480+160+0(%[ap]), %%ymm8, %%ymm15\n\t"
    "vmovdqu 192(%[bp]), %%ymm8\n\t"
    "vmovdqu 224(%[bp]), %%ymm9\n\t"
    "vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
    "vpmaddwd 192+0(%[ap]), %%ymm8, %%ymm10\n\t"
    "vpaddd %%ymm11, %%ymm7, %%ymm7\n\t"
    "vpmaddwd 4160+192+0(%[ap]), %%ymm8, %%ymm11\n\t"
    "vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
    "vpmaddwd 8320+192+0(%[ap]), %%ymm8, %%ymm12\n\t"
    "vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
    "vpmaddwd 12480+192+0(%[ap]), %%ymm8, %%ymm13\n\t"
    "vpaddd %%ymm14, %%ymm6, %%ymm6\n\t"
    "vpmaddwd 192+0(%[ap]), %%ymm9, %%ymm14\n\t"
    "vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
    "vpmaddwd 4160+192+0(%[ap]), %%ymm9, %%ymm15\n\t"
    "vpaddd %%ymm10, %%ymm0, %%ymm0\n\t"
    "vpmaddwd 8320+192+0(%[ap]), %%ymm9, %%ymm10\n\t"
    "vpaddd %%ymm11, %%ymm1, %%ymm1\n\t"
    "vpmaddwd 12480+192+0(%[ap]), %%ymm9, %%ymm11\n\t"
    "vpaddd %%ymm12, %%ymm2, %%ymm2\n\t"
    "vpmaddwd 224+0(%[ap]), %%ymm8, %%ymm12\n\t"
    "vpaddd %%ymm13, %%ymm3, %%ymm3\n\t"
    "vpmaddwd 4160+224+0(%[ap]), %%ymm8, %%ymm13\n\t"
    "vpaddd %%ymm14, %%ymm4, %%ymm4\n\t"
    "vpmaddwd 8320+224+0(%[ap]), %%ymm8, %%ymm14\n\t"
    "vpaddd %%ymm15, %%ymm5, %%ymm5\n\t"
    "vpmaddwd 12480+224+0(%[ap]), %%ymm8, %%ymm15\n\t"
    "vpaddd %%ymm10, %%ymm6, %%ymm6\n\t"
    "vpaddd %%ymm11, %%ymm7, %%ymm7\n\t"
    "vpaddd %%ymm12, %%ymm4, %%ymm4\n\t"
    "vpaddd %%ymm13, %%ymm5, %%ymm5\n\t"
    "vpaddd %%ymm14, %%ymm6, %%ymm6\n\t"
    "vpaddd %%ymm15, %%ymm7, %%ymm7\n\t"
    "add $256, %[ap]\n\t"
    "add $256, %[bp]\n\t"
    "sub $4, %[cnt]\n\t"
    "jnz 1b\n\t"
    "2:\n\t"
    "vpslld $16, %%ymm4, %%ymm4\n\t"
    "vpaddd %%ymm4, %%ymm0, %[q0]\n\t"
    "vpslld $16, %%ymm5, %%ymm5\n\t"
    "vpaddd %%ymm5, %%ymm1, %[q1]\n\t"
    "vpslld $16, %%ymm6, %%ymm6\n\t"
    "vpaddd %%ymm6, %%ymm2, %[q2]\n\t"
    "vpslld $16, %%ymm7, %%ymm7\n\t"
    "vpaddd %%ymm7, %%ymm3, %[q3]\n\t"
    ""
    : [ap] "+r" (ap), [bp] "+r" (bp), [cnt] "+r" (npair),
      [q0] "=x" (q0), [q1] "=x" (q1), [q2] "=x" (q2), [q3] "=x" (q3)
    : [la] "r" (&LA[0])
    : "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5", "ymm6", "ymm7", "ymm8", "ymm9", "ymm10", "ymm15", "memory");
  { int32_t *cp = c + (size_t)0 * ldc + 0;
    __m256i rs = q0;
    if (first) _mm256_store_si256((__m256i *)cp, rs);
    else _mm256_store_si256((__m256i *)cp, _mm256_add_epi32(_mm256_load_si256((const __m256i *)cp), rs)); }
  { int32_t *cp = c + (size_t)1 * ldc + 0;
    __m256i rs = q1;
    if (first) _mm256_store_si256((__m256i *)cp, rs);
    else _mm256_store_si256((__m256i *)cp, _mm256_add_epi32(_mm256_load_si256((const __m256i *)cp), rs)); }
  { int32_t *cp = c + (size_t)2 * ldc + 0;
    __m256i rs = q2;
    if (first) _mm256_store_si256((__m256i *)cp, rs);
    else _mm256_store_si256((__m256i *)cp, _mm256_add_epi32(_mm256_load_si256((const __m256i *)cp), rs)); }
  { int32_t *cp = c + (size_t)3 * ldc + 0;
    __m256i rs = q3;
    if (first) _mm256_store_si256((__m256i *)cp, rs);
    else _mm256_store_si256((__m256i *)cp, _mm256_add_epi32(_mm256_load_si256((const __m256i *)cp), rs)); }
}
/* MIKRO-END */
/* ---- ld-aware packing + tiled macro-kernel (base case) ---- */
AVX2 static inline void packA(int ic, int kk, int n, int lda, const T *A) {
  const __m256i m16 = _mm256_set1_epi32(0xFFFF);
  const __m256i half = _mm256_set1_epi32(0x8000);
  const __m256i i0 = _mm256_set1_epi32(0), i1 = _mm256_set1_epi32(1);
  const __m256i i4 = _mm256_set1_epi32(4), i5 = _mm256_set1_epi32(5);
  (void)n;
  for (int r = 0; r < MC; r++) {
    const T *src = A + (size_t)(ic + r) * lda + kk;
    T *dst = Apk + (size_t)r * APSTR;
    for (int m = 0; m < KC / 8; m++) {
      __m256i x = LD(src + 8 * m);
      __m256i al = _mm256_and_si256(x, m16);
      __m256i ah = _mm256_srli_epi32(_mm256_add_epi32(x, half), 16);
      __m256i pl = _mm256_packus_epi32(al, al);
      __m256i ph = _mm256_packus_epi32(ah, ah);
      T *d = dst + 64 * m;
      ST(d + 0, _mm256_permutevar8x32_epi32(pl, i0));
      ST(d + 8, _mm256_permutevar8x32_epi32(ph, i0));
      ST(d + 16, _mm256_permutevar8x32_epi32(pl, i1));
      ST(d + 24, _mm256_permutevar8x32_epi32(ph, i1));
      ST(d + 32, _mm256_permutevar8x32_epi32(pl, i4));
      ST(d + 40, _mm256_permutevar8x32_epi32(ph, i4));
      ST(d + 48, _mm256_permutevar8x32_epi32(pl, i5));
      ST(d + 56, _mm256_permutevar8x32_epi32(ph, i5));
    }
  }
}

AVX2 static inline void packB(int jc, int kk, int n, int ldb, const T *B) {
  const __m256i m16 = _mm256_set1_epi32(0xFFFF);
  const __m256i half = _mm256_set1_epi32(0x8000);
  (void)n;
  for (int m = 0; m < NPAIR; m++) {
    const T *s0 = B + (size_t)(kk + 2 * m) * ldb + jc;
    const T *s1 = B + (size_t)(kk + 2 * m + 1) * ldb + jc;
    T *d0 = Bpk + (size_t)m * 16;
    T *d1 = d0 + 8;
    for (int j = 0; j < NC; j += 8) {
      __m256i x0 = LD(s0 + j), x1 = LD(s1 + j);
      __m256i l0 = _mm256_and_si256(x0, m16);
      __m256i l1 = _mm256_and_si256(x1, m16);
      __m256i h0 = _mm256_srli_epi32(_mm256_add_epi32(x0, half), 16);
      __m256i h1 = _mm256_srli_epi32(_mm256_add_epi32(x1, half), 16);
      ST(d0 + (size_t)(j / 8) * (NPAIR * 16 + 16), _mm256_or_si256(l0, _mm256_slli_epi32(l1, 16)));
      ST(d1 + (size_t)(j / 8) * (NPAIR * 16 + 16), _mm256_or_si256(h0, _mm256_slli_epi32(h1, 16)));
    }
  }
}


AVX2 static inline void packA2(int ic, int kk, int n, int lda, const T *A, int sg, const T *A2) {
  const __m256i m16 = _mm256_set1_epi32(0xFFFF);
  const __m256i half = _mm256_set1_epi32(0x8000);
  const __m256i i0 = _mm256_set1_epi32(0), i1 = _mm256_set1_epi32(1);
  const __m256i i4 = _mm256_set1_epi32(4), i5 = _mm256_set1_epi32(5);
  (void)n;
  for (int r = 0; r < MC; r++) {
    const T *src = A + (size_t)(ic + r) * lda + kk;
    const T *src2 = A2 ? (A2 + (size_t)(ic + r) * lda + kk) : 0;
    T *dst = Apk + (size_t)r * APSTR;
    for (int m = 0; m < KC / 8; m++) {
      __builtin_prefetch(src + lda + 8 * m, 0, 1);
      __builtin_prefetch(src + 8 * m + 128, 0, 1);
      if (src2) { __builtin_prefetch(src2 + lda + 8 * m, 0, 1);
                  __builtin_prefetch(src2 + 8 * m + 128, 0, 1); }
      __m256i x = LD(src + 8 * m);
      if (src2) { __m256i y = LD(src2 + 8 * m);
                  x = (sg > 0) ? _mm256_add_epi32(x, y) : _mm256_sub_epi32(x, y); }
      __m256i al = _mm256_and_si256(x, m16);
      __m256i ah = _mm256_srli_epi32(_mm256_add_epi32(x, half), 16);
      __m256i pl = _mm256_packus_epi32(al, al);
      __m256i ph = _mm256_packus_epi32(ah, ah);
      T *d = dst + 64 * m;
      ST(d + 0, _mm256_permutevar8x32_epi32(pl, i0));
      ST(d + 8, _mm256_permutevar8x32_epi32(ph, i0));
      ST(d + 16, _mm256_permutevar8x32_epi32(pl, i1));
      ST(d + 24, _mm256_permutevar8x32_epi32(ph, i1));
      ST(d + 32, _mm256_permutevar8x32_epi32(pl, i4));
      ST(d + 40, _mm256_permutevar8x32_epi32(ph, i4));
      ST(d + 48, _mm256_permutevar8x32_epi32(pl, i5));
      ST(d + 56, _mm256_permutevar8x32_epi32(ph, i5));
    }
  }
}

AVX2 static inline void packB2(int jc, int kk, int n, int ldb, const T *B, int sg, const T *B2) {
  const __m256i m16 = _mm256_set1_epi32(0xFFFF);
  const __m256i half = _mm256_set1_epi32(0x8000);
  (void)n;
  const size_t DSTR = (size_t)(NPAIR * 16 + 16);
  const int sgn = sg > 0;
  for (int m = 0; m < NPAIR; m++) {
    const T *s0 = B + (size_t)(kk + 2 * m) * ldb + jc;
    const T *s1 = s0 + ldb;
    const T *t0 = B2 ? (B2 + (size_t)(kk + 2 * m) * ldb + jc) : 0;
    const T *t1 = t0 ? t0 + ldb : 0;
    T *d0 = Bpk + (size_t)m * 16;
    T *d1 = d0 + 8;
    for (int j = 0; j < NC; j += 8, d0 += DSTR, d1 += DSTR) {
      __m256i x0 = LD(s0 + j), x1 = LD(s1 + j);
      if (t0) {
        __m256i y0 = LD(t0 + j), y1 = LD(t1 + j);
        if (sgn) { x0 = _mm256_add_epi32(x0, y0); x1 = _mm256_add_epi32(x1, y1); }
        else     { x0 = _mm256_sub_epi32(x0, y0); x1 = _mm256_sub_epi32(x1, y1); }
      }
      __m256i l0 = _mm256_and_si256(x0, m16);
      __m256i l1 = _mm256_and_si256(x1, m16);
      __m256i h0 = _mm256_srli_epi32(_mm256_add_epi32(x0, half), 16);
      __m256i h1 = _mm256_srli_epi32(_mm256_add_epi32(x1, half), 16);
      ST(d0, _mm256_or_si256(l0, _mm256_slli_epi32(l1, 16)));
      ST(d1, _mm256_or_si256(h0, _mm256_slli_epi32(h1, 16)));
    }
  }
}


AVX2 __attribute__((always_inline)) static inline void block_mikro(int ic, int jc, int n, int ldc, T *C, int first) {
  T *cp0 = C + (size_t)ic * ldc + jc;
  for (int r = 0; r < MC; r += MR) {
    const int32_t *ap = Apk + (size_t)r * APSTR;
    T *cp = cp0 + (size_t)r * ldc;
    const int32_t *bp = Bpk;
    for (int jg = 0; jg < NC; jg += NR) {
      mikro(NPAIR, ap, bp, cp + jg, ldc, first);
      bp += (NPAIR * 16 + 16);
    }
  }
}

/* LEAF SPECIALISED: with BASE==KC==NC==128 every tiled2 leaf call has h==128 and
   ldc==h==128, so both become compile-time constants -> the jc/kk/ic loops fully fold. */
AVX2 static void tiled2_128(const T *A1, const T *A2, int sgA, int lda,
                            const T *B1, const T *B2, int sgB, int ldb, T *C) {
  const int n = 128;
  const int ldc = 128;
  for (int jc = 0; jc < n; jc += NC) {
    for (int kk = 0; kk < n; kk += KC) {
      packB2(jc, kk, n, ldb, B1, sgB, B2);
      for (int ic = 0; ic < n; ic += MC) {
        packA2(ic, kk, n, lda, A1, sgA, A2);
        block_mikro(ic, jc, n, ldc, C, kk == 0);
      }
    }
  }
}

/* correct for any n, just slow */
AVX2 static void generic(int n, const T *A, int lda, const T *B, int ldb, T *C, int ldc) {
  for (int i = 0; i < n; i++)
    for (int j = 0; j < n; j++) {
      uint32_t s = 0;
      for (int k = 0; k < n; k++)
        s += (uint32_t)A[(size_t)i * lda + k] * (uint32_t)B[(size_t)k * ldb + j];
      C[(size_t)i * ldc + j] = (T)s;
    }
}

/* tiled base case: requires n % NC == 0, n % KC == 0, n >= MC */
AVX2 static void tiled(int n, const T *A, int lda, const T *B, int ldb, T *C, int ldc) {
  /* C is fully written by the kk==0 pass (micro stores instead of RMWs) */
  for (int jc = 0; jc < n; jc += NC) {
    for (int kk = 0; kk < n; kk += KC) {
      packB(jc, kk, n, ldb, B);
      for (int ic = 0; ic < n; ic += MC) {
        packA(ic, kk, n, lda, A);
        block_mikro(ic, jc, n, ldc, C, kk == 0);
      }
    }
  }
}

AVX2 static int ok_tiled(int n) {
  return (n % NC) == 0 && (n % KC) == 0 && (n % MC) == 0 && (MC % MR) == 0 && (NC % NR) == 0;
}

/* ---- Strassen over Z/2^32.  Exact: all ops are ring ops of Z/2^32. ---- */
#define STR_DEPTH 6

/* P = X + s*Y in ONE pass over h x h  (replaces copyblk + axpy: one write instead of two) */
AVX2 static void addblk(int h, T *P, const T *X, int ldx, const T *Y, int ldy, int s) {
  for (int i = 0; i < h; i++) {
    const T *xp = X + (size_t)i * ldx, *yp = Y + (size_t)i * ldy;
    T *pp = P + (size_t)i * h;
    __builtin_prefetch(X + (size_t)(i+2)*ldx, 0, 3); __builtin_prefetch(Y + (size_t)(i+2)*ldy, 0, 3);
    int j = 0;
    if (s > 0) for (; j + 8 <= h; j += 8) STA(pp + j, _mm256_add_epi32(LDA(xp + j), LDA(yp + j)));
    else       for (; j + 8 <= h; j += 8) STA(pp + j, _mm256_sub_epi32(LDA(xp + j), LDA(yp + j)));
    for (; j < h; j++) pp[j] = (T)(s > 0 ? xp[j] + yp[j] : xp[j] - yp[j]);
  }
}
/* dst = s1*m1 + s2*m2 , ONE pass (M blocks are contiguous, row stride h) */
AVX2 static void comb2(int h, T *dst, int ldd, const T *m1, int s1, const T *m2, int s2) {
  (void)s1;
  for (int i = 0; i < h; i++) {
    const T *p1 = m1 + (size_t)i * h, *p2 = m2 + (size_t)i * h;
    T *dp = dst + (size_t)i * ldd;
    int j = 0;
    for (; j + 8 <= h; j += 8) {
      __m256i v = LDA(p1 + j);
      v = (s2 > 0) ? _mm256_add_epi32(v, LDA(p2 + j)) : _mm256_sub_epi32(v, LDA(p2 + j));
      _mm256_stream_si256((__m256i*)(dp + j), v);
    }
    for (; j < h; j++) dp[j] = (T)(p1[j] + (s2 > 0 ? p2[j] : -p2[j]));
  }
}
/* dst = s1*m1 + s2*m2 + s3*m3 + s4*m4 , ONE pass */
AVX2 static void comb4(int h, T *dst, int ldd, const T *m1, int s1, const T *m2, int s2,
                       const T *m3, int s3, const T *m4, int s4) {
  (void)s1;
  for (int i = 0; i < h; i++) {
    const T *p1 = m1 + (size_t)i * h, *p2 = m2 + (size_t)i * h,
            *p3 = m3 + (size_t)i * h, *p4 = m4 + (size_t)i * h;
    T *dp = dst + (size_t)i * ldd;
    int j = 0;
    for (; j + 8 <= h; j += 8) {
      __m256i v = LDA(p1 + j);
      v = (s2 > 0) ? _mm256_add_epi32(v, LDA(p2 + j)) : _mm256_sub_epi32(v, LDA(p2 + j));
      v = (s3 > 0) ? _mm256_add_epi32(v, LDA(p3 + j)) : _mm256_sub_epi32(v, LDA(p3 + j));
      v = (s4 > 0) ? _mm256_add_epi32(v, LDA(p4 + j)) : _mm256_sub_epi32(v, LDA(p4 + j));
      _mm256_stream_si256((__m256i*)(dp + j), v);
    }
    for (; j < h; j++)
      dp[j] = (T)(p1[j] + (s2>0?p2[j]:-p2[j]) + (s3>0?p3[j]:-p3[j]) + (s4>0?p4[j]:-p4[j]));
  }
}

#define B11 (A)
#define B12 (A + h)
#define B21 (A + (size_t)h * lda)
#define B22 (A + (size_t)h * lda + h)
#define D11 (B)
#define D12 (B + h)
#define D21 (B + (size_t)h * ldb)
#define D22 (B + (size_t)h * ldb + h)


/* ---- FUSED combine: one pass instead of four (l_sub1pct) --------------------
   The shipped rec() makes FOUR passes over the seven M blocks (12 block-reads
   for 4 outputs: M1..M5 are each read twice).  Holding the seven row vectors in
   registers lets all four quadrants be formed from 7 loads, so the combine's
   issue cost per 8 outputs goes 24 -> 19 uops and its read volume 12h^2 -> 7h^2.
   Same stores, same NT form, byte-exact.                                     */
static inline void STPAIR(int nt, T *P, __m256i u, __m256i v) {
      _mm256_stream_si256((__m256i *)(P), u); _mm256_stream_si256((__m256i *)(P) + 1, v);
}
AVX2 static void comb_all(int h, T *C, int ldc,
     const T *M1, const T *M2, const T *M3, const T *M4,
     const T *M5, const T *M6, const T *M7) {
  for (int i = 0; i < h; i++) {
    const T *a = M1 + (size_t)i * h, *b = M2 + (size_t)i * h, *c = M3 + (size_t)i * h,
            *d = M4 + (size_t)i * h, *e = M5 + (size_t)i * h, *f = M6 + (size_t)i * h,
            *g = M7 + (size_t)i * h;
    T *d1 = C + (size_t)i * ldc, *d2 = d1 + h,
      *d3 = d1 + (size_t)h * ldc, *d4 = d3 + h;
    int j = 0;
    for (; j + 16 <= h; j += 16) {
      __builtin_prefetch(a + j + 256, 0, 3); __builtin_prefetch(b + j + 256, 0, 3);
      __builtin_prefetch(c + j + 256, 0, 3); __builtin_prefetch(d + j + 256, 0, 3);
      __builtin_prefetch(e + j + 256, 0, 3); __builtin_prefetch(f + j + 256, 0, 3);
      __builtin_prefetch(g + j + 256, 0, 3);
      __m256i a1 = LDA(a + j), a2 = LDA(a + j + 8),
              b1 = LDA(b + j), b2 = LDA(b + j + 8),
              c1 = LDA(c + j), c2 = LDA(c + j + 8),
              d1v = LDA(d + j), d2v = LDA(d + j + 8),
              e1 = LDA(e + j), e2 = LDA(e + j + 8),
              f1 = LDA(f + j), f2 = LDA(f + j + 8),
              g1 = LDA(g + j), g2 = LDA(g + j + 8);
      STPAIR(NTSEL, d1 + j, _mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(a1, d1v), e1), g1),
                    _mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(a2, d2v), e2), g2));
      STPAIR(NTSEL, d2 + j, _mm256_add_epi32(c1, e1), _mm256_add_epi32(c2, e2));
      STPAIR(NTSEL, d3 + j, _mm256_add_epi32(b1, d1v), _mm256_add_epi32(b2, d2v));
      STPAIR(NTSEL, d4 + j, _mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(a1, c1), b1), f1),
                    _mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(a2, c2), b2), f2));
    }
    for (; j + 8 <= h; j += 8) {
      __m256i m1 = LDA(a + j), m2 = LDA(b + j), m3 = LDA(c + j), m4 = LDA(d + j),
              m5 = LDA(e + j), m6 = LDA(f + j), m7 = LDA(g + j);
      _mm256_stream_si256((__m256i *)(d1 + j),
        _mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(m1, m4), m5), m7));
      _mm256_stream_si256((__m256i *)(d2 + j), _mm256_add_epi32(m3, m5));
      _mm256_stream_si256((__m256i *)(d3 + j), _mm256_add_epi32(m2, m4));
      _mm256_stream_si256((__m256i *)(d4 + j),
        _mm256_add_epi32(_mm256_sub_epi32(_mm256_add_epi32(m1, m3), m2), m6));
    }
    for (; j < h; j++) {
      d1[j] = a[j] + d[j] - e[j] + g[j];
      d2[j] = c[j] + e[j];
      d3[j] = b[j] + d[j];
      d4[j] = a[j] - b[j] + c[j] + f[j];
    }
  }
}

/* ---- WINOGRAD COMBINE: C11 = M1+M2 ; C12 = M1+M3+M5+M6 ; C21 = M1+M4+M6+M7 ;
   C22 = M1+M5+M6+M7.  Same 8 additions and same 7 block-reads as comb_all;
   only the association differs (the shared M1+M6 form is computed once).      */
AVX2 static void comb_all_w(int h, T *C, int ldc,
     const T *M1, const T *M2, const T *M3, const T *M4,
     const T *M5, const T *M6, const T *M7) {
  for (int i = 0; i < h; i++) {
    const T *a = M1 + (size_t)i * h, *b = M2 + (size_t)i * h, *c = M3 + (size_t)i * h,
            *d = M4 + (size_t)i * h, *e = M5 + (size_t)i * h, *f = M6 + (size_t)i * h,
            *g = M7 + (size_t)i * h;
    T *d1 = C + (size_t)i * ldc, *d2 = d1 + h,
      *d3 = d1 + (size_t)h * ldc, *d4 = d3 + h;
    int j = 0;
    for (; j + 16 <= h; j += 16) {
      __builtin_prefetch(a + j + 256, 0, 3); __builtin_prefetch(b + j + 256, 0, 3);
      __builtin_prefetch(c + j + 256, 0, 3); __builtin_prefetch(d + j + 256, 0, 3);
      __builtin_prefetch(e + j + 256, 0, 3); __builtin_prefetch(f + j + 256, 0, 3);
      __builtin_prefetch(g + j + 256, 0, 3);
      __m256i a1 = LDA(a + j), a2 = LDA(a + j + 8),
              b1 = LDA(b + j), b2 = LDA(b + j + 8),
              c1 = LDA(c + j), c2 = LDA(c + j + 8),
              d1v = LDA(d + j), d2v = LDA(d + j + 8),
              e1 = LDA(e + j), e2 = LDA(e + j + 8),
              f1 = LDA(f + j), f2 = LDA(f + j + 8),
              g1 = LDA(g + j), g2 = LDA(g + j + 8);
      __m256i w1 = _mm256_add_epi32(a1, f1), w2 = _mm256_add_epi32(a2, f2);
      STPAIR(NTSEL, d1 + j, _mm256_add_epi32(a1, b1), _mm256_add_epi32(a2, b2));
      STPAIR(NTSEL, d2 + j, _mm256_add_epi32(_mm256_add_epi32(w1, c1), e1),
                           _mm256_add_epi32(_mm256_add_epi32(w2, c2), e2));
      STPAIR(NTSEL, d3 + j, _mm256_add_epi32(_mm256_add_epi32(w1, d1v), g1),
                           _mm256_add_epi32(_mm256_add_epi32(w2, d2v), g2));
      STPAIR(NTSEL, d4 + j, _mm256_add_epi32(_mm256_add_epi32(w1, e1), g1),
                           _mm256_add_epi32(_mm256_add_epi32(w2, e2), g2));
    }
    for (; j + 8 <= h; j += 8) {
      __m256i m1 = LDA(a + j), m2 = LDA(b + j), m3 = LDA(c + j), m4 = LDA(d + j),
              m5 = LDA(e + j), m6 = LDA(f + j), m7 = LDA(g + j);
      __m256i w = _mm256_add_epi32(m1, m6);
      _mm256_stream_si256((__m256i *)(d1 + j), _mm256_add_epi32(m1, m2));
      _mm256_stream_si256((__m256i *)(d2 + j), _mm256_add_epi32(_mm256_add_epi32(w, m3), m5));
      _mm256_stream_si256((__m256i *)(d3 + j), _mm256_add_epi32(_mm256_add_epi32(w, m4), m7));
      _mm256_stream_si256((__m256i *)(d4 + j), _mm256_add_epi32(_mm256_add_epi32(w, m5), m7));
    }
    for (; j < h; j++) {
      d1[j] = a[j] + b[j];
      d2[j] = a[j] + c[j] + e[j] + f[j];
      d3[j] = a[j] + d[j] + f[j] + g[j];
      d4[j] = a[j] + e[j] + f[j] + g[j];
    }
  }
}



/* ============ PER-LOP SPECIALISATION: constant signs, constant 2nd-operand presence ============ */
template<int SA>
AVX2 static inline void packA2_t(int ic, int kk, int n, int lda, const T *A, const T *A2) {
  const __m256i m16 = _mm256_set1_epi32(0xFFFF);
  const __m256i half = _mm256_set1_epi32(0x8000);
  const __m256i i0 = _mm256_set1_epi32(0), i1 = _mm256_set1_epi32(1);
  const __m256i i4 = _mm256_set1_epi32(4), i5 = _mm256_set1_epi32(5);
  (void)n;
  for (int r = 0; r < MC; r++) {
    const T *src = A + (size_t)(ic + r) * lda + kk;
    const T *src2 = SA ? (A2 + (size_t)(ic + r) * lda + kk) : (const T *)0;
    T *dst = Apk + (size_t)r * APSTR;
    for (int m = 0; m < KC / 8; m++) {
      __builtin_prefetch(src + (size_t)lda + 8 * m, 0, 1);
      __builtin_prefetch(src + (size_t)2 * lda + 8 * m, 0, 1);
      __m256i x = LD(src + 8 * m);
      if (SA) { __builtin_prefetch(src2 + (size_t)lda + 8 * m, 0, 1); __m256i y = LD(src2 + 8 * m);
                x = (SA > 0) ? _mm256_add_epi32(x, y) : _mm256_sub_epi32(x, y); }
      __m256i al = _mm256_and_si256(x, m16);
      __m256i ah = _mm256_srli_epi32(_mm256_add_epi32(x, half), 16);
      __m256i pl = _mm256_packus_epi32(al, al);
      __m256i ph = _mm256_packus_epi32(ah, ah);
      T *d = dst + 64 * m;
      ST(d + 0, _mm256_permutevar8x32_epi32(pl, i0));
      ST(d + 8, _mm256_permutevar8x32_epi32(ph, i0));
      ST(d + 16, _mm256_permutevar8x32_epi32(pl, i1));
      ST(d + 24, _mm256_permutevar8x32_epi32(ph, i1));
      ST(d + 32, _mm256_permutevar8x32_epi32(pl, i4));
      ST(d + 40, _mm256_permutevar8x32_epi32(ph, i4));
      ST(d + 48, _mm256_permutevar8x32_epi32(pl, i5));
      ST(d + 56, _mm256_permutevar8x32_epi32(ph, i5));
    }
  }
}

template<int SB>
AVX2 static inline void packB2_t(int jc, int kk, int n, int ldb, const T *B, const T *B2) {
  const __m256i m16 = _mm256_set1_epi32(0xFFFF);
  const __m256i half = _mm256_set1_epi32(0x8000);
  (void)n;
  const size_t DSTR = (size_t)(NPAIR * 16 + 16);
  for (int m = 0; m < NPAIR; m++) {
    const T *s0 = B + (size_t)(kk + 2 * m) * ldb + jc;
    const T *s1 = s0 + ldb;
    const T *t0 = SB ? (B2 + (size_t)(kk + 2 * m) * ldb + jc) : (const T *)0;
    const T *t1 = t0 ? t0 + ldb : (const T *)0;
    T *d0 = Bpk + (size_t)m * 16;
    T *d1 = d0 + 8;
    for (int j = 0; j < NC; j += 8, d0 += DSTR, d1 += DSTR) {
      __m256i x0 = LD(s0 + j), x1 = LD(s1 + j);
      if (SB) {
        __m256i y0 = LD(t0 + j), y1 = LD(t1 + j);
        if (SB > 0) { x0 = _mm256_add_epi32(x0, y0); x1 = _mm256_add_epi32(x1, y1); }
        else        { x0 = _mm256_sub_epi32(x0, y0); x1 = _mm256_sub_epi32(x1, y1); }
      }
      __m256i l0 = _mm256_and_si256(x0, m16);
      __m256i l1 = _mm256_and_si256(x1, m16);
      __m256i h0 = _mm256_srli_epi32(_mm256_add_epi32(x0, half), 16);
      __m256i h1 = _mm256_srli_epi32(_mm256_add_epi32(x1, half), 16);
      ST(d0, _mm256_or_si256(l0, _mm256_slli_epi32(l1, 16)));
      ST(d1, _mm256_or_si256(h0, _mm256_slli_epi32(h1, 16)));
    }
  }
}

template<int SA, int SB>
AVX2 static void tiled2_128_t(const T *A1, const T *A2, int lda,
                              const T *B1, const T *B2, int ldb, T *C) {
  const int n = 128;
  const int ldc = 128;
  for (int jc = 0; jc < n; jc += NC) {
    for (int kk = 0; kk < n; kk += KC) {
      packB2_t<SB>(jc, kk, n, ldb, B1, B2);
      for (int ic = 0; ic < n; ic += MC) {
        packA2_t<SA>(ic, kk, n, lda, A1, A2);
        block_mikro(ic, jc, n, ldc, C, kk == 0);
      }
    }
  }
}

AVX2 static void rec(int n, const T *A, int lda, const T *B, int ldb, T *C, int ldc,
                     T *scr, int depth) {
  if (depth >= STR_DEPTH || (n & 1) || n <= BASE || !ok_tiled(n)) {
    if (ok_tiled(n)) tiled(n, A, lda, B, ldb, C, ldc);
    else             generic(n, A, lda, B, ldb, C, ldc);
    return;
  }
  const int h = n / 2;
  const size_t h2 = (size_t)h * h;
  T *S = scr;            /* 7 result blocks M1..M7 then 2 scratch S blocks */
  T *M1 = S, *M2 = M1 + h2 + 4608, *M3 = M2 + h2 + 4608, *M4 = M3 + h2 + 4608,
    *M5 = M4 + h2 + 4608, *M6 = M5 + h2 + 4608, *M7 = M6 + h2 + 4608;
  T *P = M7 + h2 + 4608;        /* operand scratch block 1 (h x h, ldp = h) */
  T *Q = P + h2 + 4608;         /* operand scratch block 2 */
  T *sub = Q + h2 + 4608;       /* recursion arena */

  /* Each of the 7 products uses an operand that is A1 (+/- A2) on the A side and
     B1 (+/- B2) on the B side.  When the children are LEAVES (h <= BASE) the summed
     operand is consumed only by the packers, so we fuse the sum INTO the packer and
     never materialise P/Q: that removes 2 of every 4 accesses per summed operand
     element (addblk's write + the packer's read of it, replaced by a second load of
     the original block, which is exactly what addblk did anyway). */
  if (h <= BASE) {
#define LOP(a1, a2, sgA, b1, b2, sgB, Mi)                                        \
  do {                                                                           \
    if (h <= BASE) {                                                             \
      tiled2_128((a1), (a2), (sgA), lda,                               \
                 (b1), (b2), (sgB), ldb, (Mi));                       \
    } else {                                                                     \
      const T *pa = (a1); int la = lda;                                          \
      const T *pb = (b1); int lb = ldb;                                          \
      if (a2) { addblk(h, P, (a1), lda, (a2), lda, (sgA)); pa = P; la = h; }     \
      if (b2) { addblk(h, Q, (b1), ldb, (b2), ldb, (sgB)); pb = Q; lb = h; }     \
      rec(h, pa, la, pb, lb, (Mi), h, sub, depth + 1);                           \
    }                                                                            \
  } while (0)

  do { if (h <= BASE) { tiled2_128_t<1,1>((B11), (B22), lda, (D11), (D22), ldb, (M1)); } else { const T *pa = (B11); int la = lda; const T *pb = (D11); int lb = ldb; if (1) { addblk(h, P, (B11), lda, (B22), lda, 1); pa = P; la = h; } if (1) { addblk(h, Q, (D11), ldb, (D22), ldb, 1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M1), h, sub, depth + 1); } } while (0);
  do { if (h <= BASE) { tiled2_128_t<1,0>((B21), (B22), lda, (D11), (0), ldb, (M2)); } else { const T *pa = (B21); int la = lda; const T *pb = (D11); int lb = ldb; if (1) { addblk(h, P, (B21), lda, (B22), lda, 1); pa = P; la = h; } if (0) { addblk(h, Q, (D11), ldb, (D11), ldb, 1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M2), h, sub, depth + 1); } } while (0);
  do { if (h <= BASE) { tiled2_128_t<0,-1>((B11), (0), lda, (D12), (D22), ldb, (M3)); } else { const T *pa = (B11); int la = lda; const T *pb = (D12); int lb = ldb; if (0) { addblk(h, P, (B11), lda, (B11), lda, 1); pa = P; la = h; } if (1) { addblk(h, Q, (D12), ldb, (D22), ldb, -1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M3), h, sub, depth + 1); } } while (0);
  do { if (h <= BASE) { tiled2_128_t<0,-1>((B22), (0), lda, (D21), (D11), ldb, (M4)); } else { const T *pa = (B22); int la = lda; const T *pb = (D21); int lb = ldb; if (0) { addblk(h, P, (B22), lda, (B22), lda, 1); pa = P; la = h; } if (1) { addblk(h, Q, (D21), ldb, (D11), ldb, -1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M4), h, sub, depth + 1); } } while (0);
  do { if (h <= BASE) { tiled2_128_t<1,0>((B11), (B12), lda, (D22), (0), ldb, (M5)); } else { const T *pa = (B11); int la = lda; const T *pb = (D22); int lb = ldb; if (1) { addblk(h, P, (B11), lda, (B12), lda, 1); pa = P; la = h; } if (0) { addblk(h, Q, (D22), ldb, (D22), ldb, 1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M5), h, sub, depth + 1); } } while (0);
  do { if (h <= BASE) { tiled2_128_t<-1,1>((B21), (B11), lda, (D11), (D12), ldb, (M6)); } else { const T *pa = (B21); int la = lda; const T *pb = (D11); int lb = ldb; if (1) { addblk(h, P, (B21), lda, (B11), lda, -1); pa = P; la = h; } if (1) { addblk(h, Q, (D11), ldb, (D12), ldb, 1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M6), h, sub, depth + 1); } } while (0);
  do { if (h <= BASE) { tiled2_128_t<-1,1>((B12), (B22), lda, (D21), (D22), ldb, (M7)); } else { const T *pa = (B12); int la = lda; const T *pb = (D21); int lb = ldb; if (1) { addblk(h, P, (B12), lda, (B22), lda, -1); pa = P; la = h; } if (1) { addblk(h, Q, (D21), ldb, (D22), ldb, 1); pb = Q; lb = h; } rec(h, pa, la, pb, lb, (M7), h, sub, depth + 1); } } while (0);
#undef LOP

  /* C11 = M1 + M4 - M5 + M7 ; C12 = M3 + M5 ; C21 = M2 + M4 ; C22 = M1 - M2 + M3 + M6 */
  comb_all(h, C, ldc, M1, M2, M3, M4, M5, M6, M7);
  } else {
    /* ---- STRASSEN-WINOGRAD: 8 materialised operand passes per node (classical: 10) ----
       The chains run IN PLACE and the child calls are ordered so that only one A-side
       and one B-side form are live at a time -> the SAME two scratch blocks P and Q
       suffice, and the arena layout is unchanged.  The leaf path (h <= BASE) above is
       the shipped classical bytes, untouched.                                          */
    addblk(h, P, (B21), lda, (B22), lda, 1);            /* S1 = A21 + A22       */
    addblk(h, Q, (D12), ldb, (D11), ldb, -1);           /* T1 = B12 - B11       */
    rec(h, P, h, Q, h, (M5), h, sub, depth + 1);        /* M5 = S1 * T1         */

    addblk(h, P, P, h, (B11), lda, -1);                 /* S2 = S1 - A11   (ip) */
    addblk(h, Q, (D22), ldb, Q, h, -1);                 /* T2 = B22 - T1   (ip) */
    rec(h, P, h, Q, h, (M6), h, sub, depth + 1);        /* M6 = S2 * T2         */

    addblk(h, Q, (D21), ldb, Q, h, -1);                 /* T4 = B21 - T2   (ip) */
    rec(h, (B22), lda, Q, h, (M4), h, sub, depth + 1);  /* M4 = A22 * T4        */

    addblk(h, P, (B12), lda, P, h, -1);                 /* S4 = A12 - S2   (ip) */
    rec(h, P, h, (D22), ldb, (M3), h, sub, depth + 1);  /* M3 = S4 * B22        */

    addblk(h, P, (B11), lda, (B21), lda, -1);           /* S3 = A11 - A21       */
    addblk(h, Q, (D22), ldb, (D12), ldb, -1);           /* T3 = B22 - B12       */
    rec(h, P, h, Q, h, (M7), h, sub, depth + 1);        /* M7 = S3 * T3         */

    rec(h, (B11), lda, (D11), ldb, (M1), h, sub, depth + 1);  /* M1 = A11 * B11 */
    rec(h, (B12), lda, (D21), ldb, (M2), h, sub, depth + 1);  /* M2 = A12 * B21 */

    comb_all_w(h, C, ldc, M1, M2, M3, M4, M5, M6, M7);
  }

}

AVX2 void matrix_multiply(int n, const T *A, const T *B, T *C) {
  unsigned long long al = (unsigned long long)(size_t)A | (unsigned long long)(size_t)B | (unsigned long long)(size_t)C;
  if (n <= BASE || !ok_tiled(n) || (al & 31)) {
    if (ok_tiled(n) && !(al & 31)) tiled(n, A, n, B, n, C, n);
    else             generic(n, A, n, B, n, C, n);
    return;
  }

  /* depth-first arena:  9 h^2 blocks at each level along the path */
  size_t need = 0, h = (size_t)n;
  while (h > (size_t)BASE) { h /= 2; need += 9 * (h * h + 4608); }
  static T *scr = 0; static size_t cap = 0;
  if (need > cap) { if (scr) free(scr); cap = need;
    void *pv = 0; if (posix_memalign(&pv, 4096, cap * sizeof(T))) return; scr = (T *)pv; }
  rec(n, A, n, B, n, C, n, scr, 0);
}

CompilationN/AN/ACompile OKScore: N/A

Testcase #11.778 s255 MB + 896 KBAcceptedScore: 100


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