这是How to count character occurrences using SIMD 和c=0 的一个特例,要计算匹配的字符(字节)值。请参阅问答,了解 char_count (char const* vector, size_t size, char c); 的经过良好优化的手动向量化 AVX2 实现,其内部循环比这更紧密,避免将每个 0/-1 匹配向量分别减少为标量。
这将是 O(n) 所以你能做的最好的就是减少常数。一种快速解决方法是删除分支。如果零是随机分布的,这将给出与我下面的 SSE 版本一样快的结果。这可能是由于 GCC 对这个循环进行了矢量化。但是,对于零的长时间运行或零的随机密度小于 1%,以下 SSE 版本仍然更快。
int countZeroBytes_fix(char* values, int length) {
int zeroCount = 0;
for(int i=0; i<length; i++) {
zeroCount += values[i] == 0;
}
return zeroCount;
}
我最初认为零的密度很重要。事实证明并非如此,至少在 SSE 中是这样。与密度无关,使用 SSE 的速度要快得多。
编辑:实际上,它确实取决于密度,它只是零的密度必须小于我的预期。 1/64 个零(1.5% 的零)是 1/4 中的一个零SSE 注册,因此分支预测不能很好地工作。但是,1/1024 个零(0.1% 的零)更快(参见时间表)。
如果数据有长时间的零运行,SIMD 会更快。
您可以将 16 个字节打包到 SSE 寄存器中。然后,您可以使用_mm_cmpeq_epi8 一次将所有 16 个字节与零进行比较。然后要处理零运行,您可以在结果上使用_mm_movemask_epi8,大多数情况下它将为零。在这种情况下,您可以获得高达 16 的加速(对于前半部分 1 和后半部分零,我获得了超过 12 倍的加速)。
这是 2^16 字节(重复 10000 次)以秒为单位的时间表。
1.5% zeros 50% zeros 0.1% zeros 1st half 1, 2nd half 0
countZeroBytes 0.8s 0.8s 0.8s 0.95s
countZeroBytes_fix 0.16s 0.16s 0.16s 0.16s
countZeroBytes_SSE 0.2s 0.15s 0.10s 0.07s
您可以在http://coliru.stacked-crooked.com/a/67a169ddb03d907a查看最后 1/2 个零的结果
#include <stdio.h>
#include <stdlib.h>
#include <emmintrin.h> // SSE2
#include <omp.h>
int countZeroBytes(char* values, int length) {
int zeroCount = 0;
for(int i=0; i<length; i++) {
if (!values[i])
++zeroCount;
}
return zeroCount;
}
int countZeroBytes_SSE(char* values, int length) {
int zeroCount = 0;
__m128i zero16 = _mm_set1_epi8(0);
__m128i and16 = _mm_set1_epi8(1);
for(int i=0; i<length; i+=16) {
__m128i values16 = _mm_loadu_si128((__m128i*)&values[i]);
__m128i cmp = _mm_cmpeq_epi8(values16, zero16);
int mask = _mm_movemask_epi8(cmp);
if(mask) {
if(mask == 0xffff) zeroCount += 16;
else {
cmp = _mm_and_si128(and16, cmp); //change -1 values to 1
//hortiontal sum of 16 bytes
__m128i sum1 = _mm_sad_epu8(cmp,zero16);
__m128i sum2 = _mm_shuffle_epi32(sum1,2);
__m128i sum3 = _mm_add_epi16(sum1,sum2);
zeroCount += _mm_cvtsi128_si32(sum3);
}
}
}
return zeroCount;
}
int main() {
const int n = 1<<16;
const int repeat = 10000;
char *values = (char*)_mm_malloc(n, 16);
for(int i=0; i<n; i++) values[i] = rand()%64; //1.5% zeros
//for(int i=0; i<n/2; i++) values[i] = 1;
//for(int i=n/2; i<n; i++) values[i] = 0;
int zeroCount = 0;
double dtime;
dtime = omp_get_wtime();
for(int i=0; i<repeat; i++) zeroCount = countZeroBytes(values,n);
dtime = omp_get_wtime() - dtime;
printf("zeroCount %d, time %f\n", zeroCount, dtime);
dtime = omp_get_wtime();
for(int i=0; i<repeat; i++) zeroCount = countZeroBytes_SSE(values,n);
dtime = omp_get_wtime() - dtime;
printf("zeroCount %d, time %f\n", zeroCount, dtime);
}