【问题标题】:How can I count the occurrence of a byte in array using SIMD?如何使用 SIMD 计算数组中某个字节的出现次数?
【发布时间】:2018-09-08 05:02:33
【问题描述】:

给定以下输入字节:

var vBytes = new Vector<byte>(new byte[] {72, 101, 55, 08, 108, 111, 55, 87, 111, 114, 108, 55, 100, 55, 55, 20});

以及给定的掩码:

var mask = new Vector<byte>(55);

如何在输入数组中找到字节 55 的计数?

我已经尝试 xoring vBytesmask:

var xored = Vector.Xor(mask, vBytes);

给出:

但我不知道如何从中获得计数。

为简单起见,我们假设输入字节长度始终等于Vector&lt;byte&gt;.Count 的大小。

【问题讨论】:

  • 你的意思是没有简单的for循环?
  • 仅供参考 - Vector.Equals(vBytes,mask) 可能比 xor 更直观 - 它返回一个 255s/0s 的向量。但是,如何计算它们...
  • @MarcGravell 太棒了!我知道了!,将更新答案。
  • Vector.Dot(Vector.Negate(Vector.Equals(vBytes, new Vector&lt;byte&gt;(55))), new Vector&lt;byte&gt;(1)) 会这样做。但是,我没有使用 SIMD 的经验,我不知道这是否是一种合理的方法。
  • @MarcGravell:是的,压缩字节比较,然后使用psadbw 将这些结果水平求和为 64 位元素。

标签: c# .net simd system.numerics


【解决方案1】:

感谢 Marc Gravell 的提示,以下工作有效:

var areEqual = Vector.Equals(vBytes, mask);
var negation = Vector.Negate(areEqual);
var count = Vector.Dot(negation, Vector<byte>.One);

Marc 有一个blog post,提供有关该主题的更多信息。

【讨论】:

  • 不错; var count = Vector.Dot(-Vector.Equals(vBytes, mask), Vector&lt;byte&gt;.One); 更容易阅读,但是:喜欢它;注意:您需要真正小心“加载”SIMD 的向量;如果您不小心,您可能会因负载开销而失去所有的好处。 Span&lt;T&gt; 是加载它们的好方法 - 原始数组:通常没有那么多
  • 同意,这是一个人为的例子,目的是让它的核心工作,在生产中它将被进一步优化。不过感谢您的灯泡时刻!
【解决方案2】:

