【问题标题】:Fast vectorized function to check if a value is in an interval用于检查值是否在区间内的快速矢量化函数
【发布时间】:2012-12-23 00:54:29
【问题描述】:

R 中是否有一个函数可以有效地检查一个值是否大于一个并且小于另一个数字?它也应该适用于向量。

本质上,我正在寻找以下函数的更快版本:

> in.interval <- function(x, lo, hi) (x > lo & x < hi)
> in.interval(c(2,4,6), 3, 5)
[1] FALSE  TRUE FALSE

这里的问题是x 必须被触摸两次,与更有效的方法相比,计算消耗了两倍的内存。在内部,我会假设它是这样工作的:

  1. 计算tmp1 &lt;- (x &gt; lo)
  2. 计算tmp2 &lt;- (x &lt; hi)
  3. 计算retval &lt;- tmp1 &amp; tmp2

现在,在第 2 步之后,内存中有两个布尔向量,x 必须查看两次。我的问题是:是否有一个(内置?)函数可以一步完成所有这些操作,而无需分配额外的内存?

跟进这个问题:R: Select values from data table in range

编辑:我已经根据https://gist.github.com/4344844https://gist.github.com/4344844的 CauchyDistributedRV 的回答设置了一个要点

【问题讨论】:

  • 对于 1e8 值,该函数在我的计算机上需要大约 12 秒。你希望它快多少?您将如何通过仅访问一次 x 来检查 2 个条件?您能否向我们指出您想到的“更有效的方法”?
  • @JorisMeys:大约 6 秒就好了 :-) 稍后会编辑问题。
  • 也许findInterval 是您问题的矢量化版本?
  • @JorisMeys abs(x-(hi+lo)/2)-(hi-lo)/2 &lt; 0
  • @James Thx。我已经想通了,但希望 OP 会努力使用填充他/她头部的灰质 :)。当你把它放在评论中时,如果你愿意,你也可以给出答案。

标签: performance r vectorization


【解决方案1】:

正如@James 在 cmets 中所说,诀窍是从 x 中减去 low 和 high 之间的中间值,然后检查该差是否小于 low 和 high 之间距离的一半。或者,在代码中:

in.interval2 <- function(x, lo, hi) {
    abs(x-(hi+lo)/2) < (hi-lo)/2 
}

这与.bincode hack 的速度差不多,并且是您正在寻找的算法的实现。你可以把它翻译成 C 或 C++ 并尝试如果你得到加速。

与其他解决方案的比较:

x <- runif(1e6,1,10)
require(rbenchmark)
benchmark(
  in.interval(x, 3, 5),
  in.interval2(x, 3, 5),
  findInterval(x, c(3, 5)) == 1,
  !is.na(.bincode(x, c(3, 5))),
  order='relative',
  columns=c("test", "replications", "elapsed", "relative")
) 

给予

                           test replications elapsed relative
4  !is.na(.bincode(x, c(3, 5)))          100    1.88    1.000
2         in.interval2(x, 3, 5)          100    1.95    1.037
3 findInterval(x, c(3, 5)) == 1          100    3.42    1.819
1          in.interval(x, 3, 5)          100    3.54    1.883

【讨论】:

  • 这个想法很不错,但在我的机器上它比.bincode 慢得多,而且Rcpp 版本的行为就像其他最好的Rcpp 版本一样,它在内部使用&amp;。有关结果,请参见 Gist(测试 7 和 8)。
  • (x-lo)*(hi-x) &gt; 0 呢?
  • @Roland:(x-lo)*(x-hi) &lt;= 0 可能会更好。但我想我会选择.bincode,除非有人想出更快的选择。
  • 虽然我在这里:您和 Roland 的方法似乎都不允许测试左包含右排除(或相反)。 .bincode 可以做任何事情,除了左独右独占(在 include.lowest 参数的帮助下)。
  • 确实如此。但是left-inclusive right-exclusive不能一次性测试,因为你不能将条件重写为一个条件。为此,您需要以相同的方式对待双方(包括或排除)
【解决方案2】:

findInterval 比 in.interval 长 x 快。

library(microbenchmark)

set.seed(123L)
x <- runif(1e6, 1, 10)
in.interval <- function(x, lo, hi) (x > lo & x < hi)

