// wc2017b2 AVX2 v6: complement + nibble popcount, inner loop in inline asm (no reloads).
#include <stdint.h>
#include <string.h>
#include <immintrin.h>
#pragma GCC target("avx2")
typedef uint64_t u64;
#define NW_MAX 4690
#define PAD_MAX (NW_MAX+32)
#define TOTAL_MAX (3*NW_MAX+64)
static u64 A1[TOTAL_MAX] __attribute__((aligned(32)));
static u64 A2[TOTAL_MAX] __attribute__((aligned(32)));
static u64 P1[8][TOTAL_MAX] __attribute__((aligned(32)));
static u64 P2[8][TOTAL_MAX] __attribute__((aligned(32)));
static inline unsigned popc32(unsigned x){ unsigned r; __asm__ __volatile__("popcnt %1,%0":"=r"(r):"r"(x)); return r; }
static inline unsigned popc(u64 x){ return popc32((unsigned)x)+popc32((unsigned)(x>>32)); }
static inline u64 loadu64(const void* p){ u64 v; __builtin_memcpy(&v,p,8); return v; }
static int NW, PAD, TOTAL;
static inline __m256i run_quads(const __m256i* va1, const __m256i* va2, const __m256i* vb1, const __m256i* vb2, int n, __m256i acc){
__m256i m0f = _mm256_set1_epi8(0x0f);
__m256i tab = _mm256_broadcastsi128_si256(_mm_setr_epi8(0,1,1,2,1,2,2,3,1,2,2,3,2,3,3,4));
#ifdef __x86_64__
__asm__ __volatile__(
"movl %[n], %%ecx\n\t"
"xorl %%edi, %%edi\n\t"
"testl %%ecx, %%ecx\n\t"
"jz 1f\n\t"
"0:\n\t"
"vmovdqu (%[va1],%%rdi), %%ymm0\n\t"
"vmovdqu (%[va2],%%rdi), %%ymm1\n\t"
"vmovdqu (%[vb1],%%rdi), %%ymm2\n\t"
"vmovdqu (%[vb2],%%rdi), %%ymm3\n\t"
"addq $32, %%rdi\n\t"
"vpand %%ymm3, %%ymm0, %%ymm7\n\t"
"vpor %%ymm1, %%ymm0, %%ymm0\n\t"
"vpandn %%ymm2, %%ymm0, %%ymm0\n\t"
"vpor %%ymm3, %%ymm2, %%ymm3\n\t"
"vpandn %%ymm1, %%ymm3, %%ymm1\n\t"
"vpor %%ymm0, %%ymm7, %%ymm7\n\t"
"vpor %%ymm1, %%ymm7, %%ymm7\n\t"
"vpand %[m0f], %%ymm7, %%ymm0\n\t"
"vpsrlw $4, %%ymm7, %%ymm1\n\t"
"vpand %[m0f], %%ymm1, %%ymm1\n\t"
"vpshufb %%ymm0, %[tab], %%ymm0\n\t"
"vpshufb %%ymm1, %[tab], %%ymm1\n\t"
"vpaddb %%ymm1, %%ymm0, %%ymm0\n\t"
"vpxor %%ymm1, %%ymm1, %%ymm1\n\t"
"vpsadbw %%ymm1, %%ymm0, %%ymm0\n\t"
"vpaddq %%ymm0, %[acc], %[acc]\n\t"
"decl %%ecx\n\t"
"jnz 0b\n\t"
"1:\n\t"
: [acc] "+x"(acc)
: [va1] "r"(va1), [va2] "r"(va2), [vb1] "r"(vb1), [vb2] "r"(vb2), [n] "m"(n), [tab] "x"(tab), [m0f] "x"(m0f)
: "ecx", "rdi", "ymm0", "ymm1", "ymm2", "ymm3", "ymm7", "memory"
);
#else
__asm__ __volatile__(
"movl %[n], %%ecx\n\t"
"xorl %%edi, %%edi\n\t"
"testl %%ecx, %%ecx\n\t"
"jz 1f\n\t"
"0:\n\t"
"vmovdqu (%[va1],%%edi), %%ymm0\n\t"
"vmovdqu (%[va2],%%edi), %%ymm1\n\t"
"vmovdqu (%[vb1],%%edi), %%ymm2\n\t"
"vmovdqu (%[vb2],%%edi), %%ymm3\n\t"
"addl $32, %%edi\n\t"
"vpand %%ymm3, %%ymm0, %%ymm7\n\t"
"vpor %%ymm1, %%ymm0, %%ymm0\n\t"
"vpandn %%ymm2, %%ymm0, %%ymm0\n\t"
"vpor %%ymm3, %%ymm2, %%ymm3\n\t"
"vpandn %%ymm1, %%ymm3, %%ymm1\n\t"
"vpor %%ymm0, %%ymm7, %%ymm7\n\t"
"vpor %%ymm1, %%ymm7, %%ymm7\n\t"
"vpand %[m0f], %%ymm7, %%ymm0\n\t"
"vpsrlw $4, %%ymm7, %%ymm1\n\t"
"vpand %[m0f], %%ymm1, %%ymm1\n\t"
"vpshufb %%ymm0, %[tab], %%ymm0\n\t"
"vpshufb %%ymm1, %[tab], %%ymm1\n\t"
"vpaddb %%ymm1, %%ymm0, %%ymm0\n\t"
"vpxor %%ymm1, %%ymm1, %%ymm1\n\t"
"vpsadbw %%ymm1, %%ymm0, %%ymm0\n\t"
"vpaddq %%ymm0, %[acc], %[acc]\n\t"
"decl %%ecx\n\t"
"jnz 0b\n\t"
"1:\n\t"
: [acc] "+x"(acc)
: [va1] "r"(va1), [va2] "r"(va2), [vb1] "r"(vb1), [vb2] "r"(vb2), [n] "m"(n), [tab] "x"(tab), [m0f] "x"(m0f)
: "ecx", "edi", "ymm0", "ymm1", "ymm2", "ymm3", "ymm7", "memory"
);
#endif
return acc;
}
void solve(int n,int q,char*s1,char*s2,int*q_x,int*q_y,int*q_len,unsigned*ans){
NW = (n+63)>>6;
PAD = NW+32;
TOTAL = 3*NW+64;
for(int j=0;j<TOTAL;j++){ A1[j]=0; A2[j]=0; for(int r=0;r<8;r++){ P1[r][j]=0; P2[r][j]=0; } }
for(int i=0;i<n;i++){
int a=(unsigned char)s1[i]; if(a>2)a-='0';
int b=(unsigned char)s2[i]; if(b>2)b-='0';
int w=i>>6; u64 bit=1ULL<<(i&63);
if(a==1)A1[PAD+w]|=bit; else if(a==2)A2[PAD+w]|=bit;
if(b==1)P1[0][PAD+w]|=bit; else if(b==2)P2[0][PAD+w]|=bit;
}
for(int r=1;r<8;r++){
u64 *s=P1[0], *d=P1[r]; int sh=64-r;
for(int j=-NW; j<=2*NW+1; j++) d[PAD+j]=(s[PAD+j+1]<<sh)|(s[PAD+j]>>r);
s=P2[0]; d=P2[r];
for(int j=-NW; j<=2*NW+1; j++) d[PAD+j]=(s[PAD+j+1]<<sh)|(s[PAD+j]>>r);
}
for(int qi=0;qi<q;qi++){
int x=q_x[qi], y=q_y[qi], l=q_len[qi];
long long d=(long long)y-x;
int off=(int)(d&63), woff=(int)(d>>6);
int r=off&7, qb=off>>3;
int w_start=x>>6, boff=x&63;
int w_end=(x+l-1)>>6, beoff=(x+l-1)&63;
int nw=w_end-w_start+1;
const char *b1=(const char*)&P1[r][0] + (size_t)(PAD+w_start+woff)*8 + qb;
const char *b2=(const char*)&P2[r][0] + (size_t)(PAD+w_start+woff)*8 + qb;
u64 *a1=A1+PAD+w_start, *a2=A2+PAD+w_start;
unsigned cnt=0;
if(nw<=6){
for(int w=0;w<nw;w++){
u64 s1=loadu64(b1+(size_t)w*8), s2=loadu64(b2+(size_t)w*8);
u64 a12=a1[w]|a2[w], b12=s1|s2;
u64 win=((~a12)&s1)|(a1[w]&s2)|(a2[w]&(~b12));
if(w==0 && boff) win &= (~0ULL)<<boff;
if(w==nw-1 && beoff!=63) win &= (1ULL<<(beoff+1))-1;
cnt+=popc(win);
}
ans[qi]=cnt;
continue;
}
{
u64 s1=loadu64(b1), s2=loadu64(b2);
u64 a12=a1[0]|a2[0], b12=s1|s2;
u64 win=((~a12)&s1)|(a1[0]&s2)|(a2[0]&(~b12));
if(boff) win &= (~0ULL)<<boff;
cnt=popc(win);
}
int mid_start=1, mid_end=nw-2;
int mcount=mid_end-mid_start+1;
int w=mid_start;
int quads=mcount>>2;
__m256i acc=_mm256_setzero_si256();
acc = run_quads((const __m256i*)(a1+w), (const __m256i*)(a2+w), (const __m256i*)(b1+(size_t)w*8), (const __m256i*)(b2+(size_t)w*8), quads, acc);
int done=quads*4;
for(int wl=mid_start+done; wl<=mid_end; wl++){
u64 s1=loadu64(b1+(size_t)wl*8), s2=loadu64(b2+(size_t)wl*8);
u64 a12=a1[wl]|a2[wl], b12=s1|s2;
u64 win=((~a12)&s1)|(a1[wl]&s2)|(a2[wl]&(~b12));
cnt+=popc(win);
}
{
__m128i lo128=_mm256_castsi256_si128(acc);
__m128i hi128=_mm256_extracti128_si256(acc,1);
u64 lo, hi;
_mm_storel_epi64((__m128i*)&lo, lo128);
_mm_storel_epi64((__m128i*)&hi, _mm_srli_si128(lo128,8));
cnt += (unsigned)(lo+hi);
_mm_storel_epi64((__m128i*)&lo, hi128);
_mm_storel_epi64((__m128i*)&hi, _mm_srli_si128(hi128,8));
cnt += (unsigned)(lo+hi);
}
{
int wl=nw-1;
u64 s1=loadu64(b1+(size_t)wl*8), s2=loadu64(b2+(size_t)wl*8);
u64 a12=a1[wl]|a2[wl], b12=s1|s2;
u64 win=((~a12)&s1)|(a1[wl]&s2)|(a2[wl]&(~b12));
if(beoff!=63) win &= (1ULL<<(beoff+1))-1;
cnt+=popc(win);
}
ans[qi]=cnt;
}
}
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 183.46 us | 128 KB | Accepted | Score: 50 | 显示更多 |
| Testcase #2 | 238.295 ms | 7 MB + 116 KB | Accepted | Score: 50 | 显示更多 |