【问题标题】:The indices of non-zero bytes of an SSE/AVX registerSSE/AVX 寄存器的非零字节索引
【发布时间】:2016-06-11 09:48:53
【问题描述】:

如果一个 SSE/AVX 寄存器的值是这样的,它的所有字节都是 0 或 1,有没有办法有效地获取所有非零元素的索引?

例如,如果 xmm 值为 | r0=0 | r1=1 | r2=0 | r3=1 | r4=0 | r5=1 | r6=0 |...| r14=0 | r15=1 | 结果应该类似于 (1, 3, 5, ... , 15)。结果应放在另一个 _m128i 变量或 char[16] 数组中。

如果有帮助,我们可以假设寄存器的值是所有字节都是 0 或某个恒定的非零值(不一定是 1)。

我非常想知道是否有针对该指令的指令,或者最好是 C/C++ 内在指令。在任何 SSE 或 AVX 指令集中。

编辑 1:

原来的问题不够清楚observed by @zx485 是正确的。我正在寻找任何“连续”的解决方案。

上面的示例0 1 0 1 0 1 0 1... 应该会产生以下任一结果:

  • 如果我们假设索引从 1 开始,那么 0 将是一个终止字节,结果可能是

002 004 006 008 010 012 014 016 000 000 000 000 000 000 000 000

  • 如果我们假设负字节是终止字节,结果可能是

001 003 005 007 009 011 013 015 0xFF 0xFF 0xFF 0xFF 0xFF 0xFF 0xFF 0xFF

  • 任何以连续字节形式给出的内容,我们可以将其解释为原始值中非零元素的索引

编辑 2:

确实,正如@harold@Peter Cordes 在原始帖子的 cmets 中所建议的那样,一种可能的解决方案是首先创建一个掩码(例如使用pmovmskb)并在那里检查非零索引。但这会导致循环。

【问题讨论】:

  • 您可以使用 pmovmskb 和巨大的 lut 来完成(但这不一定很快)。顺便说一句,你想在没有索引的车道上做什么?说,0xFF?
  • 你真的只想遍历有非零元素的位置吗?因为您可以使用pcmpeqb 对全零向量(如zx485 指出)来做到这一点,然后使用pmovmskb。因此,您将 0/1 向量转换为整数寄存器中的反转位图(1,其中元素为 0)。您可以遍历位图中的零。也许最容易通过反转它,并使用bsftzcnt 循环设置位。有一条 BMI1 指令可以清除最低设置位,或者您可以使用常规 2 的补码 bithacks IIRC 执行几条指令。
  • 谢谢@harold。你们俩都是对的。事实是,如果有可用的掩码,就无法避免额外的循环。我想知道是否有办法在没有循环的情况下做到这一点。我更新了我原来的帖子(见 EDIT 2 部分)。
  • @TruLa 我的建议没有循环。但是,我很好奇你打算对结果做什么,这是一个相当“烦人”的问题要解决,也许有不同的捷径?
  • @harold : BMI2 指令pext 在这里非常有用。然而,计算结果需要相当多的指令(​​没有 LUT)。请参阅下面的答案。

标签: c++ c sse simd avx


【解决方案1】:

更新的答案:新解决方案的效率略高。

您可以使用 Bit Manipulation Instruction Set 2 中的 pext 指令在没有循环的情况下执行此操作, 结合其他一些 SSE 指令。

/*
gcc -O3 -Wall -m64 -mavx2 -march=broadwell ind_nonz_avx.c
*/

#include <stdio.h>
#include <immintrin.h>
#include <stdint.h>

__m128i nonz_index(__m128i x){
   /* Set some constants that will (hopefully) be hoisted out of a loop after inlining. */
   uint64_t  indx_const   = 0xFEDCBA9876543210;                       /* 16 4-bit integers, all possible indices from 0 o 15                                                            */
   __m128i   cntr         = _mm_set_epi8(64,60,56,52,48,44,40,36,32,28,24,20,16,12,8,4);
   __m128i   pshufbcnst   = _mm_set_epi8(0x80,0x80,0x80,0x80,0x80,0x80,0x80,0x80,  0x0E,0x0C,0x0A,0x08,0x06,0x04,0x02,0x00);
   __m128i   cnst0F       = _mm_set1_epi8(0x0F);

   __m128i   msk          = _mm_cmpeq_epi8(x,_mm_setzero_si128());    /* Generate 16x8 bit mask.                                                                                        */
             msk          = _mm_srli_epi64(msk,4);                    /* Pack 16x8 bit mask to 16x4 bit mask.                                                                           */
             msk          = _mm_shuffle_epi8(msk,pshufbcnst);         /* Pack 16x8 bit mask to 16x4 bit mask, continued.                                                                */
   uint64_t  msk64        = ~ _mm_cvtsi128_si64x(msk);                 /* Move to general purpose register and invert 16x4 bit mask.                                                     */

                                                                      /* Compute the termination byte nonzmsk separately.                                                               */
   int64_t   nnz64        = _mm_popcnt_u64(msk64);                    /* Count the nonzero bits in msk64.                                                                               */
   __m128i   nnz          = _mm_set1_epi8(nnz64);                     /* May generate vmovd + vpbroadcastb if AVX2 is enabled.                                                          */
   __m128i   nonzmsk      = _mm_cmpgt_epi8(cntr,nnz);                 /* nonzmsk is a mask of the form 0xFF, 0xFF, ..., 0xFF, 0, 0, ...,0 to mark the output positions without an index */

   uint64_t  indx64       = _pext_u64(indx_const,msk64);              /* parallel bits extract. pext shuffles indx_const such that indx64 contains the nnz64 4-bit indices that we want.*/
   __m128i   indx         = _mm_cvtsi64x_si128(indx64);               /* Use a few integer instructions to unpack 4-bit integers to 8-bit integers.                                     */
   __m128i   indx_024     = indx;                                     /* Even indices.                                                                                                  */
   __m128i   indx_135     = _mm_srli_epi64(indx,4);                   /* Odd indices.                                                                                                   */
             indx         = _mm_unpacklo_epi8(indx_024,indx_135);     /* Merge odd and even indices.                                                                                    */
             indx         = _mm_and_si128(indx,cnst0F);               /* Mask out the high bits 4,5,6,7 of every byte.                                                                  */

             return _mm_or_si128(indx,nonzmsk);                       /* Merge indx with nonzmsk .                                                                                      */
}