microbenchmark(
    findInterval(x, c(3, 5)) == 1L,
    in.interval(x, 3, 5),
    times=100)

与

Unit: milliseconds
                            expr      min       lq   median       uq      max
1 findInterval(x, c(3, 5)) == 1L 23.40665 25.13308 25.17272 25.25361 27.04032
2           in.interval(x, 3, 5) 42.91647 45.51040 45.60424 45.75144 46.38389

如果不需要== 1L,则速度更快,如果要找到的“间隔”大于 1,则很有用

> system.time(findInterval(x, 0:10))
   user  system elapsed 
  3.644   0.112   3.763 

如果速度至关重要,那么这个 C 实现虽然不能容忍整数而不是数字参数,但速度很快

library(inline)
in.interval_c <- cfunction(c(x="numeric", lo="numeric", hi="numeric"),
'    int len = Rf_length(x);
     double lower = REAL(lo)[0], upper = REAL(hi)[0],
            *xp = REAL(x);
     SEXP out = PROTECT(NEW_LOGICAL(len));
     int *outp = LOGICAL(out);

     for (int i = 0; i < len; ++i)
         outp[i] = (xp[i] - lower) * (xp[i] - upper) <= 0;

     UNPROTECT(1);
     return out;')

其他答案中提出的一些解决方案的时间安排是

microbenchmark(
    findInterval(x, c(3, 5)) == 1L,
    in.interval.abs(x, 3, 5),
    in.interval(x, 3, 5),
    in.interval_c(x, 3, 5),
    !is.na(.bincode(x, c(3, 5))),
    times=100)

与

Unit: milliseconds
                            expr       min        lq    median        uq
1 findInterval(x, c(3, 5)) == 1L 23.419117 23.495943 23.556524 23.670907
2       in.interval.abs(x, 3, 5) 12.018486 12.056290 12.093279 12.161213
3         in.interval_c(x, 3, 5)  1.619649  1.641119  1.651007  1.679531
4           in.interval(x, 3, 5) 42.946318 43.050058 43.171480 43.407930
5   !is.na(.bincode(x, c(3, 5))) 15.421340 15.468946 15.520298 15.600758
        max
1 26.360845
2 13.178126
3  2.785939
4 46.187129
5 18.558425

在 bin.cpp 文件中重新审视速度问题

#include <Rcpp.h>

using namespace Rcpp;

// [[Rcpp::export]]
SEXP bin1(SEXP x, SEXP lo, SEXP hi)
{
    const int len = Rf_length(x);
    const double lower = REAL(lo)[0], upper = REAL(hi)[0];
    SEXP out = PROTECT(Rf_allocVector(LGLSXP, len));

    double *xp = REAL(x);
    int *outp = LOGICAL(out);
    for (int i = 0; i < len; ++i)
    outp[i] = (xp[i] - lower) * (xp[i] - upper) <= 0;

    UNPROTECT(1);
    return out;
}

// [[Rcpp::export]]
LogicalVector bin2(NumericVector x, NumericVector lo, NumericVector hi)
{
    NumericVector xx(x);
    double lower = as<double>(lo);
    double upper = as<double>(hi); 

    LogicalVector out(x);
    for( int i=0; i < out.size(); i++ )
        out[i] = ( (xx[i]-lower) * (xx[i]-upper) ) <= 0;

    return out;
}

// [[Rcpp::export]]
LogicalVector bin3(NumericVector x, const double lower, const double upper)
{
    const int len = x.size();
    LogicalVector out(len);

    for (int i=0; i < len; i++)
        out[i] = ( (x[i]-lower) * (x[i]-upper) ) <= 0;

    return out;
}

有时间

> library(Rcpp)
> sourceCpp("bin.cpp")
> microbenchmark(bin1(x, 3, 5), bin2(x, 3, 5), bin3(x, 3, 5),                   
+                in.interval_c(x, 3, 5), times=1000)                            
Unit: milliseconds                                                              
                    expr       min        lq    median        uq      max       
1          bin1(x, 3, 5)  1.546703  2.668171  2.785255  2.839225 144.9574       
2          bin2(x, 3, 5) 12.547456 13.583808 13.674477 13.792773 155.6594       
3          bin3(x, 3, 5)  2.238139  3.318293  3.357271  3.540876 144.1249       
4 in.interval_c(x, 3, 5)  1.545139  2.654809  2.767784  2.822722 143.7500       

通过使用常量len 而不是out.size() 作为循环边界并分配逻辑向量而不初始化它(LogicalVector(len),因为它将在循环中初始化)。

【讨论】:

  • 我已将您的解决方案嵌入到gist.github.com/4344844 的要点中。对于 1e6 元素,它比 &amp; 方法运行得更快,但是使用 C++ 仍然比它快两倍。
  • Rcpp中大对象的复制
  • 现在让我感到莫名其妙。您提供的 C 解决方案实际上比我系统上的简单 x &lt; hi 运行得更快(大约 2 倍)(尝试在基准测试中添加 x &gt; lo &amp; x &lt; hi、x &lt; hi 以查看)。 that 是如何发生的——我认为 R 中运算符的底层 C 实现已经非常优化了?或者与我编译那个 C 函数时可能发生的任何事情相比,R 的二进制版本是否以“安全”的方式编译?
  • @CauchyDistributedRV x &lt; hi 需要分配与我的代码相同的内存量(用于返回逻辑),这两个函数都需要遍历所有值,并且 C 编译器可能已经优化了我的for 循环比高级语法所暗示的操作要少得多,因此两个循环的基本成本可能是可比的。 R 还会做很多我们认为理所当然的事情,例如,处理 NA、回收 hi(一般来说,不仅仅是长度为 1 的特殊情况)、检查数据类型之间是否需要强制转换等。
  • @user946850 我稍微研究了一下速度差异,并在我的答案中添加了一个部分。
【解决方案3】:

如果你可以处理NAs,你可以使用.bincode:

.bincode(c(2,4,6), c(3, 5))
[1] NA  1 NA

library(microbenchmark)
set.seed(42)
x = runif(1e8, 1, 10)
microbenchmark(in.interval(x, 3, 5),
               findInterval(x,  c(3, 5)),
               .bincode(x, c(3, 5)),
               times=5)

Unit: milliseconds
                      expr       min        lq    median       uq      max
1     .bincode(x, c(3, 5))  930.4842  934.3594  955.9276 1002.857 1047.348
2 findInterval(x, c(3, 5)) 1438.4620 1445.7131 1472.4287 1481.380 1551.419
3     in.interval(x, 3, 5) 2977.8460 3046.7720 3075.8381 3182.013 3288.020

【讨论】:

  • 内部函数的甜蜜使用。您可以通过!is.na(.bincode(...))得到答案
  • .bincode 将其参数转换为整数,因此有令人惊讶的(在当前上下文中)结果——.bincode(3.1, 3, 5) 是 'NA';测试每种方法的结果的同一性。
  • 糟糕,抱歉。
  • 适用于公寓,只是比 Rcpp-ed 解决方案稍慢。结果在 Gist 中。
【解决方案4】:

我能找到的主要加速是通过对函数进行字节编译。即使是 Rcpp 解决方案(尽管使用 Rcpp 糖,而不是更深入的 C 解决方案)也比编译后的解决方案慢。

library( compiler )
library( microbenchmark )
library( inline )

in.interval <- function(x, lo, hi) (x > lo & x < hi)
in.interval2 <- cmpfun( in.interval )
in.interval3 <- function(x, lo, hi) {
  sapply( x, function(xx) { 
    xx > lo && xx < hi }
          )
}
in.interval4 <- cmpfun( in.interval3 )
in.interval5 <- rcpp( signature(x="numeric", lo="numeric", hi="numeric"), '
NumericVector xx(x);
double lower = Rcpp::as<double>(lo);
double upper = Rcpp::as<double>(hi);

return Rcpp::wrap( xx > lower & xx < upper );
')

x <- c(2, 4, 6)
lo <- 3
hi <- 5

microbenchmark(
  in.interval(x, lo, hi),
  in.interval2(x, lo, hi),
  in.interval3(x, lo, hi),
  in.interval4(x, lo, hi),
  in.interval5(x, lo, hi)
)

给我

Unit: microseconds
                     expr    min      lq  median      uq    max
1  in.interval(x, lo, hi)  1.575  2.0785  2.5025  2.6560  7.490
2 in.interval2(x, lo, hi)  1.035  1.4230  1.6800  2.0705 11.246
3 in.interval3(x, lo, hi) 25.439 26.2320 26.7350 27.2250 77.541
4 in.interval4(x, lo, hi) 22.479 23.3920 23.8395 24.3725 33.770
5 in.interval5(x, lo, hi)  1.425  1.8740  2.2980  2.5565 21.598


编辑:在其他 cmets 之后,这是一个更快的 Rcpp 解决方案,使用给出绝对值的技巧:
library( compiler )
library( inline )
library( microbenchmark )

in.interval.oldRcpp <- rcpp( 
  signature(x="numeric", lo="numeric", hi="numeric"), '
    NumericVector xx(x);
    double lower = Rcpp::as<double>(lo);
    double upper = Rcpp::as<double>(hi);

    return Rcpp::wrap( (xx > lower) & (xx < upper) );
    ')

in.interval.abs <- rcpp( 
  signature(x="numeric", lo="numeric", hi="numeric"), '
    NumericVector xx(x);
    double lower = as<double>(lo);
    double upper = as<double>(hi); 

    LogicalVector out(x);
    for( int i=0; i < out.size(); i++ ) {
      out[i] = ( (xx[i]-lower) * (xx[i]-upper) ) <= 0;
    }
    return wrap(out);
    ')

in.interval.abs.sugar <- rcpp( 
  signature( x="numeric", lo="numeric", hi="numeric"), '
    NumericVector xx(x);
    double lower = as<double>(lo);
    double upper = as<double>(hi); 

    return wrap( ((xx-lower) * (xx-upper)) <= 0 );
    ')

x <- runif(1E5)
lo <- 0.5
hi <- 1

microbenchmark(
  in.interval.oldRcpp(x, lo, hi),
  in.interval.abs(x, lo, hi),
  in.interval.abs.sugar(x, lo, hi)
)

all.equal( in.interval.oldRcpp(x, lo, hi), in.interval.abs(x, lo, hi) )
all.equal( in.interval.oldRcpp(x, lo, hi), in.interval.abs.sugar(x, lo, hi) )

给我

1       in.interval.abs(x, lo, hi)  662.732  666.4855  669.939  690.6585 1580.707
2 in.interval.abs.sugar(x, lo, hi)  722.789  726.0920  728.795  742.6085 1671.093
3   in.interval.oldRcpp(x, lo, hi) 1870.784 1876.4890 1892.854 1935.0445 2859.025

> all.equal( in.interval.oldRcpp(x, lo, hi), in.interval.abs(x, lo, hi) )
[1] TRUE

> all.equal( in.interval.oldRcpp(x, lo, hi), in.interval.abs.sugar(x, lo, hi) )
[1] TRUE

【讨论】:

  • 你检查过你的函数返回了什么吗?它们不一样; &amp;&amp; 只计算其操作数的第一个元素。
  • 糟糕——你说得对。可以想象将调用封装在 sapply 或 map 中,但这仍然比其他解决方案慢。
  • 我已将您的代码放入要点中:gist.github.com/4344844。但是,它不能在我的系统上编译(Ubuntu 12.10,来自 CRAN 的最新 R):Error in compileCode(f, code, language = language, verbose = verbose) : ...,error: no match for ‘operator&amp;’ in ...
  • @user946850 您使用的是最新版本的Rcpp 吗?我相信 &amp; 运算符是作为语法糖添加到 Rcpp 0.10.0 中的;见cran.r-project.org/web/packages/Rcpp/vignettes/Rcpp-sugar.pdf。 FWIW,它在 Mac OS、R 2.15.2、Rcpp_0.10.1 上编译得很好。
  • @user946850 我为要点添加了一个潜在的解决方案。应该使用 0.10 之前的 Rcpp 版本进行编译,但 可能 会慢一点。或者,您应该能够在 R 会话中使用 install.packages("Rcpp", type="source") 从 CRAN 获取最新版本,我想。
猜你喜欢
  • 2022-01-18
  • 1970-01-01
  • 1970-01-01
  • 2021-04-04
  • 1970-01-01
  • 2016-05-28
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
相关资源
最近更新 更多