// This code is AI-generated. (AI 生成的代码)
// WC2017 挑战-任务2 (rock-paper-scissors range queries).
//
// Map s1[i] -> "120"[s1[i]-'0'], i.e. the hand that s1[i] beats. Then a win at
// offset i is exactly s1[x+i] == s2[y+i]. Count equal bytes with AVX2:
// vpcmpeqb gives 0xFF for equal, and subtracting it from an 8-bit accumulator
// adds 1 per equal byte. Every 255 blocks the 8-bit accumulators are widened
// into 16-bit ones so they never wrap; a scalar tail finishes the query.
#include <immintrin.h>
#include <string.h>
#pragma GCC optimize("O3,unroll-loops")
#pragma GCC target("avx2")
void solve(int n, int q, char *s1, char *s2, int *q_x, int *q_y, int *q_len, unsigned *ans) {
for (int i = 0; i < n; ++i) s1[i] = "120"[s1[i] - '0'];
const __m256i zero = _mm256_setzero_si256();
for (int k = 0; k < q; ++k) {
const char *p = s1 + q_x[k];
const char *r = s2 + q_y[k];
int len = q_len[k];
int i = 0;
__m256i s0 = zero, s1a = zero, s2a = zero, s3a = zero; // 16-bit lanes
while (i + 128 <= len) {
__m256i accA = zero, accB = zero, accC = zero, accD = zero;
int chunk = (len - i) / 128;
if (chunk > 255) chunk = 255;
for (int t = 0; t < chunk; ++t) {
__m256i a0 = _mm256_loadu_si256((const __m256i *)(p + i));
__m256i b0 = _mm256_loadu_si256((const __m256i *)(r + i));
__m256i a1 = _mm256_loadu_si256((const __m256i *)(p + i + 32));
__m256i b1 = _mm256_loadu_si256((const __m256i *)(r + i + 32));
__m256i a2 = _mm256_loadu_si256((const __m256i *)(p + i + 64));
__m256i b2 = _mm256_loadu_si256((const __m256i *)(r + i + 64));
__m256i a3 = _mm256_loadu_si256((const __m256i *)(p + i + 96));
__m256i b3 = _mm256_loadu_si256((const __m256i *)(r + i + 96));
accA = _mm256_sub_epi8(accA, _mm256_cmpeq_epi8(a0, b0));
accB = _mm256_sub_epi8(accB, _mm256_cmpeq_epi8(a1, b1));
accC = _mm256_sub_epi8(accC, _mm256_cmpeq_epi8(a2, b2));
accD = _mm256_sub_epi8(accD, _mm256_cmpeq_epi8(a3, b3));
i += 128;
}
s0 = _mm256_add_epi16(s0, _mm256_cvtepu8_epi16(_mm256_castsi256_si128(accA)));
s1a = _mm256_add_epi16(s1a, _mm256_cvtepu8_epi16(_mm256_extracti128_si256(accA, 1)));
s2a = _mm256_add_epi16(s2a, _mm256_cvtepu8_epi16(_mm256_castsi256_si128(accB)));
s3a = _mm256_add_epi16(s3a, _mm256_cvtepu8_epi16(_mm256_extracti128_si256(accB, 1)));
s0 = _mm256_add_epi16(s0, _mm256_cvtepu8_epi16(_mm256_castsi256_si128(accC)));
s1a = _mm256_add_epi16(s1a, _mm256_cvtepu8_epi16(_mm256_extracti128_si256(accC, 1)));
s2a = _mm256_add_epi16(s2a, _mm256_cvtepu8_epi16(_mm256_castsi256_si128(accD)));
s3a = _mm256_add_epi16(s3a, _mm256_cvtepu8_epi16(_mm256_extracti128_si256(accD, 1)));
}
unsigned count = 0;
unsigned short buf[64];
_mm256_storeu_si256((__m256i *)(buf + 0), s0);
_mm256_storeu_si256((__m256i *)(buf + 16), s1a);
_mm256_storeu_si256((__m256i *)(buf + 32), s2a);
_mm256_storeu_si256((__m256i *)(buf + 48), s3a);
for (int t = 0; t < 64; ++t) count += buf[t];
for (; i < len; ++i) count += (p[i] == r[i]);
ans[k] = count;
}
}
#ifdef LOCAL_TEST
#include <cstdio>
int main() {
int n, q;
if (scanf("%d %d", &n, &q) != 2) return 0;
static char s1[1000005], s2[1000005];
scanf("%s %s", s1, s2);
static int qx[1000005], qy[1000005], ql[1000005];
static unsigned ans[1000005];
for (int i = 0; i < q; ++i) scanf("%d %d %d", &qx[i], &qy[i], &ql[i]);
solve(n, q, s1, s2, qx, qy, ql, ans);
for (int i = 0; i < q; ++i) printf("%u\n", ans[i]);
return 0;
}
#endif
| Compilation | N/A | N/A | Compile OK | Score: N/A | 显示更多 |
| Testcase #1 | 186.1 us | 40 KB | Accepted | Score: 50 | 显示更多 |
| Testcase #2 | 874.206 ms | 5 MB + 176 KB | Accepted | Score: 50 | 显示更多 |