// Classical float matrix multiplication: 4x24 register tile, MC120 NC48 KC256.
// All real n^3 scalar products are computed; padded columns are zero only.
#pragma GCC optimize("O3")
#pragma GCC target("avx2,fma")
#include <immintrin.h>
#include <stddef.h>
#include <stdlib.h>
#include <string.h>
namespace tile4x24 {
constexpr int MR=4,NR=24,MC=96,NC=48,KC=256,MAXN=8192;
alignas(4096) static float ap[MC*MAXN],ct[MC*NC];
__attribute__((always_inline)) static inline void kernel(const float*a,const float*b,float*c,int stride){
const float*aa=a;const float*bb=b;size_t ld=(size_t)stride*sizeof(float);
asm volatile(
"leaq (%%rdi,%%rdi,2), %%rdx\n\t"
"addq %%rcx, %%rdx\n\t"
"vxorps %%ymm4, %%ymm4, %%ymm4\n\t"
"vxorps %%ymm5, %%ymm5, %%ymm5\n\t"
"vxorps %%ymm6, %%ymm6, %%ymm6\n\t"
"vxorps %%ymm7, %%ymm7, %%ymm7\n\t"
"vxorps %%ymm8, %%ymm8, %%ymm8\n\t"
"vxorps %%ymm9, %%ymm9, %%ymm9\n\t"
"vxorps %%ymm10, %%ymm10, %%ymm10\n\t"
"vxorps %%ymm11, %%ymm11, %%ymm11\n\t"
"vxorps %%ymm12, %%ymm12, %%ymm12\n\t"
"vxorps %%ymm13, %%ymm13, %%ymm13\n\t"
"vxorps %%ymm14, %%ymm14, %%ymm14\n\t"
"vxorps %%ymm15, %%ymm15, %%ymm15\n\t"
"prefetcht0 (%%rcx)\n\t"
"prefetcht0 (%%rcx,%%rdi)\n\t"
"prefetcht0 (%%rcx,%%rdi,2)\n\t"
"prefetcht0 (%%rdx)\n\t"
"movl $64, %%esi\n\t"
".p2align 5\n\t"
"1:\n\t"
"prefetcht0 256(%%rax)\n\t"
"vmovaps 0(%%rbx), %%ymm0\n\t"
"vmovaps 32(%%rbx), %%ymm1\n\t"
"vmovaps 64(%%rbx), %%ymm2\n\t"
"vbroadcastss 0(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm4\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm5\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm6\n\t"
"vbroadcastss 4(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm7\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm8\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm9\n\t"
"vbroadcastss 8(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm10\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm11\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm12\n\t"
"vbroadcastss 12(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm13\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm14\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm15\n\t"
"vmovaps 96(%%rbx), %%ymm0\n\t"
"vmovaps 128(%%rbx), %%ymm1\n\t"
"vmovaps 160(%%rbx), %%ymm2\n\t"
"vbroadcastss 16(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm4\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm5\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm6\n\t"
"vbroadcastss 20(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm7\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm8\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm9\n\t"
"vbroadcastss 24(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm10\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm11\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm12\n\t"
"vbroadcastss 28(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm13\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm14\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm15\n\t"
"vmovaps 192(%%rbx), %%ymm0\n\t"
"vmovaps 224(%%rbx), %%ymm1\n\t"
"vmovaps 256(%%rbx), %%ymm2\n\t"
"vbroadcastss 32(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm4\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm5\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm6\n\t"
"vbroadcastss 36(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm7\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm8\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm9\n\t"
"vbroadcastss 40(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm10\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm11\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm12\n\t"
"vbroadcastss 44(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm13\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm14\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm15\n\t"
"vmovaps 288(%%rbx), %%ymm0\n\t"
"vmovaps 320(%%rbx), %%ymm1\n\t"
"vmovaps 352(%%rbx), %%ymm2\n\t"
"vbroadcastss 48(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm4\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm5\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm6\n\t"
"vbroadcastss 52(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm7\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm8\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm9\n\t"
"vbroadcastss 56(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm10\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm11\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm12\n\t"
"vbroadcastss 60(%%rax), %%ymm3\n\t"
"vfmadd231ps %%ymm0, %%ymm3, %%ymm13\n\t"
"vfmadd231ps %%ymm1, %%ymm3, %%ymm14\n\t"
"vfmadd231ps %%ymm2, %%ymm3, %%ymm15\n\t"
"addq $64, %%rax\n\t"
"addq $384, %%rbx\n\t"
"decl %%esi\n\t"
"jnz 1b\n\t"
"vaddps 0(%%rcx), %%ymm4, %%ymm4\n\t"
"vmovaps %%ymm4, 0(%%rcx)\n\t"
"vaddps 32(%%rcx), %%ymm5, %%ymm5\n\t"
"vmovaps %%ymm5, 32(%%rcx)\n\t"
"vaddps 64(%%rcx), %%ymm6, %%ymm6\n\t"
"vmovaps %%ymm6, 64(%%rcx)\n\t"
"vaddps 0(%%rcx,%%rdi), %%ymm7, %%ymm7\n\t"
"vmovaps %%ymm7, 0(%%rcx,%%rdi)\n\t"
"vaddps 32(%%rcx,%%rdi), %%ymm8, %%ymm8\n\t"
"vmovaps %%ymm8, 32(%%rcx,%%rdi)\n\t"
"vaddps 64(%%rcx,%%rdi), %%ymm9, %%ymm9\n\t"
"vmovaps %%ymm9, 64(%%rcx,%%rdi)\n\t"
"vaddps 0(%%rcx,%%rdi,2), %%ymm10, %%ymm10\n\t"
"vmovaps %%ymm10, 0(%%rcx,%%rdi,2)\n\t"
"vaddps 32(%%rcx,%%rdi,2), %%ymm11, %%ymm11\n\t"
"vmovaps %%ymm11, 32(%%rcx,%%rdi,2)\n\t"
"vaddps 64(%%rcx,%%rdi,2), %%ymm12, %%ymm12\n\t"
"vmovaps %%ymm12, 64(%%rcx,%%rdi,2)\n\t"
"vaddps 0(%%rdx), %%ymm13, %%ymm13\n\t"
"vmovaps %%ymm13, 0(%%rdx)\n\t"
"vaddps 32(%%rdx), %%ymm14, %%ymm14\n\t"
"vmovaps %%ymm14, 32(%%rdx)\n\t"
"vaddps 64(%%rdx), %%ymm15, %%ymm15\n\t"
"vmovaps %%ymm15, 64(%%rdx)\n\t"
: "+a"(aa), "+b"(bb)
: "c"(c), "D"(ld)
: "rdx", "rsi", "cc", "memory", "ymm0", "ymm1", "ymm2", "ymm3", "ymm4", "ymm5", "ymm6", "ymm7", "ymm8", "ymm9", "ymm10", "ymm11", "ymm12", "ymm13", "ymm14", "ymm15");
}
static void multiply_rows(int n,int row_count,const float*A,const float*B,float*C){
if(n%KC || row_count%MR || n>MAXN){
memset(C,0,(size_t)row_count*n*sizeof(float));
for(int i=0;i<row_count;++i)for(int k=0;k<n;++k)for(int j=0;j<n;++j)
C[(size_t)i*n+j]+=A[(size_t)i*n+k]*B[(size_t)k*n+j];
return;
}
const int padded_n=(n+NC-1)/NC*NC,groups=padded_n/NR;
float*bp=nullptr;
if(posix_memalign((void**)&bp,4096,(size_t)n*padded_n*sizeof(float)))abort();
for(int pc=0;pc<n;pc+=KC)for(int q=0;q<padded_n;q+=NR)for(int k=0;k<KC;++k){
float*out=bp+((pc/KC*groups+q/NR)*KC+k)*NR;
const float*src=B+(size_t)(pc+k)*n;
for(int j=0;j<NR;j+=8)
_mm256_store_ps(out+j,q+j<n?_mm256_loadu_ps(src+q+j):_mm256_setzero_ps());
}
for(int ic=0;ic<row_count;ic+=MC){
const int rows=row_count-ic<MC?row_count-ic:MC;
for(int pc=0;pc<n;pc+=KC)for(int r=0;r<rows;r+=MR)for(int k=0;k<KC;++k)for(int i=0;i<MR;++i)
ap[((pc/KC*(MC/MR)+r/MR)*KC+k)*MR+i]=A[(size_t)(ic+r+i)*n+pc+k];
for(int jc=0;jc<padded_n;jc+=NC){
memset(ct,0,sizeof(ct));
for(int pc=0;pc<n;pc+=KC)for(int q=0;q<NC;q+=NR)for(int r=0;r<rows;r+=MR){
const float*a=ap+((pc/KC)*(MC/MR)+r/MR)*KC*MR;
const float*b=bp+((pc/KC)*groups+(jc+q)/NR)*KC*NR;
kernel(a,b,ct+(size_t)r*NC+q,NC);
}
const int cols=n-jc<NC?n-jc:NC;
for(int i=0;i<rows;++i)for(int j=0;j<cols;j+=8)
_mm256_stream_ps(C+(size_t)(ic+i)*n+jc+j,_mm256_load_ps(ct+(size_t)i*NC+j));
}
}
_mm_sfence();free(bp);
}
}
void matrix_multiply(int n,const float*A,const float*B,float*C){tile4x24::multiply_rows(n,n,A,B,C);}
// DUCK_PUBLIC_SURROGATE_BENCHMARK
// Independent public analytic matrices; no random seed or production data.
#include <stdio.h>
#include <math.h>
int main(){
constexpr int n=8192;constexpr size_t z=(size_t)n*n;
float*A=nullptr,*B=nullptr,*C=nullptr;
if(posix_memalign((void**)&A,4096,z*4)||posix_memalign((void**)&B,4096,z*4)||posix_memalign((void**)&C,4096,z*4))return 1;
for(int i=0;i<n;++i)for(int j=0;j<n;++j){
A[(size_t)i*n+j]=float(((i*23+j*7+19)%1009)-504)/504.f;
B[(size_t)i*n+j]=float(((i*37+j*11+29)%1013)-506)/506.f;
}
matrix_multiply(n,A,B,C);
double maxerr=0;
for(int t=0;t<64;++t){int i=(t*127+31)%n,j=(t*337+53)%n;double v=0;for(int k=0;k<n;++k)v+=(double)A[(size_t)i*n+k]*B[(size_t)k*n+j];double e=fabs(v-C[(size_t)i*n+j]);if(e>maxerr)maxerr=e;if(e>0.1)return 2;}
printf("classical_full_matrix_checked max_error=%g\n",maxerr);return 0;
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 10.543 s | 1027 MB + 580 KB | Accepted | Score: 100 | 显示更多 |