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:
tcrossproduses 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

