/* m211 arm: A-panel ROW-STRIDE pad, KCPAD=8 (stride KC+8 = 2112 B = 33 lines) */
#define SPPAD 1024 /* lane c2k4_rv: non-leaf pool slot pad, in DOUBLES */
#define KCPAD 8
#ifndef FUSE67
#define FUSE67 1
#endif
#ifndef FUSE_NL
#define FUSE_NL 1
#endif
#if ILV
#endif
#ifndef ILV
#define ILV 1
#endif
#ifndef GC_NT
#define GC_NT 1
#endif
#ifndef MC_G
#define MC_G 6
#endif
#ifndef NC_G
#define NC_G 1024
#endif
#define KC_G 512
#ifndef SCUT_G
#ifndef SCUT_OVR
#define SCUT_G 512
#else
#define SCUT_G SCUT_OVR
#endif
#endif
// mmmd (double) : C = A*B -- Strassen-1 + packed AVX2/FMA 6x8 GEMM.
// Judge: g++-9 -static -O2 (no -march) => per-function target pragma.
#include <immintrin.h>
#include <stddef.h>
#include <stdlib.h>
#include <string.h>
typedef double T;
#ifndef MR
#define MR 6
#endif
#ifndef NR
#define NR 8
#endif
#ifndef KC
#define KC KC_G
#endif
#ifndef MC
#define MC MC_G
#endif
#ifndef NC
#define NC NC_G
#endif
#ifndef SCUT
#define SCUT SCUT_G
#endif
static unsigned long long g_nlm_nodes = 0;
static int g_top = 0;
#pragma GCC optimize("O3","unroll-loops")
#pragma GCC push_options
#pragma GCC target("avx2,fma")
// ★ Under the judge's plain `g++-9 -O2` (no -march) GCC defaults to
// -mavx256-split-unaligned-load/-store: EVERY _mm256_loadu_*/_mm256_storeu_*
// becomes TWO instructions (xmm load + vinserti128 / vextracti128 + xmm store),
// and the insert/extract is a port-5 uop. `#pragma GCC target("avx2")` does NOT
// turn that off. Inline asm bypasses the lowering: one instruction, no port 5.
// AT&T operand order is src,dst for BOTH forms -> "vmovupd %1, %0" in each.
#define AVX2T __attribute__((target("avx2,fma")))
AVX2T static inline __m256d LD(const double *p) {
__m256d r; __asm__("vmovupd %1, %0" : "=x"(r) : "m"(*(const __m256d *)p)); return r;
}
AVX2T static inline void ST(double *p, __m256d v) {
__asm__ volatile("vmovupd %1, %0" : "=m"(*(__m256d *)p) : "x"(v));
}
// Packed A: panel p holds rows [p*MR, p*MR+MR): dst[p*MR*KC + r*KC + k] = S(row r, k)
// Packed B: panel g holds cols [g*NR, g*NR+NR): dst[g*KC*NR + k*NR + j] = S(k, col j)
#ifndef KCPAD
#define KCPAD 0
#endif
#define KCP (KC + KCPAD) /* m211: A-panel ROW STRIDE. At KCPAD=0 the stride is KC = 2048 B,
exactly HALF the L1 set period (64 sets x 64 B = 4096 B), so all
six micro-rows land on two set-groups and the 8-way L1 is ~88 %
occupied before the C stores. KCPAD=8 -> 2112 B = 33 lines. */
/* =================== lane m4k_skel arm knobs =================== */
#define g_nt 1
#define g_pf 0
#define g_pfa 0
#define g_pfb 0
#define g_arm 0
static inline void stn(int nt, T *p, __m256d v){ if(nt) _mm256_stream_pd(p,v); else ST(p,v); }
static void blk_axpy2(int m,int n,T*dA,int ldA,T sA,int accA,T*dB,int ldB,T sB,int accB,const T*src,int lds){
const int nt = g_nt && !accA && !accB;
for(int i=0;i<m;i++){
T*a=dA+(size_t)i*ldA,*b=dB+(size_t)i*ldB; const T*x=src+(size_t)i*lds;
if(g_pf){__builtin_prefetch(x+g_pf,0,1);__builtin_prefetch(a+g_pf,0,1);__builtin_prefetch(b+g_pf,0,1);}
int j=0;
for(;j+4<=n;j+=4){
__m256d v=LD(x+j);
if(accA){ if(sA>0) ST(a+j,_mm256_add_pd(LD(a+j),v)); else ST(a+j,_mm256_sub_pd(LD(a+j),v)); }
else { if(sA>0) stn(nt,a+j,v); else stn(nt,a+j,_mm256_sub_pd(_mm256_setzero_pd(),v)); }
if(accB){ if(sB>0) ST(b+j,_mm256_add_pd(LD(b+j),v)); else ST(b+j,_mm256_sub_pd(LD(b+j),v)); }
else { if(sB>0) stn(nt,b+j,v); else stn(nt,b+j,_mm256_sub_pd(_mm256_setzero_pd(),v)); }
}
for(;j<n;j++){ T v=x[j];
a[j]=accA?(sA>0?a[j]+v:a[j]-v):(sA>0?v:-v);
b[j]=accB?(sB>0?b[j]+v:b[j]-v):(sB>0?v:-v); }
}
}
static void blk_axpy(int m,int n,T*dst,int ldd,const T*src,int lds,T s,int accumulate){
const int nt = g_nt && !accumulate;
for(int i=0;i<m;i++){
T*d=dst+(size_t)i*ldd; const T*x=src+(size_t)i*lds;
if(g_pf){__builtin_prefetch(x+g_pf,0,1);}
int j=0;
if(accumulate){
if(s>0){for(;j+4<=n;j+=4)ST(d+j,_mm256_add_pd(LD(d+j),LD(x+j)));for(;j<n;j++)d[j]+=x[j];}
else {for(;j+4<=n;j+=4)ST(d+j,_mm256_sub_pd(LD(d+j),LD(x+j)));for(;j<n;j++)d[j]-=x[j];}
} else {
if(s>0){for(;j+4<=n;j+=4)stn(nt,d+j,LD(x+j));for(;j<n;j++)d[j]=x[j];}
else {for(;j+4<=n;j+=4)stn(nt,d+j,_mm256_sub_pd(_mm256_setzero_pd(),LD(x+j)));for(;j<n;j++)d[j]=-x[j];}
}
}
}
static void sadd(int m,int n,T*dst,int ldd,const T*s1,const T*s2,int sign,int ld1,int ld2){
for(int i=0;i<m;i++){
T*d=dst+(size_t)i*ldd; const T*p=s1+(size_t)i*ld1;
const T*q=s2?s2+(size_t)i*ld2:(const T*)0;
if(g_pf){__builtin_prefetch(p+g_pf,0,1); if(q)__builtin_prefetch(q+g_pf,0,1);}
int j=0;
if(!q){for(;j+4<=n;j+=4)stn(g_nt,d+j,LD(p+j));for(;j<n;j++)d[j]=p[j];continue;}
if(sign>0){for(;j+4<=n;j+=4)stn(g_nt,d+j,_mm256_add_pd(LD(p+j),LD(q+j)));for(;j<n;j++)d[j]=p[j]+q[j];}
else {for(;j+4<=n;j+=4)stn(g_nt,d+j,_mm256_sub_pd(LD(p+j),LD(q+j)));for(;j<n;j++)d[j]=p[j]-q[j];}
}
}
static void comb5(int m,T*C11,T*C12,T*C21,T*C22,int ldc,
const T*m1,const T*m2,const T*m3,const T*m4,const T*m5,int ldm){
const int nt=g_nt;
for(int i=0;i<m;i++){
T*a=C11+(size_t)i*ldc,*b=C12+(size_t)i*ldc,*c=C21+(size_t)i*ldc,*d=C22+(size_t)i*ldc;
const T*p1=m1+(size_t)i*ldm,*p2=m2+(size_t)i*ldm,*p3=m3+(size_t)i*ldm,*p4=m4+(size_t)i*ldm,*p5=m5+(size_t)i*ldm;
if(g_pf){__builtin_prefetch(p1+g_pf,0,1);__builtin_prefetch(p2+g_pf,0,1);__builtin_prefetch(p3+g_pf,0,1);__builtin_prefetch(p4+g_pf,0,1);__builtin_prefetch(p5+g_pf,0,1);}
for(int j=0;j+4<=m;j+=4){
__m256d x1=LD(p1+j),x2=LD(p2+j),x3=LD(p3+j),x4=LD(p4+j),x5=LD(p5+j);
stn(nt,a+j,_mm256_add_pd(_mm256_sub_pd(x1,x5),x4));
stn(nt,b+j,_mm256_add_pd(x3,x5));
stn(nt,c+j,_mm256_add_pd(x2,x4));
stn(nt,d+j,_mm256_sub_pd(_mm256_add_pd(x1,x3),x2));
}
for(int j=(m&~3);j<m;j++){
double x1=p1[j],x2=p2[j],x3=p3[j],x4=p4[j],x5=p5[j];
a[j]=x1-x5+x4; b[j]=x3+x5; c[j]=x2+x4; d[j]=x1-x2+x3;
}
}
}
/* =============================================================== */
static T Apack[(size_t)((MC + MR - 1) / MR) * MR * KCP + 64] __attribute__((aligned(64)));
static T Bpack[(size_t)((NC + NR - 1) / NR) * NR * KC + 64] __attribute__((aligned(64)));
// 8-wide op: dst[0..7] = s1 or s1+s2 or s1-s2
static inline void vop8(T *d, const T *s1, const T *s2, int op) {
if (op == 0) {
ST(d, LD(s1));
ST(d + 4, LD(s1 + 4));
} else if (op == 1) {
ST(d, _mm256_add_pd(LD(s1), LD(s2)));
ST(d + 4, _mm256_add_pd(LD(s1 + 4), LD(s2 + 4)));
} else {
ST(d, _mm256_sub_pd(LD(s1), LD(s2)));
ST(d + 4, _mm256_sub_pd(LD(s1 + 4), LD(s2 + 4)));
}
}
static void packA(int mc, int kc, const T *a1, const T *a2, int op, int lda, T *Ap) {
for (int p = 0; p < mc; p += MR) {
int mr = mc - p < MR ? mc - p : MR;
T *dst = Ap + (size_t)(p / MR) * MR * KCP;
for (int r = 0; r < mr; r++) {
const T *s1 = a1 + (size_t)(p + r) * lda;
if (r + 1 < mr) __builtin_prefetch(s1 + lda, 0, 1); /* E8PF */
else if (p + MR + r < mc) __builtin_prefetch(s1 + (size_t)MR * lda, 0, 1);
for (int _q = 1; _q <= g_pfa; _q++) {
int _r = r + _q;
if (_r < mr) __builtin_prefetch(s1 + (size_t)_q * lda, 0, 1);
else if (p + MR + _r < mc) __builtin_prefetch(s1 + (size_t)(MR + _q) * lda, 0, 1);
}
const T *s2 = a2 ? a2 + (size_t)(p + r) * lda : (const T *)0;
T *d = dst + (size_t)r * KCP;
int k = 0;
if (op == 0) { for (; k + 8 <= kc; k += 8) vop8(d + k, s1 + k, 0, 0); }
else if (op == 1) { for (; k + 8 <= kc; k += 8) vop8(d + k, s1 + k, s2 + k, 1); }
else { for (; k + 8 <= kc; k += 8) vop8(d + k, s1 + k, s2 + k, 2); }
for (; k < kc; k++) d[k] = op == 0 ? s1[k] : (op == 1 ? s1[k] + s2[k] : s1[k] - s2[k]);
}
}
}
static void packB(int kc, int nc, const T *b1, const T *b2, int op, int lda, T *Bp) {
/* E8SWAP (ported from mmmd2k, where it is board-priced at -2.21 %): the shipped loop is
g OUTER / k INNER, so it reads 8 doubles (ONE FULL 64-B LINE) per lda*8 = 32 KB jump and
kc of them in a row. THE HW PREFETCHER CANNOT FOLLOW A 32 KB STRIDE, so every line was a
DEMAND MISS. k OUTER / g INNER makes the source read nc*8 bytes CONTIGUOUS per row. */
for (int k0 = 0; k0 < kc; k0 += 8)
for (int g = 0; g + NR <= nc; g += NR) {
if (g_pfb) __builtin_prefetch(b1 + (size_t)(k0 + 8*g_pfb) * lda + g, 0, 1);
for (int k = k0; k < k0 + 8 && k < kc; k++)
vop8(Bp + (size_t)(g / NR) * KC * NR + (size_t)k * NR,
b1 + (size_t)k * lda + g, b2 ? b2 + (size_t)k * lda + g : (const T *)0, op);
}
for (int g = nc - (nc % NR); g < nc; g++) {
T *d = Bp + (size_t)(g / NR) * KC * NR;
for (int k = 0; k < kc; k++) {
T *dd = d + (size_t)k * NR;
for (int j = 0; j < NR; j++) dd[j] = (g + j < nc) ? (op == 0 ? b1[(size_t)k*lda+g+j] : (op == 1 ? b1[(size_t)k*lda+g+j]+b2[(size_t)k*lda+g+j] : b1[(size_t)k*lda+g+j]-b2[(size_t)k*lda+g+j])) : 0.0;
}
}
}
#define KSTEP(U) \
do { \
__m256d b0 = LD(bp + (size_t)(k + (U)) * NR); \
__m256d b1 = LD(bp + (size_t)(k + (U)) * NR + 4); \
__m256d t; \
t = _mm256_broadcast_sd(ap + 0 * KCP + (k + (U))); \
c00 = _mm256_fmadd_pd(t, b0, c00); c01 = _mm256_fmadd_pd(t, b1, c01); \
t = _mm256_broadcast_sd(ap + 1 * KCP + (k + (U))); \
c10 = _mm256_fmadd_pd(t, b0, c10); c11 = _mm256_fmadd_pd(t, b1, c11); \
t = _mm256_broadcast_sd(ap + 2 * KCP + (k + (U))); \
c20 = _mm256_fmadd_pd(t, b0, c20); c21 = _mm256_fmadd_pd(t, b1, c21); \
t = _mm256_broadcast_sd(ap + 3 * KCP + (k + (U))); \
c30 = _mm256_fmadd_pd(t, b0, c30); c31 = _mm256_fmadd_pd(t, b1, c31); \
t = _mm256_broadcast_sd(ap + 4 * KCP + (k + (U))); \
c40 = _mm256_fmadd_pd(t, b0, c40); c41 = _mm256_fmadd_pd(t, b1, c41); \
t = _mm256_broadcast_sd(ap + 5 * KCP + (k + (U))); \
c50 = _mm256_fmadd_pd(t, b0, c50); c51 = _mm256_fmadd_pd(t, b1, c51); \
} while (0)
// ncols guard: B panel is zero-padded past ncols, so only the store needs masking.
__attribute__((always_inline)) static inline void micro(int kc, const T *__restrict ap, const T *__restrict bp,
T *__restrict C, int ldc, int accumulate, int mrows, int ncols, int nt) {
__m256d c00, c01, c10, c11, c20, c21, c30, c31, c40, c41, c50, c51;
c00 = c01 = c10 = c11 = c20 = c21 = c30 = c31 = c40 = c41 = c50 = c51 = _mm256_setzero_pd();
/* lane m164c SPF: the C destination is written in NR-double row segments ldc*8 bytes apart
(16 KB at ldc=2048) and an ACCUMULATING store is a read-modify-write, so its first touch is
a cold line fetch paid at the very end of the call, where the OoO window is closing.
One prefetch per row, issued before the k-loop, is hidden by ~kc*6.5 cycles of mikro work.
Output-preserving by construction (a prefetch cannot change a value). */
if (accumulate) {
for (int r = 0; r < mrows; r++)
_mm_prefetch((const char *)(C + (size_t)r * ldc), _MM_HINT_T0);
}
int k = 0;
for (; k + 4 <= kc; k += 4) {
__builtin_prefetch(bp + (size_t)(k + 16) * NR);
__builtin_prefetch(bp + (size_t)(k + 17) * NR);
__builtin_prefetch(bp + (size_t)(k + 18) * NR);
__builtin_prefetch(bp + (size_t)(k + 19) * NR);
KSTEP(0); KSTEP(1); KSTEP(2); KSTEP(3);
}
for (; k < kc; k++) KSTEP(0);
if (ncols == NR) {
#define STO(r, a, b) \
if (r < mrows) { \
T *p = C + (size_t)(r) * ldc; \
if (accumulate) { \
ST(p, _mm256_add_pd(LD(p), a)); \
ST(p + 4, _mm256_add_pd(LD(p + 4), b)); \
} else if (nt) { \
_mm256_stream_pd(p, a); _mm256_stream_pd(p + 4, b); \
} else { ST(p, a); ST(p + 4, b); } \
}
STO(0, c00, c01) STO(1, c10, c11) STO(2, c20, c21)
STO(3, c30, c31) STO(4, c40, c41) STO(5, c50, c51)
#undef STO
} else {
__m256i ml = _mm256_set_epi64x(ncols > 3 ? -1LL : 0, ncols > 2 ? -1LL : 0,
ncols > 1 ? -1LL : 0, ncols > 0 ? -1LL : 0);
__m256i mh = _mm256_set_epi64x(ncols > 7 ? -1LL : 0, ncols > 6 ? -1LL : 0,
ncols > 5 ? -1LL : 0, ncols > 4 ? -1LL : 0);
#define STO2(r, a, b) \
if (r < mrows) { \
T *p = C + (size_t)(r) * ldc; \
if (accumulate) { \
_mm256_maskstore_pd(p, ml, _mm256_add_pd(_mm256_maskload_pd(p, ml), a)); \
_mm256_maskstore_pd(p + 4, mh, _mm256_add_pd(_mm256_maskload_pd(p + 4, mh), b)); \
} else { _mm256_maskstore_pd(p, ml, a); _mm256_maskstore_pd(p + 4, mh, b); } \
}
STO2(0, c00, c01) STO2(1, c10, c11) STO2(2, c20, c21)
STO2(3, c30, c31) STO2(4, c40, c41) STO2(5, c50, c51)
#undef STO2
}
}
/* E8FUSE (ported from mmmd2k, board-priced at -0.6 % there): the two combine axpys of one
Strassen product each read the WHOLE m x m scratch M; one pass reads M once, not twice. */
/* moved */
static void gemm_core(int m, int n, int k,
const T *a1, const T *a2, int oa, int lda,
const T *b1, const T *b2, int ob, int ldb,
T *C, int ldc, int accumulate, int nt_allowed) {
if (m <= 0 || n <= 0 || k <= 0) return;
/* non-temporal C stores need every written 64 B line to be complete and aligned */
int ntok = nt_allowed && ((ldc & 7) == 0) && (((size_t)C & 63) == 0);
for (int jc = 0; jc < n; jc += NC) {
int nc = n - jc < NC ? n - jc : NC;
for (int pc = 0; pc < k; pc += KC) {
int kc = k - pc < KC ? k - pc : KC;
packB(kc, nc, b1 + (size_t)pc * ldb + jc,
b2 ? b2 + (size_t)pc * ldb + jc : (const T *)0, ob, ldb, Bpack);
int acc0 = accumulate || pc > 0;
for (int ic = 0; ic < m; ic += MC) {
int mc = m - ic < MC ? m - ic : MC;
packA(mc, kc, a1 + (size_t)ic * lda + pc,
a2 ? a2 + (size_t)ic * lda + pc : (const T *)0, oa, lda, Apack);
T *Cp = C + (size_t)ic * ldc + jc;
int acc = acc0; /* (pc,ic) blocks cover DISJOINT rows: only pc>0 accumulates */
for (int jr = 0; jr < nc; jr += NR) {
int ncc = nc - jr < NR ? nc - jr : NR;
for (int ir = 0; ir + MR <= mc; ir += MR) {
const T *ap = Apack + (size_t)(ir / MR) * MR * KCP;
const T *bp = Bpack + (size_t)(jr / NR) * KC * NR;
micro(kc, ap, bp, Cp + (size_t)ir * ldc + jr, ldc, acc, MR, ncc, ntok && !acc);
}
int ir = mc - (mc % MR);
if (ir < mc) { /* row tail: phantom rows must not be stored */
const T *ap = Apack + (size_t)(ir / MR) * MR * KCP;
const T *bp = Bpack + (size_t)(jr / NR) * KC * NR;
micro(kc, ap, bp, Cp + (size_t)ir * ldc + jr, ldc, acc, mc - ir, ncc, 0);
}
}
}
}
}
}
/* moved */
/* ============ LEAF-FUSED RECURSIVE STRASSEN (spliced driver) ============
* Same 7-product Strassen decomposition, but the LAST level does not
* materialise its operand sums: gemm_core already takes (s2, op) and folds the
* elementwise add/sub into its packing loop. That deletes the whole
* bottom-level sadd pass (2 reads + 1 write of m*m per operand) and the
* re-read of the materialised sum.
*/
#define NIL ((const T *)0)
#if GC_NT
#define GC(a1,a2,oa,lda,b1,b2,ob,ldb,c,ldc,acc) gemm_core(m,m,m,a1,a2,oa,lda,b1,b2,ob,ldb,c,ldc,acc,1)
#define GC0(sz,a1,a2,oa,lda,b1,b2,ob,ldb,c,ldc,acc) gemm_core(sz,sz,sz,a1,a2,oa,lda,b1,b2,ob,ldb,c,ldc,acc,0)
#else
#define GC(a1,a2,oa,lda,b1,b2,ob,ldb,c,ldc,acc) gemm_core(m,m,m,a1,a2,oa,lda,b1,b2,ob,ldb,c,ldc,acc)
#define GC0(sz,a1,a2,oa,lda,b1,b2,ob,ldb,c,ldc,acc) gemm_core(sz,sz,sz,a1,a2,oa,lda,b1,b2,ob,ldb,c,ldc,acc)
#endif
/* moved */
/* ---- DA1: fused single-pass Strassen combine (lane_da1_mmmd24k) ----
The shipped leaf combine is 5 sequential blk_axpy2 passes: each C quadrant is read+written
2-4 times, 9 quadrant touches where 4 would do. Holding M1..M5 LIVE lets one pass produce all
four quadrants. MEASURED on the real judge, whole row, interleaved base/fuse/fuse/base at
n=2048: 1 071 333 406 / 1 071 273 606 -> 1 062 728 456 / 1 062 600 628 ticks = -0.807 %.
CORRECTNESS: comb5 OVERWRITES all four quadrants and FUSE67's two direct products are applied
AFTERWARDS. That ordering is forced: at a CHILD level the destination `C` is the parent's
recycled `M` plane, so ACCUMULATING into C11/C22 would pick up the previous product -- which is
exactly how this arm returned a WRONG ANSWER (max|diff| 2.36e+02) before the fix.
Max |diff| vs the shipped engine at n=2048 = 2.84e-14 against the problem's 3n^2*eps
= 2.79e-09 => ratio 1.0e-5. */
/* moved */
/* ==================== lane m4k_nl: NLM helpers ==================== */
/* ONE pass over the four operand quadrants produces the five Strassen sums.
base: 5 sadd calls x (2 reads + 1 write) = 15 m^2 of traffic per side.
NLM : 4 reads + 5 writes = 9 m^2 per side. */
static void sum5A(int m,const T*p11,const T*p12,const T*p21,const T*p22,int ld,
T*d1,T*d2,T*d5,T*d6,T*d7,int ldo,int lq,int nt){
nt = nt && ((((size_t)d1|(size_t)d2|(size_t)d5|(size_t)d6|(size_t)d7) & 31)==0);
for(int i=0;i<m;i++){
const T*a11=p11+(size_t)i*ld,*a12=p12+(size_t)i*ld,*a21=p21+(size_t)i*ld,*a22=p22+(size_t)i*ld;
T*r1=d1+(size_t)i*ldo,*r2=d2+(size_t)i*ldo,*r5=d5+(size_t)i*ldo,*r6=d6+(size_t)i*lq,*r7=d7+(size_t)i*lq;
int j=0;
for(;j+4<=m;j+=4){
__m256d x11=LD(a11+j),x12=LD(a12+j),x21=LD(a21+j),x22=LD(a22+j);
stn(nt,r1+j,_mm256_add_pd(x11,x22));
stn(nt,r2+j,_mm256_add_pd(x21,x22));
stn(nt,r5+j,_mm256_add_pd(x11,x12));
stn(nt,r6+j,_mm256_sub_pd(x21,x11));
stn(nt,r7+j,_mm256_sub_pd(x12,x22));
}
for(;j<m;j++){ T x11=a11[j],x12=a12[j],x21=a21[j],x22=a22[j];
r1[j]=x11+x22; r2[j]=x21+x22; r5[j]=x11+x12; r6[j]=x21-x11; r7[j]=x12-x22; }
}
}
static void sum5B(int m,const T*p11,const T*p12,const T*p21,const T*p22,int ld,
T*d1,T*d3,T*d4,T*d6,T*d7,int ldo,int lq,int nt){
nt = nt && ((((size_t)d1|(size_t)d3|(size_t)d4|(size_t)d6|(size_t)d7) & 31)==0);
for(int i=0;i<m;i++){
const T*b11=p11+(size_t)i*ld,*b12=p12+(size_t)i*ld,*b21=p21+(size_t)i*ld,*b22=p22+(size_t)i*ld;
T*r1=d1+(size_t)i*ldo,*r3=d3+(size_t)i*ldo,*r4=d4+(size_t)i*ldo,*r6=d6+(size_t)i*lq,*r7=d7+(size_t)i*lq;
int j=0;
for(;j+4<=m;j+=4){
__m256d x11=LD(b11+j),x12=LD(b12+j),x21=LD(b21+j),x22=LD(b22+j);
stn(nt,r1+j,_mm256_add_pd(x11,x22));
stn(nt,r3+j,_mm256_sub_pd(x12,x22));
stn(nt,r4+j,_mm256_sub_pd(x21,x11));
stn(nt,r6+j,_mm256_add_pd(x11,x12));
stn(nt,r7+j,_mm256_add_pd(x21,x22));
}
for(;j<m;j++){ T x11=b11[j],x12=b12[j],x21=b21[j],x22=b22[j];
r1[j]=x11+x22; r3[j]=x12-x22; r4[j]=x21-x11; r6[j]=x11+x12; r7[j]=x21+x22; }
}
}
/* SEVEN live M planes, ONE pass, all four quadrants out (base: comb5 = 9 m^2 and TWO
materialised FUSE67 products added by blk_axpy = 6 m^2 => 15 m^2). */
static void comb7(int m,T*C11,T*C12,T*C21,T*C22,int ldc,
const T*m1,const T*m2,const T*m3,const T*m4,const T*m5,const T*m6,const T*m7,int ldm){
const int nt=g_nt;
for(int i=0;i<m;i++){
T*a=C11+(size_t)i*ldc,*b=C12+(size_t)i*ldc,*c=C21+(size_t)i*ldc,*d=C22+(size_t)i*ldc;
const T*q1=m1+(size_t)i*ldm,*q2=m2+(size_t)i*ldm,*q3=m3+(size_t)i*ldm,*q4=m4+(size_t)i*ldm,
*q5=m5+(size_t)i*ldm,*q6=m6+(size_t)i*ldm,*q7=m7+(size_t)i*ldm;
int j=0;
for(;j+4<=m;j+=4){
__m256d x1=LD(q1+j),x2=LD(q2+j),x3=LD(q3+j),x4=LD(q4+j),x5=LD(q5+j),x6=LD(q6+j),x7=LD(q7+j);
stn(nt,a+j,_mm256_add_pd(_mm256_add_pd(_mm256_sub_pd(x1,x5),x4),x7));
stn(nt,b+j,_mm256_add_pd(x3,x5));
stn(nt,c+j,_mm256_add_pd(x2,x4));
stn(nt,d+j,_mm256_add_pd(_mm256_add_pd(_mm256_sub_pd(x1,x2),x3),x6));
}
for(;j<m;j++){
double x1=q1[j],x2=q2[j],x3=q3[j],x4=q4[j],x5=q5[j],x6=q6[j],x7=q7[j];
a[j]=x1-x5+x4+x7; b[j]=x3+x5; c[j]=x2+x4; d[j]=x1-x2+x3+x6;
}
}
}
static void dsgemm(int n, const T *A, const T *B, T *C, int lda, int ldb, int ldc, T *scratch) {
if (n <= SCUT) { GC0(n, A, 0, 0, lda, B, 0, 0, ldb, C, ldc, 0); return; }
int m = n >> 1;
int leaf = (m <= SCUT);
const T *A11 = A, *A12 = A + m, *A21 = A + (size_t)m * lda, *A22 = A + (size_t)m * lda + m;
const T *B11 = B, *B12 = B + m, *B21 = B + (size_t)m * ldb, *B22 = B + (size_t)m * ldb + m;
T *C11 = C, *C12 = C + m, *C21 = C + (size_t)m * ldc, *C22 = C + (size_t)m * ldc + m;
T *M = scratch;
#if FUSE_NL
T *TA = leaf ? (T *)0 : scratch + 5 * (size_t)m * m;
T *TB = leaf ? (T *)0 : scratch + 6 * (size_t)m * m;
T *sub = scratch + (leaf ? (size_t)m * m : 7 * (size_t)m * m);
#else
T *TA = leaf ? (T *)0 : scratch + (size_t)m * m;
T *TB = leaf ? (T *)0 : scratch + 2 * (size_t)m * m;
T *sub = scratch + (leaf ? (size_t)m * m : 3 * (size_t)m * m);
#endif
/* one product: A-operand = Aa <op> Ab (op 0=copy,1=+,2=-), B likewise, into Dst */
#define PROD(oa, Aa, Ab, ob, Ba, Bb, Dst, lddst, acc) \
do { \
if (leaf) { \
GC(Aa, Ab, oa, lda, Ba, Bb, ob, ldb, Dst, lddst, acc); \
} else { \
const T *_a = Aa, *_b = Ba; int _la = lda, _lb = ldb; \
if (oa) { sadd(m, m, TA, m, Aa, Ab, oa, lda, lda); _a = TA; _la = m; } \
if (ob) { sadd(m, m, TB, m, Ba, Bb, ob, ldb, ldb); _b = TB; _lb = m; } \
dsgemm(m, _a, _b, (Dst), _la, _lb, (lddst), sub); \
} \
} while (0)
if (leaf) {
/* DA1: 5 live M planes, ONE fused combine pass (all four quadrants produced together). */
#if ILV
/* m211 ILV: the five live M planes INTERLEAVED by row (plane i at base+i*m, row stride
5*m). comb5's five read streams become ONE contiguous 5*m-double stream per output
row -- one prefetch stream, one page set -- while the mikro's per-row NT stores only
change their stride. BIT-IDENTICAL (same values, same order, addresses only). */
const int ld5 = 5 * m;
T *M1 = scratch, *M2 = M1 + m, *M3 = M2 + m, *M4 = M3 + m, *M5 = M4 + m;
GC(A11, A12, 1, lda, B22, NIL, 0, ldb, M5, ld5, 0);
GC(A11, NIL, 0, lda, B12, B22, 2, ldb, M3, ld5, 0);
GC(A11, A22, 1, lda, B11, B22, 1, ldb, M1, ld5, 0);
GC(A21, A22, 1, lda, B11, NIL, 0, ldb, M2, ld5, 0);
GC(A22, NIL, 0, lda, B21, B11, 2, ldb, M4, ld5, 0);
comb5(m, C11, C12, C21, C22, ldc, M1, M2, M3, M4, M5, ld5);
#else
T *M1 = scratch, *M2 = M1 + (size_t)m * m, *M3 = M2 + (size_t)m * m,
*M4 = M3 + (size_t)m * m, *M5 = M4 + (size_t)m * m;
GC(A11, A22, 1, lda, B11, B22, 1, ldb, M1, m, 0);
GC(A21, A22, 1, lda, B11, NIL, 0, ldb, M2, m, 0);
GC(A11, NIL, 0, lda, B12, B22, 2, ldb, M3, m, 0);
GC(A22, NIL, 0, lda, B21, B11, 2, ldb, M4, m, 0);
GC(A11, A12, 1, lda, B22, NIL, 0, ldb, M5, m, 0);
comb5(m, C11, C12, C21, C22, ldc, M1, M2, M3, M4, M5, m);
#endif
PROD(-1, A21, A11, +1, B11, B12, C22, ldc, 1);
PROD(-1, A12, A22, +1, B21, B22, C11, ldc, 1);
return;
}
{
/* ---------- lane m4k_nl MS5: THE C-DEAD WINDOW AT EVERY NON-LEAF NODE ----------
A node's own C is dead from its entry until its own combine (z4 SS1), and with a
7-input combine there are NO post-combine products, so nothing pool-resident needs
that window. C is n^2 = 4*m^2 doubles = EXACTLY four m^2 slots. The four slots
used are s3/s4/s8/s9 -- the sums consumed by M6/M7, whose OUTPUTS land in pool
slots -- so comb7 never reads a C-resident plane while it writes C.
=> 7 slots from the pool per non-leaf level, i.e. THE BASE'S OWN pool size. */
int istop = g_top; g_top = 0;
if (istop) {
/* ---------- lane m4k_rst NLT: the top node's ORIGINAL pre-FUSE_NL branch ---------- */
const int ld5 = 5 * m;
T *F1 = scratch, *F2 = F1 + m, *F3 = F2 + m, *F4 = F3 + m, *F5 = F4 + m;
PROD(+1, A11, A22, +1, B11, B22, F1, ld5, 0);
PROD(+1, A21, A22, 0, B11, NIL, F2, ld5, 0);
PROD( 0, A11, NIL, -1, B12, B22, F3, ld5, 0);
PROD( 0, A22, NIL, -1, B21, B11, F4, ld5, 0);
PROD(+1, A11, A12, 0, B22, NIL, F5, ld5, 0);
comb5(m, C11, C12, C21, C22, ldc, F1, F2, F3, F4, F5, ld5);
PROD(-1, A21, A11, +1, B11, B12, F1, m, 0); blk_axpy(m, m, C22, ldc, F1, m, +1, 1);
PROD(-1, A12, A22, +1, B21, B22, F2, m, 0); blk_axpy(m, m, C11, ldc, F2, m, +1, 1);
return;
}
g_nlm_nodes++;
const int ldo = m;
const size_t SP = (size_t)m*m + (size_t)SPPAD;
T *s3, *s4, *s8, *s9, *subn;
T *s0 = scratch + 0*SP, *s1 = scratch + 1*SP, *s2 = scratch + 2*SP;
T *s5 = scratch + 3*SP, *s6 = scratch + 4*SP, *s7 = scratch + 5*SP, *sX = scratch + 6*SP;
/* m4k_rst: the four C slots ARE the four quadrant buffers, at the C row stride ldc.
They exactly tile the node's own C region for ANY parent plane layout, and they
are the addresses comb7 writes -- so s3/s4/s8/s9 are dead exactly when it overwrites
them. Identical to the shipped placement when ldc == 2*m (an MS5-form parent). */
const int lq = ldc;
s3 = C; s4 = C + (size_t)m; s8 = C + (size_t)m*ldc; s9 = C + (size_t)m*ldc + m;
subn = scratch + 7*SP;
const int nts = g_nt;
sum5A(m, A11, A12, A21, A22, lda, s0, s1, s2, s3, s4, ldo, lq, nts);
sum5B(m, B11, B12, B21, B22, ldb, s5, s6, s7, s8, s9, ldo, lq, nts);
dsgemm(m, s0, s5, sX, ldo, ldo, m, subn); /* M1 = (A11+A22)(B11+B22) -> sX */
dsgemm(m, s1, B11, s0, ldo, ldb, m, subn); /* M2 = (A21+A22) B11 -> s0 */
dsgemm(m, A11, s6, s5, lda, ldo, m, subn); /* M3 = A11 (B12-B22) -> s5 */
dsgemm(m, A22, s7, s1, lda, ldo, m, subn); /* M4 = A22 (B21-B11) -> s1 */
dsgemm(m, s2, B22, s6, ldo, ldb, m, subn); /* M5 = (A11+A12) B22 -> s6 */
dsgemm(m, s3, s8, s7, lq, lq, m, subn); /* M6 = (A21-A11)(B11+B12) -> s7 */
dsgemm(m, s4, s9, s2, lq, lq, m, subn); /* M7 = (A12-A22)(B21+B22) -> s2 */
comb7(m, C11, C12, C21, C22, ldc, sX, s0, s5, s1, s6, s7, s2, ldo);
return;
}
#undef PROD
}
static T *g_pool = 0;
static size_t g_pool_n = 0;
/* m211: the pool must hold 5 live M planes + TA + TB at EVERY non-leaf level
(FUSE_NL), or 5 at a leaf level; the levels below share the tail sequentially. */
static size_t pool_need(int n) {
size_t need = 0;
int m = n >> 1;
for (;;) {
if (m <= SCUT) { need += 5 * (size_t)m * m; break; }
need += 7 * ((size_t)m * m + SPPAD);
m >>= 1;
}
return need + 1024;
}
static void strassen1(int n, const T *A, const T *B, T *C, int ldc) {
#if FUSE_NL
size_t need = pool_need(n);
#else
size_t need = 2 * (size_t)n * n + 1024; /* DA1: 5 live m^2 planes at a leaf level */
#endif
if (g_pool_n < need) {
free(g_pool);
if (posix_memalign((void **)&g_pool, 64, need * sizeof(T))) g_pool = 0;
g_pool_n = g_pool ? need : 0;
}
if (!g_pool) { GC0(n, A, 0, 0, n, B, 0, 0, n, C, n, 0); return; }
g_top = 1;
dsgemm(n, A, B, C, n, n, ldc, g_pool); g_top = 0;
}
#pragma GCC pop_options
#if defined(LOCAL_DRIVER) || defined(LOCAL_ONLY)
#else
void matrix_multiply(int n, const double *A, const double *B, double *C) {
if (n > SCUT) strassen1(n, A, B, C, n);
else gemm_core(n, n, n, A, 0, 0, n, B, 0, 0, n, C, n, 0, 1);
}
#endif
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 2.003 s | 420 MB + 72 KB | Accepted | Score: 100 | 显示更多 |