【问题标题】:Faster way to find the first TRUE value in a vector在向量中找到第一个 TRUE 值的更快方法
【发布时间】:2013-11-27 15:47:54
【问题描述】:

在一个函数中,我经常需要使用如下代码:

which(x==1)[1]
which(x>1)[1]
x[x>10][1]

其中x 是一个数字向量。 summaryRprof() 表明我在关系运算符上花费了超过 80% 的时间。我想知道是否有一个函数只在达到第一个 TRUE 值之前进行比较以加速我的代码。 For-loop 比上面提供的选项慢。

【问题讨论】:

  • which.minwhich.max 考虑类似的问题。
  • @James:不完全是,因为它们需要一个逻辑向量,而创建逻辑向量很耗时。

标签: r performance


【解决方案1】:

我不知道纯 R 方法可以做到这一点,所以我写了一个 C function 来为 quantstrat 包做这件​​事。这个函数是为特定目的而编写的,所以它不像我想要的那样通用。例如,您可能会注意到它仅适用于实数/双精度/数字数据,因此请务必在调用 .firstCross 函数之前将 Data 强制转换为该值。

#include <R.h>
#include <Rinternals.h>

SEXP firstCross(SEXP x, SEXP th, SEXP rel, SEXP start)
{
    int i, int_rel, int_start;
    double *real_x=NULL, real_th;

    if(ncols(x) > 1)
        error("only univariate data allowed");

    /* this currently only works for real x and th arguments
     * support for other types may be added later */
    real_th = asReal(th);
    int_rel = asInteger(rel);
    int_start = asInteger(start)-1;

    switch(int_rel) {
        case 1:  /* >  */
            real_x = REAL(x);
            for(i=int_start; i<nrows(x); i++)
                if(real_x[i] >  real_th)
                    return(ScalarInteger(i+1));
            break;
        case 2:  /* <  */
            real_x = REAL(x);
            for(i=int_start; i<nrows(x); i++)
                if(real_x[i] <  real_th)
                    return(ScalarInteger(i+1));
            break;
        case 3:  /* == */
            real_x = REAL(x);
            for(i=int_start; i<nrows(x); i++)
                if(real_x[i] == real_th)
                    return(ScalarInteger(i+1));
            break;
        case 4:  /* >= */
            real_x = REAL(x);
            for(i=int_start; i<nrows(x); i++)
                if(real_x[i] >= real_th)
                    return(ScalarInteger(i+1));
            break;
        case 5:  /* <= */
            real_x = REAL(x);
            for(i=int_start; i<nrows(x); i++)
                if(real_x[i] <= real_th)
                    return(ScalarInteger(i+1));
            break;
        default:
            error("unsupported relationship operator");
  }
  /* return number of observations if relationship is never TRUE */
  return(ScalarInteger(nrows(x)));
}

这是调用它的 R 函数:

.firstCross <- function(Data, threshold=0, relationship, start=1) {
    rel <- switch(relationship[1],
            '>'    =  ,
            'gt'   = 1,
            '<'    =  ,
            'lt'   = 2,
            '=='   =  ,
            'eq'   = 3,
            '>='   =  ,
            'gte'  =  ,
            'gteq' =  ,
            'ge'   = 4,
            '<='   =  ,
            'lte'  =  ,
            'lteq' =  ,
            'le'   = 5)
    .Call('firstCross', Data, threshold, rel, start)
}

一些基准测试,只是为了好玩。

> library(quantstrat)
> library(microbenchmark)
> firstCross <- quantstrat:::.firstCross
> set.seed(21)
> x <- rnorm(1e6)
> microbenchmark(which(x > 3)[1], firstCross(x,3,">"), times=10)
Unit: microseconds
                  expr      min       lq    median       uq      max neval
       which(x > 3)[1] 9482.081 9578.072 9597.3870 9690.448 9820.176    10
 firstCross(x, 3, ">")   11.370   11.675   31.9135   34.443   38.614    10
> which(x>3)[1]
[1] 919
> firstCross(x,3,">")
[1] 919

请注意,firstCross 将产生更大的相对加速,Data 越大(因为 R 的关系运算符必须完成整个向量的比较)。

> x <- rnorm(1e7)
> microbenchmark(which(x > 3)[1], firstCross(x,3,">"), times=10)
Unit: microseconds
                  expr      min        lq    median        uq        max neval
       which(x > 3)[1] 94536.21 94851.944 95799.857 96154.756 113962.794    10
 firstCross(x, 3, ">")     5.08     5.507    25.845    32.164     34.183    10
> which(x>3)[1]
[1] 97
> firstCross(x,3,">")
[1] 97

...如果第一个 TRUE 值接近向量的末尾,它不会明显更快。

> microbenchmark(which(x==last(x))[1], firstCross(x,last(x),"eq"),times=10)
Unit: milliseconds
                         expr      min       lq   median       uq       max neval
       which(x == last(x))[1] 92.56311 93.85415 94.38338 98.18422 106.35253    10
 firstCross(x, last(x), "eq") 86.55415 86.70980 86.98269 88.32168  92.97403    10
