【问题标题】:How do I write a branchless std::vector scan?如何编写无分支 std::vector 扫描?
【发布时间】:2016-12-12 10:28:17
【问题描述】:

我想对数组进行简单的扫描。我有一个std::vector<int> data,我想找到元素小于 9 的所有数组索引并将它们添加到结果向量中。我可以用一个分支来写这个:

for (int i = 0; i < data.size(); ++i)
    if (data[i] < 9)
        r.push_back(i);

这给出了正确答案,但我想将其与无分支版本进行比较。

使用原始数组 - 并假设 data 是一个 int 数组,length 是其中元素的数量,r 是一个有足够空间的结果数组 - 我可以这样写:

int current_write_point = 0;
for (int i = 0; i < length; ++i){
    r[current_write_point] = i;
    current_write_point += (data[i] < 9);
}

如何使用data 的向量获得类​​似的行为?

【问题讨论】:

  • data[i] &lt; 9 通常是汇编级别的分支(尽管与 push_back 相比,它肯定是一些 cmov 魔术的更好候选者,但肯定不是)
  • 第二个解决方案为什么比第一个更好?
  • 我希望current_write_point += 行生成与if (data[i] &lt; 9) { current_write_point++; } 相同的代码
  • @DimChtz 我想这就是他想要找出的——他想比较两种方法生成的代码。
  • 您是否在分割点对std::partition()std::copy() 进行了分析?

标签: c++ arrays vector conditional-statements


【解决方案1】:

让我们看看实际的compiler output

auto scan_branch(const std::vector<int>& v)
{
  std::vector<int> res;
  int insert_index = 0;
  for(int i = 0; i < v.size(); ++i)
  {
    if (v[i] < 9)
    {
       res.push_back(i);
    } 
  }
  return res;
}

这段代码显然在disassembly 的第26 行有一个分支。如果它大于或等于 9,它只会继续下一个元素,但是如果小于 9,则会为 push_back 执行一些可怕的代码,然后我们继续。没有什么意外。

auto scan_nobranch(const std::vector<int>& v)
{
  std::vector<int> res;
  res.resize(v.size());

  int insert_index = 0;
  for(int i = 0; i < v.size(); ++i)
  {
    res[insert_index] = i;
    insert_index += v[i] < 9;
  }

  res.resize(insert_index);
  return res;
}

然而,这个只有一个条件移动,你可以在disassembly 的第 190 行看到。看起来我们有一个赢家。由于条件移动不会导致流水线停顿,因此在这个中没有分支(for 条件检查除外)。

【讨论】:

  • 您能否将反汇编发布在答案本身中?谢谢:)
  • @Rakete1111,当然可以,但是godbolt颜色与c++代码和汇编代码的实际行匹配,在我看来这更容易理解。
  • 您可以离开链接,但如果将来某个时候链接失效,则答案不完整:)
  • 你的编译器设置是什么?循环是否展开? v.size() 被常量替换了吗?
  • 现在看看结果……第二个版本快了多少?
【解决方案2】:
std::copy_if(std::begin(data), std::end(data), std::back_inserter(r));

【讨论】:

  • 虽然这段代码可能有助于解决问题,但它并没有解释为什么和/或如何回答问题。提供这种额外的背景将显着提高其长期价值。请edit您的答案添加解释,包括适用的限制和假设。
【解决方案3】:

好吧,您可以事先调整向量的大小并保留您的算法:

// Resize the vector so you can index it normally
r.resize(length);

// Do your algorithm like before
int current_write_point = 0;
for (int i = 0; i < length; ++i){
    r[current_write_point] = i;
    current_write_point += (data[i] < 9);
}

// Afterwards, current_write_point can be used to shrink the vector, so
// there are no excess elements not written to
r.resize(current_write_point + 1);

如果您不希望比较,您可以使用一些按位和带有短路的布尔运算来确定。

首先,我们知道所有负整数都小于 9。其次,如果是正数,我们可以使用位掩码来判断一个整数是否在 0-15 范围内(实际上,我们会检查如果不在该范围内,则大于 15)。那么,我们知道如果从那个数减去 8 的结果是负数,那么结果小于 9: 其实,我只是想出了一个更好的方法。由于我们可以很容易地确定是否为x &lt; 0,因此我们只需将x 减去9 即可确定是否为x &lt; 9

#include <iostream>

// Use bitwise operations to determine if x is negative
int n(int x) {
    return x & (1 << 31);
}

int main() {
    int current_write_point = 0;
    for (int i = 0; i < length; ++i){
        r[current_write_point] = i;
        current_write_point += n(data[i] - 9);
    }
}

【讨论】:

  • 之后别忘了用current_write_point缩小vector
  • 还有 2 个比较:(data[i] &lt; 9)i &lt; length
  • @ThomasMatthews 好吧,我只是为 OP 提供了一种在他的算法中使用向量的方法,因为他/她想知道如何做到这一点。据我了解,OP希望将这种代码与带有if的代码进行比较
  • @ThomasMatthews 比较不是分支。
  • @NicolBolas,但据我所知,他的第二个 sn-p 不包含任何分支,因此他似乎知道这一点?连标题都这么说……
猜你喜欢
  • 2012-02-07
  • 1970-01-01
  • 2022-12-07
  • 2011-07-15
  • 2019-11-15
  • 2021-04-03
  • 2021-09-27
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多