And*_*ell 3 c++ r rcpp desctools
我有一个数字向量v(已经省略了NA),并希望获得第n个最大值及其各自的频率.
我发现 http://gallery.rcpp.org/articles/top-elements-from-vectors-using-priority-queue/ 非常快.
// [[Rcpp::export]]
std::vector<int> top_i_pq(NumericVector v, unsigned int n)
{
typedef pair<double, int> Elt;
priority_queue< Elt, vector<Elt>, greater<Elt> > pq;
vector<int> result;
for (int i = 0; i != v.size(); ++i) {
if (pq.size() < n)
pq.push(Elt(v[i], i));
else {
Elt elt = Elt(v[i], i);
if (pq.top() < elt) {
pq.pop();
pq.push(elt);
}
}
}
result.reserve(pq.size());
while (!pq.empty()) {
result.push_back(pq.top().second + 1);
pq.pop();
}
return result ;
}
Run Code Online (Sandbox Code Playgroud)
但是,关系不会得到尊重.实际上我不需要索引,返回值也可以.
我想得到的是一个包含值和频率的列表,例如:
numv <- c(4.2, 4.2, 4.5, 0.1, 4.4, 2.0, 0.9, 4.4, 3.3, 2.4, 0.1)
top_i_pq(numv, 3)
$lengths
[1] 2 2 1
$values
[1] 4.2 4.4 4.5
Run Code Online (Sandbox Code Playgroud)
获得一个独特的向量,一个表,一个(完整)排序都不是一个好主意,因为与v的长度(可能很容易> 1e6)相比,n通常很小.
迄今解决方案是:
library(microbenchmark)
library(data.table)
library(DescTools)
set.seed(1789)
x <- sample(round(rnorm(1000), 3), 1e5, replace = TRUE)
n <- 5
microbenchmark(
BaseR = tail(table(x), n),
data.table = data.table(x)[, .N, keyby = x][(.N - n + 1):.N],
DescTools = Large(x, n, unique=TRUE),
Coatless = ...
)
Unit: milliseconds
expr min lq mean median uq max neval
BaseR 188.09662 190.830975 193.189422 192.306297 194.02815 253.72304 100
data.table 11.23986 11.553478 12.294456 11.768114 12.25475 15.68544 100
DescTools 4.01374 4.174854 5.796414 4.410935 6.70704 64.79134 100
Run Code Online (Sandbox Code Playgroud)
嗯,DescTools仍然是最快的,但我相信它可以通过Rcpp显着改善(因为它是纯R)!
我想用另一个基于Rcpp的解决方案将我的帽子戴在戒指上,使用上面的1e5长度和样本数据,它比DescTools方法快约7倍,比接近快约13倍.实施有点冗长,所以我将带领基准:data.tablexn = 5
fn.dt <- function(v, n) {
data.table(v = v)[
,.N, keyby = v
][(.N - n + 1):.N]
}
microbenchmark(
"DescTools" = Large(x, n, unique=TRUE),
"top_n" = top_n(x, 5),
"data.table" = fn.dt(x, n),
times = 500L
)
# Unit: microseconds
# expr min lq mean median uq max neval
# DescTools 3330.527 3790.035 4832.7819 4070.573 5323.155 54921.615 500
# top_n 566.207 587.590 633.3096 593.577 640.832 3568.299 500
# data.table 6920.636 7380.786 8072.2733 7764.601 8585.472 14443.401 500
Run Code Online (Sandbox Code Playgroud)
更新
如果您的编译器支持C++ 11,您可以利用std::priority_queue::emplace(令人惊讶的)显着的性能提升(与下面的C++ 98版本相比).我不会发布这个版本,因为它几乎是相同的,除了几个调用std::move和emplace,但这里是一个链接.
对前三个函数进行测试,并使用data.table1.9.7(比1.9.6快一点)产生
print(res2, order = "median", signif = 3)
# Unit: relative
# expr min lq mean median uq max neval cld
# top_n2 1.0 1.00 1.000000 1.00 1.00 1.00 1000 a
# top_n 1.6 1.58 1.666523 1.58 1.75 2.75 1000 b
# DescTools 10.4 10.10 8.512887 9.68 7.19 12.30 1000 c
# data.table-1.9.7 16.9 16.80 14.164139 15.50 10.50 43.70 1000 d
Run Code Online (Sandbox Code Playgroud)
哪里top_n2是C++ 11版本.
该top_n功能实现如下:
#include <Rcpp.h>
#include <utility>
#include <queue>
class histogram {
private:
struct paired {
typedef std::pair<double, unsigned int> pair_t;
pair_t pair;
unsigned int is_set;
paired()
: pair(pair_t()),
is_set(0)
{}
paired(double x)
: pair(std::make_pair(x, 1)),
is_set(1)
{}
bool operator==(const paired& other) const {
return pair.first == other.pair.first;
}
bool operator==(double other) const {
return is_set && (pair.first == other);
}
bool operator>(double other) const {
return is_set && (pair.first > other);
}
bool operator<(double other) const {
return is_set && (pair.first < other);
}
paired& operator++() {
++pair.second;
return *this;
}
paired operator++(int) {
paired tmp(*this);
++(*this);
return tmp;
}
};
struct greater {
bool operator()(const paired& lhs, const paired& rhs) const {
if (!lhs.is_set) return false;
if (!rhs.is_set) return true;
return lhs.pair.first > rhs.pair.first;
}
};
typedef std::priority_queue<
paired,
std::vector<paired>,
greater
> queue_t;
unsigned int sz;
queue_t queue;
void insert(double x) {
if (queue.empty()) {
queue.push(paired(x));
return;
}
if (queue.top() > x && queue.size() >= sz) return;
queue_t qtmp;
bool matched = false;
while (queue.size()) {
paired elem = queue.top();
if (elem == x) {
qtmp.push(++elem);
matched = true;
} else {
qtmp.push(elem);
}
queue.pop();
}
if (!matched) {
if (qtmp.size() >= sz) qtmp.pop();
qtmp.push(paired(x));
}
std::swap(queue, qtmp);
}
public:
histogram(unsigned int sz_)
: sz(sz_),
queue(queue_t())
{}
template <typename InputIt>
void insert(InputIt first, InputIt last) {
for ( ; first != last; ++first) {
insert(*first);
}
}
Rcpp::List get() const {
Rcpp::NumericVector values(sz);
Rcpp::IntegerVector freq(sz);
R_xlen_t i = 0;
queue_t tmp(queue);
while (tmp.size()) {
values[i] = tmp.top().pair.first;
freq[i] = tmp.top().pair.second;
++i;
tmp.pop();
}
return Rcpp::List::create(
Rcpp::Named("value") = values,
Rcpp::Named("frequency") = freq);
}
};
// [[Rcpp::export]]
Rcpp::List top_n(Rcpp::NumericVector x, int n = 5) {
histogram h(n);
h.insert(x.begin(), x.end());
return h.get();
}
Run Code Online (Sandbox Code Playgroud)
histogram上面的课程中有很多内容,但只是触及一些关键点:
paired类型本质上是围绕a的包装类std::pair<double, unsigned int>,它将值与计数相关联,提供一些便利功能,例如operator++()/ operator++(int)用于计数的直接前/后递增,以及修改的比较运算符.histogram类包装一种"管理"优先级队列的,在这个意义上的大小std::priority_queue是在一个特定的值上限sz.std::less排序std::priority_queue,而是使用大于比较器,以便可以检查候选值std::priority_queue::top()以快速确定是否应该(a)丢弃它们,(b)替换队列中的当前最小值,或者(c)更新队列中现有值之一的计数.这是唯一可能的,因为队列的大小被限制为<= sz.