// Four-dimensional strict dominance via event divide-and-conquer and SIMD leaves.
// Algorithm based on Duck accepted public solution #101855; #120527 also reviewed.
// Event layout in 128 bits: key (22 bits), insert flag (1), next key (22), ID (22).
// Array/counter capacities support n <= 3,000,000 without field truncation.
// This source contains the solver alone.
#ifndef MF_CTRPF
#define MF_CTRPF 0
#endif
#ifndef MF_CTR5
#define MF_CTR5 0
#endif
#ifndef MF_LEAFLP
#define MF_LEAFLP 64
#endif
#ifndef MF_LEAFG
#define MF_LEAFG 4
#endif
#pragma GCC optimize("O3","unroll-loops")
#pragma GCC target("avx2,bmi,bmi2,popcnt,lzcnt")
#include <cstring>
#include <immintrin.h>
#pragma GCC push_options
#pragma GCC target("avx2,popcnt,bmi2")
#include <cstdint>
typedef unsigned u32;
typedef __uint128_t u64;
typedef unsigned char u8;
typedef unsigned short u16;
#ifndef KG_COND
#define KG_COND 0
#endif
#ifndef KG_PF
#define KG_PF 0
#endif
#ifndef KG_BASE
#define KG_BASE 2048
#endif
#ifndef LEAFNB
#define LEAFNB 2048
#endif
#define KG_MAXK 5
#define KG_ARENA (48u << 20)
namespace kg {
#define CT_MAXV 4194308u
static u8 CT0[CT_MAXV];
static u16 CT1[CT_MAXV / 32 + 8];
static u16 CT2[CT_MAXV / 256 + 8];
static u32 CT3[CT_MAXV / 2048 + 8];
static u32 CT4[CT_MAXV / 16384 + 8];
static u32 CT5[CT_MAXV / 131072 + 8];
static int FUSE_OK;
static inline u32 su8_32(const u8 *p, int cnt) {
__m256i v = _mm256_loadu_si256((const __m256i *)p);
__m256i m = _mm256_cmpgt_epi8(_mm256_set1_epi8((char)cnt),
_mm256_setr_epi8(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15,
16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31));
__m256i x = _mm256_and_si256(v, m);
__m256i s = _mm256_sad_epu8(x, _mm256_setzero_si256());
__m128i t = _mm_add_epi64(_mm256_castsi256_si128(s), _mm256_extracti128_si256(s, 1));
return (u32)((unsigned long long)_mm_cvtsi128_si64(t) + (unsigned long long)_mm_extract_epi64(t, 1));
}
static inline u32 su16_8(const u16 *p, int cnt) {
__m128i v = _mm_loadu_si128((const __m128i *)p);
__m128i x = _mm_and_si128(v, _mm_cmpgt_epi16(_mm_set1_epi16((short)cnt),
_mm_setr_epi16(0, 1, 2, 3, 4, 5, 6, 7)));
__m128i s = _mm_madd_epi16(x, _mm_set1_epi16(1));
s = _mm_add_epi32(s, _mm_shuffle_epi32(s, _MM_SHUFFLE(1, 0, 3, 2)));
s = _mm_add_epi32(s, _mm_shuffle_epi32(s, _MM_SHUFFLE(2, 3, 0, 1)));
return (u32)_mm_cvtsi128_si32(s);
}
static inline u32 su32_8(const u32 *p, int cnt) {
__m256i v = _mm256_loadu_si256((const __m256i *)p);
__m256i x = _mm256_and_si256(v, _mm256_cmpgt_epi32(_mm256_set1_epi32(cnt),
_mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7)));
__m128i t = _mm_add_epi32(_mm256_castsi256_si128(x), _mm256_extracti128_si256(x, 1));
t = _mm_add_epi32(t, _mm_shuffle_epi32(t, _MM_SHUFFLE(1, 0, 3, 2)));
t = _mm_add_epi32(t, _mm_shuffle_epi32(t, _MM_SHUFFLE(2, 3, 0, 1)));
return (u32)_mm_cvtsi128_si32(t);
}
static inline void ctr_ins(u32 v) {
CT0[v]++; CT1[v >> 5]++; CT2[v >> 8]++; CT3[v >> 11]++; CT4[v >> 14]++; CT5[v >> 17]++;
}
static inline void ctr_clr(u32 v) {
CT0[v]--; CT1[v >> 5]--; CT2[v >> 8]--; CT3[v >> 11]--; CT4[v >> 14]--; CT5[v >> 17]--;
}
static inline u32 ctr_qry(u32 v) {
u32 k5 = v >> 17, s = 0;
#if MF_CTR5
{
__m256i tv = _mm256_loadu_si256((const __m256i *)CT5);
__m256i tm = _mm256_cmpgt_epi32(_mm256_set1_epi32((int)k5),
_mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7));
__m128i ta = _mm_add_epi32(_mm256_castsi256_si128(_mm256_and_si256(tv, tm)),
_mm256_extracti128_si256(_mm256_and_si256(tv, tm), 1));
ta = _mm_add_epi32(ta, _mm_shuffle_epi32(ta, _MM_SHUFFLE(1, 0, 3, 2)));
ta = _mm_add_epi32(ta, _mm_shuffle_epi32(ta, _MM_SHUFFLE(2, 3, 0, 1)));
s += (u32)_mm_cvtsi128_si32(ta);
}
#else
for (u32 j = 0; j < k5; j++) s += CT5[j];
#endif
s += su32_8(CT4 + ((v >> 17) << 3), (int)((v >> 14) & 7));
s += su32_8(CT3 + ((v >> 14) << 3), (int)((v >> 11) & 7));
s += su16_8(CT2 + ((v >> 11) << 3), (int)((v >> 8) & 7));
s += su16_8(CT1 + ((v >> 8) << 3), (int)((v >> 5) & 7));
s += su8_32(CT0 + (v & ~31u), (int)(v & 31));
return s;
}
#ifndef LFAUTO
#define LFAUTO 0
#endif
static int LFSH;
static int LFD;
static int LFB;
static const u32 *XD[KG_MAXK];
static u32 *OUT;
static u64 *AB, *ABtop;
static inline int eid(u64 v) { return (int)(v & 0x3FFFFFu); }
static inline int lfbucket(u64 v) { return (int)(((u32)(v >> 45)) >> LFSH); }
static inline int eq_(u64 v) { return (int)((v >> 44) & 1); }
static inline int keep_test(int takeA, int role) { return 1 - (takeA ^ role); }
static inline u32 embnk(u64 v) { return (u32)((v >> 22) & 0x3FFFFFu); }
static inline u64 mk(u32 key, int q, int id, u32 nk) {
return ((u64)key << 45) | ((u64)(1 - q) << 44) | ((u64)nk << 22) | (u64)id;
}
template <bool B> struct TB {};
template <int V> struct NCL { static const int v = (V < 0) ? 0 : ((V > 3) ? 3 : V); };
template <int KT, int D, int NO> static inline void solve_t(u64 *seq, int m, u64 *out);
static void ctr_solve(u64 *cross, int cx, u32 *dirt) {
int nd = 0;
for (int i = 0; i < cx; i++) {
#if MF_CTRPF > 0
if (i + MF_CTRPF < cx) {
u32 vf = (u32)(cross[i + MF_CTRPF] >> 45);
__builtin_prefetch(&CT0[vf], 1, 0);
__builtin_prefetch(&CT1[vf >> 5], 1, 0);
__builtin_prefetch(&CT2[vf >> 8], 1, 0);
}
#endif
u64 v = cross[i];
u32 val = (u32)(v >> 45);
if ((v >> 44) & 1) { ctr_ins(val); dirt[nd++] = val; }
else OUT[eid(v)] += ctr_qry(val);
}
for (int i = 0; i < nd; i++) {
#if MF_CTRPF > 0
if (i + MF_CTRPF < nd) {
u32 vf = dirt[i + MF_CTRPF];
__builtin_prefetch(&CT0[vf], 1, 1);
__builtin_prefetch(&CT1[vf >> 5], 1, 1);
__builtin_prefetch(&CT2[vf >> 8], 1, 1);
}
#endif
ctr_clr(dirt[i]);
}
}
static const u32 VMASK[8][8] = {
{0,0,0,0,0,0,0,0},
{~0u,0,0,0,0,0,0,0},
{~0u,~0u,0,0,0,0,0,0},
{~0u,~0u,~0u,0,0,0,0,0},
{~0u,~0u,~0u,~0u,0,0,0,0},
{~0u,~0u,~0u,~0u,~0u,0,0,0},
{~0u,~0u,~0u,~0u,~0u,~0u,0,0},
{~0u,~0u,~0u,~0u,~0u,~0u,~0u,0}};
static inline u32 hsum256(__m256i v) {
__m128i t = _mm_add_epi32(_mm256_castsi256_si128(v), _mm256_extracti128_si256(v, 1));
t = _mm_add_epi32(t, _mm_shuffle_epi32(t, _MM_SHUFFLE(1, 0, 3, 2)));
t = _mm_add_epi32(t, _mm_shuffle_epi32(t, _MM_SHUFFLE(2, 3, 0, 1)));
return (u32)_mm_cvtsi128_si32(t);
}
template<int ND>
static void base_case_t(u64 *seq, int m, int d, u64 *out, int needOut) {
const int nd = ND;
__attribute__((aligned(32)))
u32 IK[KG_BASE + 8], IC0[KG_BASE + 8], IC1[KG_BASE + 8], IC2[KG_BASE + 8];
u32 QK[KG_BASE + 8], QC0[KG_BASE + 8], QC1[KG_BASE + 8], QC2[KG_BASE + 8];
u32 QID[KG_BASE + 8], QN[KG_BASE + 8];
const u32 *X1 = (ND > 1) ? XD[d + 3] : 0, *X2 = (ND > 2) ? XD[d + 4] : 0;
int ni = 0, nq = 0;
for (int t = 0; t < m; t++) {
#if MF_LEAFLP > 0
if ((ND > 1) && t + MF_LEAFLP < m) {
int idf = eid(seq[t + MF_LEAFLP]);
if (ND > 2) __builtin_prefetch(&X2[idf], 0, 1);
__builtin_prefetch(&X1[idf], 0, 1);
}
#endif
u64 v = seq[t];
int id = eid(v);
u32 k = (u32)(v >> 45);
u32 c0 = 0, c1 = 0, c2 = 0;
if (ND > 0) c0 = embnk(v);
if (ND > 1) c1 = X1[id];
if (ND > 2) c2 = X2[id];
int isins = (int)((v >> 44) & 1u);
IK[ni] = k; IC0[ni] = c0; IC1[ni] = c1; IC2[ni] = c2; ni += isins;
QK[nq] = k; QC0[nq] = c0; QC1[nq] = c1; QC2[nq] = c2; QID[nq] = (u32)id; QN[nq] = (u32)ni; nq += 1 - isins;
}
#if !defined(MF_G34_LDA)
#define MF_G34_LDA(p) _mm256_load_si256((const __m256i *)(p))
#define MF_G34_LDU(p) _mm256_loadu_si256((const __m256i *)(p))
#endif
#define LDA(p) MF_G34_LDA(p)
#define LDU_MASK(k) MF_G34_LDU((const __m256i *)VMASK[k])
#define K11B_TAIL(ACC, KVQ, C0Q, C1Q, C2Q, NX, JS) do { \
int jj = (JS); \
for (; jj + 8 <= (NX); jj += 8) { \
__m256i msk = _mm256_cmpgt_epi32(KVQ, LDA(IK + jj)); \
if (nd > 0) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(C0Q, LDA(IC0 + jj))); \
if (nd > 1) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(C1Q, LDA(IC1 + jj))); \
if (nd > 2) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(C2Q, LDA(IC2 + jj))); \
ACC = _mm256_sub_epi32(ACC, msk); \
} \
if (jj < (NX)) { \
int nn_ = (NX); \
__m256i msk = _mm256_cmpgt_epi32(KVQ, LDA(IK + jj)); \
if (nd > 0) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(C0Q, LDA(IC0 + jj))); \
if (nd > 1) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(C1Q, LDA(IC1 + jj))); \
if (nd > 2) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(C2Q, LDA(IC2 + jj))); \
msk = _mm256_and_si256(msk, LDU_MASK(nn_ - jj)); \
ACC = _mm256_sub_epi32(ACC, msk); \
} \
} while (0)
int q = 0;
#if MF_LEAFG >= 4
if (nd <= 1) {
for (; q + 4 <= nq; q += 4) {
if (QN[q] > QN[q+1] || QN[q+1] > QN[q+2] || QN[q+2] > QN[q+3]) break;
const int nA4 = (int)QN[q], nB4 = (int)QN[q+1], nC4 = (int)QN[q+2], nD4 = (int)QN[q+3];
__m256i kA4 = _mm256_set1_epi32((int)QK[q]), dA4 = _mm256_set1_epi32((int)QC0[q]),
aA4 = _mm256_setzero_si256();
__m256i kB4 = _mm256_set1_epi32((int)QK[q+1]), dB4 = _mm256_set1_epi32((int)QC0[q+1]),
aB4 = _mm256_setzero_si256();
__m256i kC4 = _mm256_set1_epi32((int)QK[q+2]), dC4 = _mm256_set1_epi32((int)QC0[q+2]),
aC4 = _mm256_setzero_si256();
__m256i kD4 = _mm256_set1_epi32((int)QK[q+3]), dD4 = _mm256_set1_epi32((int)QC0[q+3]),
aD4 = _mm256_setzero_si256();
int j4 = 0;
for (; j4 + 8 <= nA4; j4 += 8) {
__m256i IKv = LDA(IK + j4);
__m256i mA = _mm256_cmpgt_epi32(kA4, IKv), mB = _mm256_cmpgt_epi32(kB4, IKv),
mC = _mm256_cmpgt_epi32(kC4, IKv), mD = _mm256_cmpgt_epi32(kD4, IKv);
if (nd > 0) {
__m256i C0v = LDA(IC0 + j4);
mA = _mm256_and_si256(mA, _mm256_cmpgt_epi32(dA4, C0v));
mB = _mm256_and_si256(mB, _mm256_cmpgt_epi32(dB4, C0v));
mC = _mm256_and_si256(mC, _mm256_cmpgt_epi32(dC4, C0v));
mD = _mm256_and_si256(mD, _mm256_cmpgt_epi32(dD4, C0v));
}
aA4 = _mm256_sub_epi32(aA4, mA); aB4 = _mm256_sub_epi32(aB4, mB);
aC4 = _mm256_sub_epi32(aC4, mC); aD4 = _mm256_sub_epi32(aD4, mD);
}
const int pA4 = j4;
for (; j4 + 8 <= nB4; j4 += 8) {
__m256i IKv = LDA(IK + j4);
__m256i mB = _mm256_cmpgt_epi32(kB4, IKv), mC = _mm256_cmpgt_epi32(kC4, IKv),
mD = _mm256_cmpgt_epi32(kD4, IKv);
if (nd > 0) {
__m256i C0v = LDA(IC0 + j4);
mB = _mm256_and_si256(mB, _mm256_cmpgt_epi32(dB4, C0v));
mC = _mm256_and_si256(mC, _mm256_cmpgt_epi32(dC4, C0v));
mD = _mm256_and_si256(mD, _mm256_cmpgt_epi32(dD4, C0v));
}
aB4 = _mm256_sub_epi32(aB4, mB); aC4 = _mm256_sub_epi32(aC4, mC);
aD4 = _mm256_sub_epi32(aD4, mD);
}
const int pB4 = j4;
for (; j4 + 8 <= nC4; j4 += 8) {
__m256i IKv = LDA(IK + j4);
__m256i mC = _mm256_cmpgt_epi32(kC4, IKv), mD = _mm256_cmpgt_epi32(kD4, IKv);
if (nd > 0) {
__m256i C0v = LDA(IC0 + j4);
mC = _mm256_and_si256(mC, _mm256_cmpgt_epi32(dC4, C0v));
mD = _mm256_and_si256(mD, _mm256_cmpgt_epi32(dD4, C0v));
}
aC4 = _mm256_sub_epi32(aC4, mC); aD4 = _mm256_sub_epi32(aD4, mD);
}
const int pC4 = j4;
for (; j4 + 8 <= nD4; j4 += 8) {
__m256i mD = _mm256_cmpgt_epi32(kD4, LDA(IK + j4));
if (nd > 0) mD = _mm256_and_si256(mD, _mm256_cmpgt_epi32(dD4, LDA(IC0 + j4)));
aD4 = _mm256_sub_epi32(aD4, mD);
}
const int pD4 = j4;
K11B_TAIL(aA4, kA4, dA4, dA4, dA4, nA4, pA4);
K11B_TAIL(aB4, kB4, dB4, dB4, dB4, nB4, pB4);
K11B_TAIL(aC4, kC4, dC4, dC4, dC4, nC4, pC4);
K11B_TAIL(aD4, kD4, dD4, dD4, dD4, nD4, pD4);
OUT[QID[q]] += hsum256(aA4);
OUT[QID[q + 1]] += hsum256(aB4);
OUT[QID[q + 2]] += hsum256(aC4);
OUT[QID[q + 3]] += hsum256(aD4);
}
}
#endif
#if MF_LEAFG >= 3
if (nd <= 2) {
for (; q + 3 <= nq; q += 3) {
if (QN[q] > QN[q + 1] || QN[q + 1] > QN[q + 2]) break;
const int nA3 = (int)QN[q], nB3 = (int)QN[q + 1], nC3 = (int)QN[q + 2];
__m256i kA3 = _mm256_set1_epi32((int)QK[q]), dA3 = _mm256_set1_epi32((int)QC0[q]),
eA3 = _mm256_set1_epi32((int)QC1[q]), aA3 = _mm256_setzero_si256();
__m256i kB3 = _mm256_set1_epi32((int)QK[q+1]), dB3 = _mm256_set1_epi32((int)QC0[q+1]),
eB3 = _mm256_set1_epi32((int)QC1[q+1]), aB3 = _mm256_setzero_si256();
__m256i kC3 = _mm256_set1_epi32((int)QK[q+2]), dC3 = _mm256_set1_epi32((int)QC0[q+2]),
eC3 = _mm256_set1_epi32((int)QC1[q+2]), aC3 = _mm256_setzero_si256();
int j3 = 0;
for (; j3 + 8 <= nA3; j3 += 8) {
__m256i IKv = LDA(IK + j3);
__m256i mA = _mm256_cmpgt_epi32(kA3, IKv), mB = _mm256_cmpgt_epi32(kB3, IKv),
mC = _mm256_cmpgt_epi32(kC3, IKv);
if (nd > 0) {
__m256i C0v = LDA(IC0 + j3);
mA = _mm256_and_si256(mA, _mm256_cmpgt_epi32(dA3, C0v));
mB = _mm256_and_si256(mB, _mm256_cmpgt_epi32(dB3, C0v));
mC = _mm256_and_si256(mC, _mm256_cmpgt_epi32(dC3, C0v));
}
if (nd > 1) {
__m256i C1v = LDA(IC1 + j3);
mA = _mm256_and_si256(mA, _mm256_cmpgt_epi32(eA3, C1v));
mB = _mm256_and_si256(mB, _mm256_cmpgt_epi32(eB3, C1v));
mC = _mm256_and_si256(mC, _mm256_cmpgt_epi32(eC3, C1v));
}
aA3 = _mm256_sub_epi32(aA3, mA); aB3 = _mm256_sub_epi32(aB3, mB);
aC3 = _mm256_sub_epi32(aC3, mC);
}
const int pA3 = j3;
for (; j3 + 8 <= nB3; j3 += 8) {
__m256i IKv = LDA(IK + j3);
__m256i mB = _mm256_cmpgt_epi32(kB3, IKv), mC = _mm256_cmpgt_epi32(kC3, IKv);
if (nd > 0) {
__m256i C0v = LDA(IC0 + j3);
mB = _mm256_and_si256(mB, _mm256_cmpgt_epi32(dB3, C0v));
mC = _mm256_and_si256(mC, _mm256_cmpgt_epi32(dC3, C0v));
}
if (nd > 1) {
__m256i C1v = LDA(IC1 + j3);
mB = _mm256_and_si256(mB, _mm256_cmpgt_epi32(eB3, C1v));
mC = _mm256_and_si256(mC, _mm256_cmpgt_epi32(eC3, C1v));
}
aB3 = _mm256_sub_epi32(aB3, mB); aC3 = _mm256_sub_epi32(aC3, mC);
}
const int pB3 = j3;
for (; j3 + 8 <= nC3; j3 += 8) {
__m256i IKv = LDA(IK + j3);
__m256i mC = _mm256_cmpgt_epi32(kC3, IKv);
if (nd > 0) mC = _mm256_and_si256(mC, _mm256_cmpgt_epi32(dC3, LDA(IC0 + j3)));
if (nd > 1) mC = _mm256_and_si256(mC, _mm256_cmpgt_epi32(eC3, LDA(IC1 + j3)));
aC3 = _mm256_sub_epi32(aC3, mC);
}
const int pC3 = j3;
K11B_TAIL(aA3, kA3, dA3, eA3, eA3, nA3, pA3);
K11B_TAIL(aB3, kB3, dB3, eB3, eB3, nB3, pB3);
K11B_TAIL(aC3, kC3, dC3, eC3, eC3, nC3, pC3);
OUT[QID[q]] += hsum256(aA3);
OUT[QID[q + 1]] += hsum256(aB3);
OUT[QID[q + 2]] += hsum256(aC3);
}
}
#endif
#if MF_LEAFG >= 2
for (; q + 2 <= nq; q += 2) {
const int n = ((int)QN[q] < (int)QN[q + 1]) ? (int)QN[q] : (int)QN[q + 1];
const int nA = (int)QN[q], nB = (int)QN[q + 1];
__m256i kvA = _mm256_set1_epi32((int)QK[q]), kvB = _mm256_set1_epi32((int)QK[q + 1]);
__m256i v0A = _mm256_set1_epi32((int)QC0[q]), v0B = _mm256_set1_epi32((int)QC0[q + 1]);
__m256i v1A = _mm256_set1_epi32((int)QC1[q]), v1B = _mm256_set1_epi32((int)QC1[q + 1]);
__m256i v2A = _mm256_set1_epi32((int)QC2[q]), v2B = _mm256_set1_epi32((int)QC2[q + 1]);
__m256i accA = _mm256_setzero_si256(), accB = _mm256_setzero_si256();
int j = 0;
if (nd >= 3) {
for (; j + 8 <= n; j += 8) {
__m256i IKv = _mm256_load_si256((const __m256i *)(IK + j));
__m256i C0v = _mm256_load_si256((const __m256i *)(IC0 + j));
__m256i C1v = _mm256_load_si256((const __m256i *)(IC1 + j));
__m256i C2v = _mm256_load_si256((const __m256i *)(IC2 + j));
__m256i mA = _mm256_and_si256(_mm256_cmpgt_epi32(kvA, IKv), _mm256_cmpgt_epi32(v0A, C0v));
__m256i mB = _mm256_and_si256(_mm256_cmpgt_epi32(kvB, IKv), _mm256_cmpgt_epi32(v0B, C0v));
mA = _mm256_and_si256(mA, _mm256_cmpgt_epi32(v1A, C1v));
mB = _mm256_and_si256(mB, _mm256_cmpgt_epi32(v1B, C1v));
mA = _mm256_and_si256(mA, _mm256_cmpgt_epi32(v2A, C2v));
mB = _mm256_and_si256(mB, _mm256_cmpgt_epi32(v2B, C2v));
accA = _mm256_sub_epi32(accA, mA);
accB = _mm256_sub_epi32(accB, mB);
}
} else {
for (; j + 8 <= n; j += 8) {
__m256i IKv = _mm256_load_si256((const __m256i *)(IK + j));
__m256i mA = _mm256_cmpgt_epi32(kvA, IKv);
__m256i mB = _mm256_cmpgt_epi32(kvB, IKv);
if (nd > 0) {
__m256i C0v = _mm256_load_si256((const __m256i *)(IC0 + j));
mA = _mm256_and_si256(mA, _mm256_cmpgt_epi32(v0A, C0v));
mB = _mm256_and_si256(mB, _mm256_cmpgt_epi32(v0B, C0v));
}
if (nd > 1) {
__m256i C1v = _mm256_load_si256((const __m256i *)(IC1 + j));
mA = _mm256_and_si256(mA, _mm256_cmpgt_epi32(v1A, C1v));
mB = _mm256_and_si256(mB, _mm256_cmpgt_epi32(v1B, C1v));
}
accA = _mm256_sub_epi32(accA, mA);
accB = _mm256_sub_epi32(accB, mB);
}
}
{
int jj = j;
__m256i acc = accA;
for (; jj + 8 <= nA; jj += 8) {
__m256i msk = _mm256_cmpgt_epi32(kvA, _mm256_load_si256((const __m256i *)(IK + jj)));
if (nd > 0) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v0A, _mm256_load_si256((const __m256i *)(IC0 + jj))));
if (nd > 1) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v1A, _mm256_load_si256((const __m256i *)(IC1 + jj))));
if (nd > 2) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v2A, _mm256_load_si256((const __m256i *)(IC2 + jj))));
acc = _mm256_sub_epi32(acc, msk);
}
if (jj < nA) {
int n = nA;
__m256i msk = _mm256_cmpgt_epi32(kvA, _mm256_load_si256((const __m256i *)(IK + jj)));
if (nd > 0) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v0A, _mm256_load_si256((const __m256i *)(IC0 + jj))));
if (nd > 1) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v1A, _mm256_load_si256((const __m256i *)(IC1 + jj))));
if (nd > 2) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v2A, _mm256_load_si256((const __m256i *)(IC2 + jj))));
msk = _mm256_and_si256(msk, _mm256_loadu_si256((const __m256i *)VMASK[n - jj]));
acc = _mm256_sub_epi32(acc, msk);
}
OUT[QID[q]] += hsum256(acc);
}
{
int jj = j;
__m256i acc = accB;
for (; jj + 8 <= nB; jj += 8) {
__m256i msk = _mm256_cmpgt_epi32(kvB, _mm256_load_si256((const __m256i *)(IK + jj)));
if (nd > 0) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v0B, _mm256_load_si256((const __m256i *)(IC0 + jj))));
if (nd > 1) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v1B, _mm256_load_si256((const __m256i *)(IC1 + jj))));
if (nd > 2) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v2B, _mm256_load_si256((const __m256i *)(IC2 + jj))));
acc = _mm256_sub_epi32(acc, msk);
}
if (jj < nB) {
int n = nB;
__m256i msk = _mm256_cmpgt_epi32(kvB, _mm256_load_si256((const __m256i *)(IK + jj)));
if (nd > 0) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v0B, _mm256_load_si256((const __m256i *)(IC0 + jj))));
if (nd > 1) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v1B, _mm256_load_si256((const __m256i *)(IC1 + jj))));
if (nd > 2) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v2B, _mm256_load_si256((const __m256i *)(IC2 + jj))));
msk = _mm256_and_si256(msk, _mm256_loadu_si256((const __m256i *)VMASK[n - jj]));
acc = _mm256_sub_epi32(acc, msk);
}
OUT[QID[q + 1]] += hsum256(acc);
}
}
#endif
for (; q < nq; q++) {
const int n = (int)QN[q];
__m256i kv = _mm256_set1_epi32((int)QK[q]);
__m256i v0 = _mm256_set1_epi32((int)QC0[q]);
__m256i v1 = _mm256_set1_epi32((int)QC1[q]);
__m256i v2 = _mm256_set1_epi32((int)QC2[q]);
__m256i acc = _mm256_setzero_si256();
int j = 0;
switch (nd) {
case 0:
for (; j + 8 <= n; j += 8)
acc = _mm256_sub_epi32(acc, _mm256_cmpgt_epi32(kv, _mm256_load_si256((const __m256i *)(IK + j))));
break;
case 1:
for (; j + 8 <= n; j += 8) {
__m256i msk = _mm256_cmpgt_epi32(kv, _mm256_load_si256((const __m256i *)(IK + j)));
msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v0, _mm256_load_si256((const __m256i *)(IC0 + j))));
acc = _mm256_sub_epi32(acc, msk);
}
break;
case 2:
for (; j + 8 <= n; j += 8) {
__m256i msk = _mm256_cmpgt_epi32(kv, _mm256_load_si256((const __m256i *)(IK + j)));
msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v0, _mm256_load_si256((const __m256i *)(IC0 + j))));
msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v1, _mm256_load_si256((const __m256i *)(IC1 + j))));
acc = _mm256_sub_epi32(acc, msk);
}
break;
default:
for (; j + 8 <= n; j += 8) {
__m256i msk = _mm256_cmpgt_epi32(kv, _mm256_load_si256((const __m256i *)(IK + j)));
msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v0, _mm256_load_si256((const __m256i *)(IC0 + j))));
msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v1, _mm256_load_si256((const __m256i *)(IC1 + j))));
msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v2, _mm256_load_si256((const __m256i *)(IC2 + j))));
acc = _mm256_sub_epi32(acc, msk);
}
break;
}
if (j < n) {
__m256i msk = _mm256_cmpgt_epi32(kv, _mm256_load_si256((const __m256i *)(IK + j)));
if (nd > 0) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v0, _mm256_load_si256((const __m256i *)(IC0 + j))));
if (nd > 1) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v1, _mm256_load_si256((const __m256i *)(IC1 + j))));
if (nd > 2) msk = _mm256_and_si256(msk, _mm256_cmpgt_epi32(v2, _mm256_load_si256((const __m256i *)(IC2 + j))));
msk = _mm256_and_si256(msk, _mm256_loadu_si256((const __m256i *)VMASK[n - j]));
acc = _mm256_sub_epi32(acc, msk);
}
u32 c = hsum256(acc);
OUT[QID[q]] += c;
}
if (!needOut) return;
#if LEAFNB
{
const int NB = (LFAUTO && m <= LFB) ? (m < 8 ? 8 : m) : LFB;
u32 cnt[LEAFNB + 2];
for (int b = 0; b <= NB; b++) cnt[b] = 0;
for (int t = 0; t < m; t++) cnt[(int)(((u32)(seq[t] >> 45)) >> (LFSH + LFD)) + 1]++;
for (int b = 0; b < NB; b++) cnt[b + 1] += cnt[b];
for (int t = 0; t < m; t++) { u64 v = seq[t]; out[cnt[(int)(((u32)(v >> 45)) >> (LFSH + LFD))]++] = v; }
int lo = 0;
for (int b = 0; b < NB; b++) {
int hi = (int)cnt[b];
if (hi - lo > 1)
for (int t = lo + 1; t < hi; t++) {
u64 v = out[t];
int q = t - 1;
while (q >= lo && v < out[q]) { out[q + 1] = out[q]; q--; }
out[q + 1] = v;
}
lo = hi;
}
}
#else
for (int t = 1; t < m; t++) {
u64 v = seq[t];
int q = t - 1;
while (q >= 0 && v < seq[q]) { seq[q + 1] = seq[q]; q--; }
seq[q + 1] = v;
}
memcpy(out, seq, sizeof(u64) * (size_t)m);
#endif
}
#define MKE_CROSS(V, EMB) ((RK ? (((((V) >> 22) & 0x3FFFFFull) << 45) | ((V) & 17592190238719ull)) \
: ((V) & 17592190238719ull)) \
| (GAT ? ((u64)(EMB)[eid(V)] << 22) : 0ull))
#define MF_EMB_NONE 0
template<int NO, int RK, int GAT>
static inline void merge_loops(const u64 *retA, int h, const u64 *retB, int r,
u64 *out, u64 *cross, const u32 *XN, const u32 *XE, int &cx) {
(void)XN;
int ia = 0, ib = 0, t = 0;
while (ia < h && ib < r) {
u64 ea = retA[ia], eb = retB[ib];
int takeA = (ea < eb);
ia += takeA; ib += 1 - takeA;
u64 v = takeA ? ea : eb;
if (NO) out[t] = v;
t++;
int role = (int)((v >> 44) & 1);
u64 c = MKE_CROSS(v, XE);
cross[cx] = c;
cx += keep_test(takeA, role);
}
while (ia < h) {
u64 v = retA[ia++];
if (NO) out[t] = v;
t++;
if ((v >> 44) & 1) cross[cx++] = MKE_CROSS(v, XE);
}
while (ib < r) {
u64 v = retB[ib++];
if (NO) out[t] = v;
t++;
if (!((v >> 44) & 1)) cross[cx++] = MKE_CROSS(v, XE);
}
}
template <int KT, int D, int NO>
static void solve_b(u64 *seq, int m, u64 *out, TB<true>) {
(void)out;
int cnt = 0;
for (int t = 0; t < m; t++) {
u64 v = seq[t];
if ((v >> 44) & 1) cnt++;
else OUT[eid(v)] += (u32)cnt;
}
}
template <int KT, int D, int NO>
static void solve_b(u64 *seq, int m, u64 *out, TB<false>) {
if (m <= KG_BASE) {
base_case_t<NCL< KT - (D + 2) >::v>(seq, m, D, out, NO);
return;
}
int h = m >> 1, r = m - h;
u64 *save = ABtop;
u64 *retA = ABtop; ABtop += h;
u64 *retB = ABtop; ABtop += r;
solve_t<KT, D, 1>(seq, h, retA);
solve_t<KT, D, 1>(seq + h, r, retB);
u64 *cross = ABtop; ABtop += m + 1;
const int RK = (D + 2 <= KT - 1) ? 1 : 0;
const int GAT = (D + 3 <= KT - 1) ? 1 : 0;
const u32 *XN = RK ? XD[D + 2] : 0;
const u32 *XE = GAT ? XD[D + 3] : 0;
int cx = 0;
merge_loops<NO, RK, GAT>(retA, h, retB, r, out, cross, XN, XE, cx);
if (FUSE_OK && D == KT - 3) {
u32 *dirt = (u32 *)ABtop; ABtop += (cx + 3) / 2 + 1;
ctr_solve(cross, cx, dirt);
} else {
u64 *cres = ABtop; ABtop += cx;
solve_t<KT, D + 1, 0>(cross, cx, cres);
}
ABtop = save;
}
template <int KT, int D, int NO>
static inline void solve_t(u64 *seq, int m, u64 *out) {
if (m <= 1) return;
solve_b<KT, D, NO>(seq, m, out, TB<(D >= KT - 1)>());
}
static u32 *g_ord, *g_x;
static void sort0(int n) {
const u32 *X0 = XD[0];
u32 *cnt = g_x;
memset(cnt, 0, sizeof(u32) * (size_t)(n + 1));
for (int i = 0; i < n; i++) cnt[X0[i] + 1]++;
for (int i = 0; i < n; i++) cnt[i + 1] += cnt[i];
for (int i = 0; i < n; i++) g_ord[cnt[X0[i]]++] = (u32)i;
}
template <int KT>
static void run_t(int n, const u32 **x, u32 *out) {
LFSH = 0; LFD = 0;
for (int d = 0; d < KT; d++) XD[d] = x[d];
OUT = out;
ABtop = AB;
if (n <= 1) return;
FUSE_OK = 0;
if (KT >= 4 && n > 2 && (u32)n < CT_MAXV) {
const u32 *XL = XD[KT - 1];
memset(CT0, 0, (size_t)n);
int ok = 1;
for (int i = 0; i < n; i++) if (++CT0[XL[i]] > 127) { ok = 0; break; }
memset(CT0, 0, (size_t)n);
FUSE_OK = ok;
}
{
int bit = 0;
while ((n >> bit) >= LEAFNB && bit < 30) bit++;
LFSH = bit; LFD = 0;
LFB = (int)((((u32)(n - 1)) >> (LFSH + LFD)) + 2u);
if (LFB > LEAFNB) LFB = LEAFNB;
if (LFB < 8) LFB = 8;
}
sort0(n);
const u32 *X0 = XD[0], *X1 = XD[1];
const u32 *nk2 = (KT >= 3) ? XD[2] : X0;
u64 *ev = AB;
int m = 0, g = 0;
while (g < n) {
int g1 = g;
u32 v = X0[g_ord[g]];
while (g1 < n && X0[g_ord[g1]] == v) g1++;
for (int t = g; t < g1; t++) ev[m++] = mk(X1[g_ord[t]], 1, (int)g_ord[t], nk2[g_ord[t]]);
for (int t = g; t < g1; t++) ev[m++] = mk(X1[g_ord[t]], 0, (int)g_ord[t], nk2[g_ord[t]]);
g = g1;
}
ABtop = AB + m;
u64 *res = ABtop; ABtop += m;
solve_t<KT, 0, 0>(ev, m, res);
}
}
static u64 g_arena[KG_ARENA];
static u32 g_ord_st[4194308], g_x_st[4194308];
static int g_ready = 0;
static void setup() {
if (g_ready) return;
g_ready = 1;
kg::AB = g_arena;
kg::g_ord = g_ord_st;
kg::g_x = g_x_st;
}
void count_4d(int n, const unsigned *x[4], unsigned *out) {
setup();
memset(out, 0, sizeof(unsigned) * (size_t)n);
kg::run_t<4>(n, (const unsigned **)x, out);
}
#pragma GCC pop_options
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 7.532 s | 452 MB + 276 KB | Accepted | Score: 100 | 显示更多 |