You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Rcpp模板优化:不同类型列的成对完整观测计算方案问询

Solution: Pairwise Complete Observations for Data Frames in Rcpp

The key issue with your initial list implementation is that you're trying to use runtime type information (RTYPEj/RTYPEk) as template parameters, which requires compile-time constants. To avoid nested switch statements, we can use per-column type dispatch to handle mixed-type columns cleanly. Below are two efficient approaches:


Approach 1: Indicator Matrix + tcrossprod

This approach converts each column to a logical vector indicating non-NA values, then uses matrix cross-product to compute pairwise counts. It leverages optimized BLAS operations for speed.

Code Implementation

// [[Rcpp::plugins(cpp11)]]
#include <Rcpp.h>
using namespace Rcpp;

template <int RTYPE>
LogicalVector non_na_indicator(const Vector<RTYPE>& vec) {
  int n = vec.size();
  LogicalVector res(n);
  
  if (RTYPE == REALSXP) {
    // For real numbers, check if not NaN/NA (x == x is FALSE for NaN)
    for (int i = 0; i < n; ++i) {
      res[i] = (vec[i] == vec[i]);
    }
  } else {
    // For other types, check against the type-specific NA value
    const typename Vector<RTYPE>::storage_type na_val = Vector<RTYPE>::get_na();
    for (int i = 0; i < n; ++i) {
      res[i] = (vec[i] != na_val);
    }
  }
  return res;
}

LogicalVector get_non_na_indicator(SEXP vec) {
  // Dispatch to the correct template based on vector type
  switch (TYPEOF(vec)) {
    case INTSXP:  return non_na_indicator<INTSXP>(vec);
    case REALSXP: return non_na_indicator<REALSXP>(vec);
    case LGLSXP:  return non_na_indicator<LGLSXP>(vec);
    case STRSXP:  return non_na_indicator<STRSXP>(vec);
    default: stop("Unsupported column type: ", Rf_type2char(TYPEOF(vec)));
  }
}

// [[Rcpp::export]]
IntegerMatrix pwNobslCpp(const List& x) {
  int p = x.size();
  if (p == 0) return IntegerMatrix(0, 0);
  
  // Validate all columns have the same length
  int n = Rf_length(x[0]);
  for (int j = 1; j < p; ++j) {
    if (Rf_length(x[j]) != n) {
      stop("All columns must have the same length.");
    }
  }
  
  // Build matrix of non-NA indicators (each column is TRUE/FALSE for non-NA)
  LogicalMatrix indicators(n, p);
  for (int j = 0; j < p; ++j) {
    indicators(_, j) = get_non_na_indicator(x[j]);
  }
  
  // Compute pairwise counts using cross-product (sum of element-wise products)
  IntegerMatrix result = tcrossprod(indicators);
  
  // Set dimnames to match original data frame
  result.attr("dimnames") = List::create(names(x), names(x));
  
  return result;
}

Key Points:

  • Single Switch Per Column: We only dispatch once per column to create the non-NA indicator, avoiding nested switches.
  • Optimized Computation: tcrossprod uses BLAS under the hood, making it fast for large datasets.
  • Memory Usage: Requires O(n*p) memory for the indicator matrix, which is manageable for most use cases.

Approach 2: Type-Erased Checker Functions

This approach uses std::function to store type-erased non-NA checkers for each column, then computes pairwise counts by iterating through rows. It uses less memory than the matrix approach.

Code Implementation

// [[Rcpp::plugins(cpp11)]]
#include <Rcpp.h>
#include <functional>
using namespace Rcpp;

std::function<bool(int)> create_non_na_checker(SEXP vec) {
  switch (TYPEOF(vec)) {
    case INTSXP: {
      IntegerVector iv(vec);
      return [iv](int i) { return iv[i] != NA_INTEGER; };
    }
    case REALSXP: {
      NumericVector nv(vec);
      return [nv](int i) { return nv[i] == nv[i]; };
    }
    case LGLSXP: {
      LogicalVector lv(vec);
      return [lv](int i) { return lv[i] != NA_LOGICAL; };
    }
    case STRSXP: {
      StringVector sv(vec);
      return [sv](int i) { return !sv[i].is_na(); };
    }
    default: stop("Unsupported column type: ", Rf_type2char(TYPEOF(vec)));
  }
}

// [[Rcpp::export]]
IntegerMatrix pwNobslCpp(const List& x) {
  int p = x.size();
  if (p == 0) return IntegerMatrix(0, 0);
  
  // Validate all columns have the same length
  int n = Rf_length(x[0]);
  for (int j = 1; j < p; ++j) {
    if (Rf_length(x[j]) != n) {
      stop("All columns must have the same length.");
    }
  }
  
  // Create a vector of type-erased non-NA checkers
  std::vector<std::function<bool(int)>> checkers(p);
  for (int j = 0; j < p; ++j) {
    checkers[j] = create_non_na_checker(x[j]);
  }
  
  IntegerMatrix out(p, p);
  
  // Compute diagonal (self-counts) and pairwise counts
  for (int j = 0; j < p; ++j) {
    int nj = 0;
    for (int i = 0; i < n; ++i) {
      if (checkers[j](i)) nj++;
    }
    out(j, j) = nj;
    
    for (int k = j + 1; k < p; ++k) {
      int count = 0;
      for (int i = 0; i < n; ++i) {
        if (checkers[j](i) && checkers[k](i)) count++;
      }
      out(j, k) = out(k, j) = count;
    }
  }
  
  // Set dimnames
  out.attr("dimnames") = List::create(names(x), names(x));
  
  return out;
}

Key Points:

  • No Nested Switches: Only a single switch per column to create the checker function.
  • Lower Memory Usage: Uses O(p) memory for the checkers instead of O(n*p) for the indicator matrix.
  • Flexibility: Easy to extend to additional column types by adding cases to the switch.

Testing the Function

Both implementations will correctly handle mixed-type data frames. For example:

df <- data.frame(
  int_col = c(1, NA, 3, 4),
  num_col = c(NA, 2.5, 3.0, NA),
  log_col = c(TRUE, FALSE, NA, TRUE),
  str_col = c("a", NA, "c", "d")
)

pwNobslCpp(df)

This should return a 4x4 matrix where each entry (j,k) is the number of rows with non-NA values in both columns j and k.

内容的提问来源于stack exchange,提问作者Sebastian

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.07 10:37:32