int main(){
   int i;
   char w[16],xa[16];
   __m128i x;

   /* Example with bytes 15, 12, 7, 5, 4, 3, 2, 1, 0 set. */
   x = _mm_set_epi8(1,0,0,1,  0,0,0,0,  1,0,1,1,  1,1,1,1);   

   /* Other examples. */
   /* 
   x = _mm_set_epi8(1,1,1,1,  1,1,1,1, 1,1,1,1, 1,1,1,1);   
   x = _mm_set_epi8(0,0,0,0,  0,0,0,0, 0,0,0,0, 0,0,0,0);   
   x = _mm_set_epi8(1,0,0,0,  0,0,0,0, 0,0,0,0, 0,0,0,0);   
   x = _mm_set_epi8(0,0,0,0,  0,0,0,0, 0,0,0,0, 0,0,0,1);   
   */   
   __m128i indices = nonz_index(x);
   _mm_storeu_si128((__m128i *)w,indices);
   _mm_storeu_si128((__m128i *)xa,x);

   printf("counter 15..0 ");for (i=15;i>-1;i--) printf(" %2d ",i);      printf("\n\n");
   printf("example xmm:  ");for (i=15;i>-1;i--) printf(" %2d ",xa[i]);  printf("\n");
   printf("result in dec ");for (i=15;i>-1;i--) printf(" %2hhd ",w[i]); printf("\n");
   printf("result in hex ");for (i=15;i>-1;i--) printf(" %2hhX ",w[i]); printf("\n");

   return 0;
}

大约需要 5 条指令才能在不需要的位置获得 0xFF(终止字节)。 请注意,函数nonz_index 返回索引和仅返回终止字节的位置,实际上没有 插入终止字节,计算成本会低得多,并且可能适用于特定应用程序。 第一个终止字节的位置是nnz64&gt;&gt;2

结果是:

$ ./a.out
counter 15..0  15  14  13  12  11  10   9   8   7   6   5   4   3   2   1   0 

example xmm:    1   0   0   1   0   0   0   0   1   0   1   1   1   1   1   1 
result in dec  -1  -1  -1  -1  -1  -1  -1  15  12   7   5   4   3   2   1   0 
result in hex  FF  FF  FF  FF  FF  FF  FF   F   C   7   5   4   3   2   1   0 

英特尔 Haswell 处理器或更新版本支持 pext 指令。

【讨论】:

    【解决方案2】:

    如果您希望结果数组被“压缩”,您的问题就不清楚了。我所说的“压缩”是指结果应该是连续的。所以,例如0 1 0 1 0 1 0 1...,有两种可能:

    非连续性:

    XMM0: 000 001 000 003 000 005 000 007 000 009 000 011 000 013 000 015

    连续:

    XMM0: 001 003 005 007 009 011 013 015 000 000 000 000 000 000 000 000

    连续方法的一个问题是:如何确定它是索引0 还是终止值?

    我正在为第一种非连续方法提供一个简单的解决方案,它应该非常快:

    .data
      ddqZeroToFifteen              db 0,1,2,3,4,5,6,7,8,9,10,11,12,13,14,15
      ddqTestValue:                 db 0,1,0,1,0,1,0,1,0,1,0,1,0,1,0,1
    .code
      movdqa xmm0, xmmword ptr [ddqTestValue]
      pxor xmm1, xmm1                             ; zero XMM1
      pcmpeqb xmm0, xmm1                          ; set to -1 for all matching
      pandn xmm0, xmmword ptr [ddqZeroToFifteen]  ; invert and apply indices
    

    为了完整起见:第二种方法,即连续方法,不在此答案中。

    【讨论】:

    • 谢谢@zx485,我更新了我原来的帖子(见EDIT 1部分)。
    猜你喜欢
    • 2013-10-31
    • 1970-01-01
    • 2013-03-23
    • 1970-01-01
    • 1970-01-01
    • 2011-08-23
    • 2016-08-27
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多