(以下想法的 AVX2 C 内在函数实现,如果有具体示例有帮助:How to count character occurrences using SIMD

在 asm 中,您希望 pcmpeqb 生成一个 0 或 0xFF 的向量。被视为有符号整数,即 0/-1。

然后使用比较结果作为整数值psubb 将 0 / 1 添加到该元素的计数器。 (减 -1 = 加 +1)

这可能会在 256 次迭代后溢出,因此在此之前的某个时间,使用 _mm_setzero_si128()_mm_setzero_si128() 将这些无符号字节(没有溢出)水平求和为 64 位整数(每组 8 个字节一个 64 位整数) .然后paddq 累加 64 位总数。

在溢出之前累积可以通过嵌套循环完成,或者在常规展开循环结束时完成。 psadbw 速度很快(因为它是视频编码运动搜索的关键构建块),因此每 4 次比较,甚至每 1 次比较累加并跳过 psubb 也不错。

有关 x86 的更多详细信息,请参阅Agner Fog's optimization guides。根据他的指令表,psadbw xmm/vpsadbw ymm 在 Skylake 上以每个时钟周期运行 1 个向量,具有 3 个周期延迟。 (前端带宽只有 1 uop。)上面提到的所有指令也是单 uop,并且运行在多个端口上(所以在吞吐量上不一定相互冲突)。他们的 128 位版本只需要 SSE2。


如果你真的一次只有一个向量可以计数,并且没有循环内存,那么可能pcmpeqb / psadbw / pshufd(将高半复制到低)/paddd / @ 987654337@ 为您提供 255 * 整数寄存器中的匹配数。一个额外的向量指令(例如从零减去,或与 1 的与,或pabsb(绝对值)将删除 x255 比例因子。


IDK 如何在 C# SIMD 中编写它,但您肯定想要一个点积!解包并转换为 FP 会比上面的慢 4 倍,这只是因为固定宽度的向量比浮点数多 4 倍,而且dpps (_mm_dp_ps) 快. Skylake 上每 1.5 个周期吞吐量 4 微指令和 1 个微指令。如果您确实必须对无符号字节以外的内容进行水平求和,请参阅Fastest way to do horizontal SSE vector sum (or other reduction)(我的答案还包括整数)。

或者如果Vector.Dotpmaddubsw / pmaddwd 用于整数向量,那么这可能没有那么糟糕,但是与psadbw 相比,对比较结果的每个向量进行多步水平求和就很糟糕了, 或者特别是字节累加器,你偶尔只会水平求和。

或者,如果 C# 优化了与 1 的常量向量的任何实际乘法。无论如何,这个答案的第一部分是您希望 CPU 运行的代码。使用任何源代码来实现它。

【讨论】:

    【解决方案3】:

    这里是 C 中的快速 SSE2 实现:

    size_t memcount_sse2(const void *s, int c, size_t n) {
       __m128i cv = _mm_set1_epi8(c), sum = _mm_setzero_si128(), acr0,acr1,acr2,acr3;
        const char *p,*pe;                                                                         
        for(p = s; p != (char *)s+(n- (n % (252*16)));) { 
          for(acr0 = acr1 = acr2 = acr3 = _mm_setzero_si128(),pe = p+252*16; p != pe; p += 64) { 
            acr0 = _mm_add_epi8(acr0, _mm_cmpeq_epi8(cv, _mm_loadu_si128((const __m128i *)p))); 
            acr1 = _mm_add_epi8(acr1, _mm_cmpeq_epi8(cv, _mm_loadu_si128((const __m128i *)(p+16)))); 
            acr2 = _mm_add_epi8(acr2, _mm_cmpeq_epi8(cv, _mm_loadu_si128((const __m128i *)(p+32)))); 
            acr3 = _mm_add_epi8(acr3, _mm_cmpeq_epi8(cv, _mm_loadu_si128((const __m128i *)(p+48))));
            __builtin_prefetch(p+1024);
          }
          sum = _mm_add_epi64(sum, _mm_sad_epu8(_mm_sub_epi8(_mm_setzero_si128(), acr0), _mm_setzero_si128()));
          sum = _mm_add_epi64(sum, _mm_sad_epu8(_mm_sub_epi8(_mm_setzero_si128(), acr1), _mm_setzero_si128()));
          sum = _mm_add_epi64(sum, _mm_sad_epu8(_mm_sub_epi8(_mm_setzero_si128(), acr2), _mm_setzero_si128()));
          sum = _mm_add_epi64(sum, _mm_sad_epu8(_mm_sub_epi8(_mm_setzero_si128(), acr3), _mm_setzero_si128()));
        }
    
        // may require SSE4, rewrite this part for actual SSE2.
        size_t count = _mm_extract_epi64(sum, 0) + _mm_extract_epi64(sum, 1);
    
        // scalar cleanup.  Could be optimized.
        while(p != (char *)s + n) count += *p++ == c;
        return count;
    }
    

    并查看:https://gist.github.com/powturbo 和 avx2 实现。

    【讨论】:

    • 对于某些编译器,_mm_extract_epi64(sum, 1) 只能使用 SSE4.1 进行编译。您可以在内循环中使用_mm_sub_epi8,以避免需要在psadbw 之前否定累加器。 acr0 -= -1acr0 += 1 相同。
    • 预取能带来多少加速?在 IvyBridge 及更高版本上,使用硬件下一页预取,它应该没有太大区别。
    • 另外,对于小的不均匀大小的缓冲区,你可以做得更好,清理循环一次去 1 个向量,然后可能是 1 movq,而不是最多 63 个一个字节- 一次迭代。或者,也许使用一直到缓冲区 end 的负载,并屏蔽掉重复计算的重叠字节。 (例如,从...,0,0,0,-1,-1,-1,-1,... 的滑动窗口加载遮罩,像这样stackoverflow.com/questions/34306933/…
    • @PeterCordes 感谢您的建议。在 i2600k 和大缓冲区上,预取速度提高了约 10%。
    • 有趣。如果我能解决它,我将在 Skylake 上进行测试。 (由于下一页预取,加速可能要低得多。)我没有 IvB 系统,但 IvB 显然对 SW 预取指令有某种主要的吞吐量瓶颈。
    【解决方案4】:

    我知道我迟到了,但到目前为止,这里的答案都没有真正提供完整的解决方案。这是我最好的尝试,来自this Gistthe DotNet source code。所有功劳都归功于 DotNet 团队和社区成员(尤其是@Peter Cordes)。

    用法:

    var bytes = Encoding.ASCII.GetBytes("The quick brown fox jumps over the lazy dog.");
    var byteCount = bytes.OccurrencesOf(32);
    
    var chars = "The quick brown fox jumps over the lazy dog.";
    var charCount = chars.OccurrencesOf(' ');
    

    代码:

    public static class VectorExtensions
    {
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        private static nuint GetByteVector128SpanLength(nuint offset, int length) =>
            ((nuint)(uint)((length - (int)offset) & ~(Vector128<byte>.Count - 1)));
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        private static nuint GetByteVector256SpanLength(nuint offset, int length) =>
            ((nuint)(uint)((length - (int)offset) & ~(Vector256<byte>.Count - 1)));
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        private static nint GetCharVector128SpanLength(nint offset, nint length) =>
            ((length - offset) & ~(Vector128<ushort>.Count - 1));
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        private static nint GetCharVector256SpanLength(nint offset, nint length) =>
            ((length - offset) & ~(Vector256<ushort>.Count - 1));
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        private static Vector128<byte> LoadVector128(ref byte start, nuint offset) =>
            Unsafe.ReadUnaligned<Vector128<byte>>(ref Unsafe.AddByteOffset(ref start, offset));
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        private static Vector256<byte> LoadVector256(ref byte start, nuint offset) =>
            Unsafe.ReadUnaligned<Vector256<byte>>(ref Unsafe.AddByteOffset(ref start, offset));
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        private static Vector128<ushort> LoadVector128(ref char start, nint offset) =>
            Unsafe.ReadUnaligned<Vector128<ushort>>(ref Unsafe.As<char, byte>(ref Unsafe.Add(ref start, offset)));
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        private static Vector256<ushort> LoadVector256(ref char start, nint offset) =>
            Unsafe.ReadUnaligned<Vector256<ushort>>(ref Unsafe.As<char, byte>(ref Unsafe.Add(ref start, offset)));
        [MethodImpl(MethodImplOptions.AggressiveOptimization)]
        private static unsafe int OccurrencesOf(ref byte searchSpace, byte value, int length) {
            var lengthToExamine = ((nuint)length);
            var offset = ((nuint)0);
            var result = 0L;
    
            if (Sse2.IsSupported || Avx2.IsSupported) {
                if (31 < length) {
                    lengthToExamine = UnalignedCountVector128(ref searchSpace);
                }
            }
    
        SequentialScan:
            while (7 < lengthToExamine) {
                ref byte current = ref Unsafe.AddByteOffset(ref searchSpace, offset);
    
                if (value == current) {
                    ++result;
                }
                if (value == Unsafe.AddByteOffset(ref current, 1)) {
                    ++result;
                }
                if (value == Unsafe.AddByteOffset(ref current, 2)) {
                    ++result;
                }
                if (value == Unsafe.AddByteOffset(ref current, 3)) {
                    ++result;
                }
                if (value == Unsafe.AddByteOffset(ref current, 4)) {
                    ++result;
                }
                if (value == Unsafe.AddByteOffset(ref current, 5)) {
                    ++result;
                }
                if (value == Unsafe.AddByteOffset(ref current, 6)) {
                    ++result;
                }
                if (value == Unsafe.AddByteOffset(ref current, 7)) {
                    ++result;
                }
    
                lengthToExamine -= 8;
                offset += 8;
            }
    
            while (3 < lengthToExamine) {
                ref byte current = ref Unsafe.AddByteOffset(ref searchSpace, offset);
    
                if (value == current) {
                    ++result;
                }
                if (value == Unsafe.AddByteOffset(ref current, 1)) {
                    ++result;
                }
                if (value == Unsafe.AddByteOffset(ref current, 2)) {
                    ++result;
                }
                if (value == Unsafe.AddByteOffset(ref current, 3)) {
                    ++result;
                }
    
                lengthToExamine -= 4;
                offset += 4;
            }
    
            while (0 < lengthToExamine) {
                if (value == Unsafe.AddByteOffset(ref searchSpace, offset)) {
                    ++result;
                }
    
                --lengthToExamine;
                ++offset;
            }
    
            if (offset < ((nuint)(uint)length)) {
                if (Avx2.IsSupported) {
                    if (0 != (((nuint)(uint)Unsafe.AsPointer(ref searchSpace) + offset) & (nuint)(Vector256<byte>.Count - 1))) {
                        var sum = Sse2.SumAbsoluteDifferences(Sse2.Subtract(Vector128<byte>.Zero, Sse2.CompareEqual(Vector128.Create(value), LoadVector128(ref searchSpace, offset))).AsByte(), Vector128<byte>.Zero).AsInt64();
    
                        offset += 16;
                        result += (sum.GetElement(0) + sum.GetElement(1));
                    }
    
                    lengthToExamine = GetByteVector256SpanLength(offset, length);
    
                    var searchMask = Vector256.Create(value);
    
                    if (127 < lengthToExamine) {
                        var sum = Vector256<long>.Zero;
    
                        do {
                            var accumulator0 = Vector256<byte>.Zero;
                            var accumulator1 = Vector256<byte>.Zero;
                            var accumulator2 = Vector256<byte>.Zero;
                            var accumulator3 = Vector256<byte>.Zero;
                            var loopIndex = ((nuint)0);
                            var loopLimit = Math.Min(255, (lengthToExamine / 128));
    
                            do {
                                accumulator0 = Avx2.Subtract(accumulator0, Avx2.CompareEqual(searchMask, LoadVector256(ref searchSpace, offset)));
                                accumulator1 = Avx2.Subtract(accumulator1, Avx2.CompareEqual(searchMask, LoadVector256(ref searchSpace, (offset + 32))));
                                accumulator2 = Avx2.Subtract(accumulator2, Avx2.CompareEqual(searchMask, LoadVector256(ref searchSpace, (offset + 64))));
                                accumulator3 = Avx2.Subtract(accumulator3, Avx2.CompareEqual(searchMask, LoadVector256(ref searchSpace, (offset + 96))));
                                loopIndex++;
                                offset += 128;
                            } while (loopIndex < loopLimit);
    
                            lengthToExamine -= (128 * loopLimit);
                            sum = Avx2.Add(sum, Avx2.SumAbsoluteDifferences(accumulator0.AsByte(), Vector256<byte>.Zero).AsInt64());
                            sum = Avx2.Add(sum, Avx2.SumAbsoluteDifferences(accumulator1.AsByte(), Vector256<byte>.Zero).AsInt64());
                            sum = Avx2.Add(sum, Avx2.SumAbsoluteDifferences(accumulator2.AsByte(), Vector256<byte>.Zero).AsInt64());
                            sum = Avx2.Add(sum, Avx2.SumAbsoluteDifferences(accumulator3.AsByte(), Vector256<byte>.Zero).AsInt64());
                        } while (127 < lengthToExamine);
    
                        var sumX = Avx2.ExtractVector128(sum, 0);
                        var sumY = Avx2.ExtractVector128(sum, 1);
                        var sumZ = Sse2.Add(sumX, sumY);
    
                        result += (sumZ.GetElement(0) + sumZ.GetElement(1));
                    }
    
                    if (31 < lengthToExamine) {
                        var sum = Vector256<long>.Zero;
    
                        do {
                            sum = Avx2.Add(sum, Avx2.SumAbsoluteDifferences(Avx2.Subtract(Vector256<byte>.Zero, Avx2.CompareEqual(searchMask, LoadVector256(ref searchSpace, offset))).AsByte(), Vector256<byte>.Zero).AsInt64());
                            lengthToExamine -= 32;
                            offset += 32;
                        } while (31 < lengthToExamine);
    
                        var sumX = Avx2.ExtractVector128(sum, 0);
                        var sumY = Avx2.ExtractVector128(sum, 1);
                        var sumZ = Sse2.Add(sumX, sumY);
    
                        result += (sumZ.GetElement(0) + sumZ.GetElement(1));
                    }
    
                    if (offset < ((nuint)(uint)length)) {
                        lengthToExamine = (((nuint)(uint)length) - offset);
    
                        goto SequentialScan;
                    }
                }
                else if (Sse2.IsSupported) {
                    lengthToExamine = GetByteVector128SpanLength(offset, length);
    
                    var searchMask = Vector128.Create(value);
    
                    if (63 < lengthToExamine) {
                        var sum = Vector128<long>.Zero;
    
                        do {
                            var accumulator0 = Vector128<byte>.Zero;
                            var accumulator1 = Vector128<byte>.Zero;
                            var accumulator2 = Vector128<byte>.Zero;
                            var accumulator3 = Vector128<byte>.Zero;
                            var loopIndex = ((nuint)0);
                            var loopLimit = Math.Min(255, (lengthToExamine / 64));
    
                            do {
                                accumulator0 = Sse2.Subtract(accumulator0, Sse2.CompareEqual(searchMask, LoadVector128(ref searchSpace, offset)));
                                accumulator1 = Sse2.Subtract(accumulator1, Sse2.CompareEqual(searchMask, LoadVector128(ref searchSpace, (offset + 16))));
                                accumulator2 = Sse2.Subtract(accumulator2, Sse2.CompareEqual(searchMask, LoadVector128(ref searchSpace, (offset + 32))));
                                accumulator3 = Sse2.Subtract(accumulator3, Sse2.CompareEqual(searchMask, LoadVector128(ref searchSpace, (offset + 48))));
                                loopIndex++;
                                offset += 64;
                            } while (loopIndex < loopLimit);
    
                            lengthToExamine -= (64 * loopLimit);
                            sum = Sse2.Add(sum, Sse2.SumAbsoluteDifferences(accumulator0.AsByte(), Vector128<byte>.Zero).AsInt64());
                            sum = Sse2.Add(sum, Sse2.SumAbsoluteDifferences(accumulator1.AsByte(), Vector128<byte>.Zero).AsInt64());
                            sum = Sse2.Add(sum, Sse2.SumAbsoluteDifferences(accumulator2.AsByte(), Vector128<byte>.Zero).AsInt64());
                            sum = Sse2.Add(sum, Sse2.SumAbsoluteDifferences(accumulator3.AsByte(), Vector128<byte>.Zero).AsInt64());
                        } while (63 < lengthToExamine);
    
                        result += (sum.GetElement(0) + sum.GetElement(1));
                    }
    
                    if (15 < lengthToExamine) {
                        var sum = Vector128<long>.Zero;
    
                        do {
                            sum = Sse2.Add(sum, Sse2.SumAbsoluteDifferences(Sse2.Subtract(Vector128<byte>.Zero, Sse2.CompareEqual(searchMask, LoadVector128(ref searchSpace, offset))).AsByte(), Vector128<byte>.Zero).AsInt64());
                            lengthToExamine -= 16;
                            offset += 16;
                        } while (15 < lengthToExamine);
    
                        result += (sum.GetElement(0) + sum.GetElement(1));
                    }
    
                    if (offset < ((nuint)(uint)length)) {
                        lengthToExamine = (((nuint)(uint)length) - offset);
    
                        goto SequentialScan;
                    }
                }
            }
    
            return ((int)result);
        }
        [MethodImpl(MethodImplOptions.AggressiveOptimization)]
        private static unsafe int OccurrencesOf(ref char searchSpace, char value, int length) {
            var lengthToExamine = ((nint)length);
            var offset = ((nint)0);
            var result = 0L;
    
            if (0 != ((int)Unsafe.AsPointer(ref searchSpace) & 1)) { }
            else if (Sse2.IsSupported || Avx2.IsSupported) {
                if (15 < length) {
                    lengthToExamine = UnalignedCountVector128(ref searchSpace);
                }
            }
    
        SequentialScan:
            while (3 < lengthToExamine) {
                ref char current = ref Unsafe.Add(ref searchSpace, offset);
    
                if (value == current) {
                    ++result;
                }
                if (value == Unsafe.Add(ref current, 1)) {
                    ++result;
                }
                if (value == Unsafe.Add(ref current, 2)) {
                    ++result;
                }
                if (value == Unsafe.Add(ref current, 3)) {
                    ++result;
                }
    
                lengthToExamine -= 4;
                offset += 4;
            }
    
            while (0 < lengthToExamine) {
                if (value == Unsafe.Add(ref searchSpace, offset)) {
                    ++result;
                }
    
                --lengthToExamine;
                ++offset;
            }
    
            if (offset < length) {
                if (Avx2.IsSupported) {
                    if (0 != (((nint)Unsafe.AsPointer(ref Unsafe.Add(ref searchSpace, offset))) & (Vector256<byte>.Count - 1))) {
                        var sum = Sse2.SumAbsoluteDifferences(Sse2.Subtract(Vector128<ushort>.Zero, Sse2.CompareEqual(Vector128.Create(value), LoadVector128(ref searchSpace, offset))).AsByte(), Vector128<byte>.Zero).AsInt64();
    
                        offset += 8;
                        result += (sum.GetElement(0) + sum.GetElement(1));
                    }
    
                    lengthToExamine = GetCharVector256SpanLength(offset, length);
    
                    var searchMask = Vector256.Create(value);
    
                    if (63 < lengthToExamine) {
                        var sum = Vector256<long>.Zero;
    
                        do {
                            var accumulator0 = Vector256<ushort>.Zero;
                            var accumulator1 = Vector256<ushort>.Zero;
                            var accumulator2 = Vector256<ushort>.Zero;
                            var accumulator3 = Vector256<ushort>.Zero;
                            var loopIndex = 0;
                            var loopLimit = Math.Min(255, (lengthToExamine / 64));
    
                            do {
                                accumulator0 = Avx2.Subtract(accumulator0, Avx2.CompareEqual(searchMask, LoadVector256(ref searchSpace, offset)));
                                accumulator1 = Avx2.Subtract(accumulator1, Avx2.CompareEqual(searchMask, LoadVector256(ref searchSpace, (offset + 16))));
                                accumulator2 = Avx2.Subtract(accumulator2, Avx2.CompareEqual(searchMask, LoadVector256(ref searchSpace, (offset + 32))));
                                accumulator3 = Avx2.Subtract(accumulator3, Avx2.CompareEqual(searchMask, LoadVector256(ref searchSpace, (offset + 48))));
                                loopIndex++;
                                offset += 64;
                            } while (loopIndex < loopLimit);
    
                            lengthToExamine -= (64 * loopLimit);
                            sum = Avx2.Add(sum, Avx2.SumAbsoluteDifferences(accumulator0.AsByte(), Vector256<byte>.Zero).AsInt64());
                            sum = Avx2.Add(sum, Avx2.SumAbsoluteDifferences(accumulator1.AsByte(), Vector256<byte>.Zero).AsInt64());
                            sum = Avx2.Add(sum, Avx2.SumAbsoluteDifferences(accumulator2.AsByte(), Vector256<byte>.Zero).AsInt64());
                            sum = Avx2.Add(sum, Avx2.SumAbsoluteDifferences(accumulator3.AsByte(), Vector256<byte>.Zero).AsInt64());
                        } while (63 < lengthToExamine);
    
                        var sumX = Avx2.ExtractVector128(sum, 0);
                        var sumY = Avx2.ExtractVector128(sum, 1);
                        var sumZ = Sse2.Add(sumX, sumY);
    
                        result += (sumZ.GetElement(0) + sumZ.GetElement(1));
                    }
    
                    if (15 < lengthToExamine) {
                        var sum = Vector256<long>.Zero;
    
                        do {
                            sum = Avx2.Add(sum, Avx2.SumAbsoluteDifferences(Avx2.Subtract(Vector256<ushort>.Zero, Avx2.CompareEqual(searchMask, LoadVector256(ref searchSpace, offset))).AsByte(), Vector256<byte>.Zero).AsInt64());
                            lengthToExamine -= 16;
                            offset += 16;
                        } while (15 < lengthToExamine);
    
                        var sumX = Avx2.ExtractVector128(sum, 0);
                        var sumY = Avx2.ExtractVector128(sum, 1);
                        var sumZ = Sse2.Add(sumX, sumY);
    
                        result += (sumZ.GetElement(0) + sumZ.GetElement(1));
                    }
    
                    if (offset < length) {
                        lengthToExamine = (length - offset);
    
                        goto SequentialScan;
                    }
                }
                else if (Sse2.IsSupported) {
                    lengthToExamine = GetCharVector128SpanLength(offset, length);
    
                    var searchMask = Vector128.Create(value);
    
                    if (31 < lengthToExamine) {
                        var sum = Vector128<long>.Zero;
    
                        do {
                            var accumulator0 = Vector128<ushort>.Zero;
                            var accumulator1 = Vector128<ushort>.Zero;
                            var accumulator2 = Vector128<ushort>.Zero;
                            var accumulator3 = Vector128<ushort>.Zero;
                            var loopIndex = 0;
                            var loopLimit = Math.Min(255, (lengthToExamine / 32));
    
                            do {
                                accumulator0 = Sse2.Subtract(accumulator0, Sse2.CompareEqual(searchMask, LoadVector128(ref searchSpace, offset)));
                                accumulator1 = Sse2.Subtract(accumulator1, Sse2.CompareEqual(searchMask, LoadVector128(ref searchSpace, (offset + 8))));
                                accumulator2 = Sse2.Subtract(accumulator2, Sse2.CompareEqual(searchMask, LoadVector128(ref searchSpace, (offset + 16))));
                                accumulator3 = Sse2.Subtract(accumulator3, Sse2.CompareEqual(searchMask, LoadVector128(ref searchSpace, (offset + 24))));
                                loopIndex++;
                                offset += 32;
                            } while (loopIndex < loopLimit);
    
                            lengthToExamine -= (32 * loopLimit);
                            sum = Sse2.Add(sum, Sse2.SumAbsoluteDifferences(accumulator0.AsByte(), Vector128<byte>.Zero).AsInt64());
                            sum = Sse2.Add(sum, Sse2.SumAbsoluteDifferences(accumulator1.AsByte(), Vector128<byte>.Zero).AsInt64());
                            sum = Sse2.Add(sum, Sse2.SumAbsoluteDifferences(accumulator2.AsByte(), Vector128<byte>.Zero).AsInt64());
                            sum = Sse2.Add(sum, Sse2.SumAbsoluteDifferences(accumulator3.AsByte(), Vector128<byte>.Zero).AsInt64());
                        } while (31 < lengthToExamine);
    
                        result += (sum.GetElement(0) + sum.GetElement(1));
                    }
    
                    if (7 < lengthToExamine) {
                        var sum = Vector128<long>.Zero;
    
                        do {
                            sum = Sse2.Add(sum, Sse2.SumAbsoluteDifferences(Sse2.Subtract(Vector128<ushort>.Zero, Sse2.CompareEqual(searchMask, LoadVector128(ref searchSpace, offset))).AsByte(), Vector128<byte>.Zero).AsInt64());
                            lengthToExamine -= 8;
                            offset += 8;
                        } while (7 < lengthToExamine);
    
                        result += (sum.GetElement(0) + sum.GetElement(1));
                    }
    
                    if (offset < length) {
                        lengthToExamine = (length - offset);
    
                        goto SequentialScan;
                    }
                }
            }
    
            return ((int)result);
        }
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        private static unsafe nuint UnalignedCountVector128(ref byte searchSpace) {
            nint unaligned = ((nint)Unsafe.AsPointer(ref searchSpace) & (Vector128<byte>.Count - 1));
    
            return ((nuint)(uint)((Vector128<byte>.Count - unaligned) & (Vector128<byte>.Count - 1)));
        }
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        private static unsafe nint UnalignedCountVector128(ref char searchSpace) {
            const int ElementsPerByte = (sizeof(ushort) / sizeof(byte));
    
            return ((nint)(uint)(-(int)Unsafe.AsPointer(ref searchSpace) / ElementsPerByte) & (Vector128<ushort>.Count - 1));
        }
    
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        public static int OccurrencesOf(this ReadOnlySpan<byte> span, byte value) =>
            OccurrencesOf(
                length: span.Length,
                searchSpace: ref MemoryMarshal.GetReference(span),
                value: value
            );
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        public static int OccurrencesOf(this Span<byte> span, byte value) =>
            ((ReadOnlySpan<byte>)span).OccurrencesOf(value);
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        public static int OccurrencesOf(this ReadOnlySpan<char> span, char value) =>
            OccurrencesOf(
                length: span.Length,
                searchSpace: ref MemoryMarshal.GetReference(span),
                value: value
            );
        [MethodImpl(MethodImplOptions.AggressiveInlining)]
        public static int OccurrencesOf(this Span<char> span, char value) =>
            ((ReadOnlySpan<char>)span).OccurrencesOf(value);
    }
    

    【讨论】:

    • 您可以避免在 PSADBW 之前进行 Avx2.Subtract(Vector256&lt;ushort&gt;.Zero, accumulator0) 清理,方法是首先使用 sub 而不是 add。即accumulator -= cmp(),因为 cmp 结果为 -1 或 0。另外,请执行 left + right 然后减少它,而不是单独提取所有 4 个元素。
    • 干杯。感谢您编写实际的 C# 实现;除了在 SO 上看到它之外,我真的不知道 C#,所以我不打算尝试这个。 How to count character occurrences using SIMD 使用嵌套循环来处理溢出,可能想看看它是如何实现的。
    • 另外,在外循环中,你不需要减少到标量,只需一个 SIMD 向量 var sum 就可以了。标量整数的 hsum 可以从外循环中消失。 (确保您使用 64 位或至少 32 位元素大小的 SIMD 加法来累积 psadbw 结果。我猜 SumAbsoluteDifferences() 返回 vector&lt;uint64_t&gt; 或任何 C# 调用它,这意味着元素类型)
    • 另外,请考虑输入 31 个字节长(或 63x uint16,因为您正在处理字符串而不是问题所问的字节)会发生什么。或 n*64 + 31。这是很多标量迭代。这是展开的缺点:除非您还提供未展开的向量循环,否则您会使最坏的情况(包括小情况)在缓慢的标量代码中花费更多时间。如果您想针对中短字符串进行调整,您可以提供一个循环,每次迭代执行一个 SSE2 向量,最多留下 7 个剩余元素。
    • 糟糕,将min(..., 255) 写入到 255 作为内部迭代次数的上限,但如果您接近缓冲区的末尾,请减少执行次数。我经常将minmax 混为一谈,以便在我不停下来思考的情况下设置某个值的最大值。 ://
    猜你喜欢
    • 2023-04-04
    • 2013-12-25
    • 2013-02-24
    • 1970-01-01
    • 2017-10-09
    • 2010-11-12
    • 1970-01-01
    相关资源
    最近更新 更多