> which(x==last(x))[1]
[1] 10000000
> firstCross(x,last(x),"eq")
[1] 10000000

【讨论】:

  • 检查:firstCross(as.numeric(1:100),as.numeric(200),"gte")。我希望出现错误或 NA,而不仅仅是最后一行。
  • @user1603038:你为什么会这样?代码清楚地写着“/* return number of observations if relationship is never TRUE */”。 ;) 只需将 C 函数更改为返回 NA_INTEGER,如果这是您想要的。正如我所说,该函数是为特定用例编写的,并不是特别通用。
  • 没注意到那条评论,抱歉。
  • @user1603038:别担心,反正我是在讽刺。您可能会遇到其他类似的情况,但您始终可以更改代码并使用 R CMD SHLIB 重新编译。
  • @user1603038: NA_REAL 是特殊的NaN,IEEE 标准基本上说与NaN 比较产生false
【解决方案2】:

Base R 提供PositionFind 分别用于定位第一个索引和值,谓词为其返回真值。这些高阶函数在第一次命中时立即返回。

f<-function(x) {
  r<-vector("list",3)
  r[[1]]<-which(x==1)[1]
  r[[2]]<-which(x>1)[1]
  r[[3]]<-x[x>10][1]
  return(r)
}

p<-function(f,b) function(a) f(a,b)
g<-function(x) {
  r<-vector("list",3)
  r[[1]]<-Position(p(`==`,1),x)
  r[[2]]<-Position(p(`>`,1),x)
  r[[3]]<-Find(p(`>`,10),x)
  return(r)
}

相对性能很大程度上取决于相对于谓词成本与Position/Find 开销的早期发现命中的概率。

library(microbenchmark)
set.seed(1)
x<-sample(1:100,1e5,replace=TRUE)
microbenchmark(f(x),g(x))

Unit: microseconds
 expr      min        lq     mean    median        uq      max neval cld
 f(x) 5034.283 5410.1205 6313.861 5798.4780 6948.5675 26735.52   100   b
 g(x)  587.463  650.4795 1013.183  734.6375  950.9845 20285.33   100  a

y<-rep(0,1e5)
microbenchmark(f(y),g(y))

Unit: milliseconds
 expr        min         lq       mean     median         uq        max neval cld
 f(y)   3.470179   3.604831   3.791592   3.718752   3.866952   4.831073   100  a 
 g(y) 131.250981 133.687454 137.199230 134.846369 136.193307 177.082128   100   b

【讨论】:

  • PositionFind 只是 for 循环的语法糖。并不是说这有什么问题,只是 OP 提到 for 循环较慢。
【解决方案3】:

这是一个很好的问题和答案...只是添加 any() 并不比 which()match() 快,但两者都比 [] 快,我猜这可能会创建一个大向量无用的 T,F。所以我猜没有..缺少上面的答案。

    v=rep('A', 10e6)
    v[5e6]='B'
    v[10e6]='B'

    microbenchmark(which(v=='B')[1])
    Unit: milliseconds
                   expr      min       lq   median       uq      max neval
     which(v == "B")[1] 332.3788 337.6718 344.4076 347.1194 503.4022   100

    microbenchmark(any(v=='B'))
    Unit: milliseconds
              expr      min      lq   median       uq      max neval
     any(v == "B") 334.4466 335.114 335.6714 347.5474 356.0261   100

    microbenchmark(v[v=='B'][1])
    Unit: milliseconds
               expr      min       lq  median       uq      max neval
     v[v == "B"][1] 601.5923 605.3331 609.191 612.0689 707.1409   100

    microbenchmark(match("B", v))
    Unit: milliseconds
               expr      min       lq   median       uq      max neval
    match("B", v) 339.2872 344.7648 350.5444 359.6746 915.6446   100

还有其他想法吗?

【讨论】:

  • 正如 OP 所说,分析表明 >80% 的时间花在关系运算符上,所以只要你仍然使用 ==,你不应该期望任何比边际速度提高更多的东西.
  • 确实,我只是好奇any 或其他一些函数是否可能被编码为在发现某些东西时停止,从而减少 ==。显然不是。即便如此,这也不是一个通用的解决方案。
  • 创建逻辑向量非常耗时,因此任何需要逻辑向量的 R 函数(例如 any)都无济于事。此外,任何无条件地对整个对象进行操作的函数都无济于事。正如我在回答中所说的那样,我想不出在纯 R 中做到这一点的方法......这并不是说我没有尝试过。就我而言,Data 已排序,而通常快速的 findInterval 并没有太大帮助。
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 2011-10-16
  • 1970-01-01
  • 2017-07-25
  • 2011-01-28
  • 2021-06-20
  • 1970-01-01
相关资源
最近更新 更多