如何在Rcpp::List中适配arma::mat与arma::sp_mat并操作混合矩阵数组?
解决方案
首先明确:Armadillo的arma::mat(稠密矩阵)和arma::sp_mat(稀疏矩阵)没有共同的直接父类,它们是独立实现的模板类,共享的操作(比如乘法、维度获取)是通过模板重载实现的,而非继承体系。针对你的需求,以下是几种最简实现方式:
方法一:使用C变体类型(推荐C17及以上)
利用C++17的std::variant存储两种矩阵类型的指针或对象,配合std::visit实现统一操作,无需每次手动判断类型。
代码示例
#include <RcppArmadillo.h> #include <variant> #include <vector> // [[Rcpp::depends(RcppArmadillo)]] // [[Rcpp::plugins(cpp17)]] // 定义变体类型,存储稠密或稀疏矩阵的指针 using MatrixVariant = std::variant<arma::mat*, arma::sp_mat*>; // 统一操作的函数对象,比如打印矩阵维度 struct PrintDim { void operator()(arma::mat* mat) const { Rcpp::Rcout << "Dense matrix: " << mat->n_rows << "x" << mat->n_cols << "\n"; } void operator()(arma::sp_mat* spmat) const { Rcpp::Rcout << "Sparse matrix: " << spmat->n_rows << "x" << spmat->n_cols << "\n"; } }; // [[Rcpp::export]] void process_mat_list(Rcpp::List mat_list) { std::vector<MatrixVariant> mat_ptrs; // 遍历输入列表,转换并存储指针 for (int i = 0; i < mat_list.size(); ++i) { SEXP elem = mat_list[i]; if (Rcpp::is<arma::mat>(elem)) { arma::mat* mat = new arma::mat(Rcpp::as<arma::mat>(elem)); mat_ptrs.emplace_back(mat); } else if (Rcpp::is<arma::sp_mat>(elem)) { arma::sp_mat* spmat = new arma::sp_mat(Rcpp::as<arma::sp_mat>(elem)); mat_ptrs.emplace_back(spmat); } } // 执行统一操作 for (auto& var : mat_ptrs) { std::visit(PrintDim(), var); // 可添加更多操作,比如矩阵乘法、元素访问等 } // 释放内存(改用std::unique_ptr可避免手动释放) for (auto& var : mat_ptrs) { std::visit([](auto ptr) { delete ptr; }, var); } }
方法二:自定义抽象基类封装
如果需要更灵活的接口控制,可以自定义抽象基类,让稠密/稀疏矩阵的封装类继承它,通过多态实现统一操作。
代码示例
#include <RcppArmadillo.h> #include <memory> #include <vector> // [[Rcpp::depends(RcppArmadillo)]] // 抽象基类 class MatrixBase { public: virtual ~MatrixBase() = default; virtual void print_dim() const = 0; virtual arma::uword n_rows() const = 0; virtual arma::uword n_cols() const = 0; // 可添加更多统一接口,比如乘法、转置等 }; // 稠密矩阵封装类 class DenseMatrix : public MatrixBase { private: arma::mat mat_; public: DenseMatrix(const arma::mat& mat) : mat_(mat) {} void print_dim() const override { Rcpp::Rcout << "Dense matrix: " << mat_.n_rows << "x" << mat_.n_cols << "\n"; } arma::uword n_rows() const override { return mat_.n_rows; } arma::uword n_cols() const override { return mat_.n_cols; } }; // 稀疏矩阵封装类 class SparseMatrix : public MatrixBase { private: arma::sp_mat spmat_; public: SparseMatrix(const arma::sp_mat& spmat) : spmat_(spmat) {} void print_dim() const override { Rcpp::Rcout << "Sparse matrix: " << spmat_.n_rows << "x" << spmat_.n_cols << "\n"; } arma::uword n_rows() const override { return spmat_.n_rows; } arma::uword n_cols() const override { return spmat_.n_cols; } }; // [[Rcpp::export]] void process_mat_list(Rcpp::List mat_list) { std::vector<std::unique_ptr<MatrixBase>> mats; for (int i = 0; i < mat_list.size(); ++i) { SEXP elem = mat_list[i]; if (Rcpp::is<arma::mat>(elem)) { mats.emplace_back(std::make_unique<DenseMatrix>(Rcpp::as<arma::mat>(elem))); } else if (Rcpp::is<arma::sp_mat>(elem)) { mats.emplace_back(std::make_unique<SparseMatrix>(Rcpp::as<arma::sp_mat>(elem))); } } // 统一调用接口 for (const auto& mat : mats) { mat->print_dim(); } }
方法三:简化版Rcpp::List配合类型判断(适合简单场景)
如果操作逻辑不复杂,也可以直接在Rcpp::List中存储原始对象,每次操作时通过Rcpp::is判断类型并转换——虽然需要显式转换,但实现最简单:
#include <RcppArmadillo.h> // [[Rcpp::depends(RcppArmadillo)]] // [[Rcpp::export]] void process_mat_list(Rcpp::List mat_list) { for (int i = 0; i < mat_list.size(); ++i) { SEXP elem = mat_list[i]; if (Rcpp::is<arma::mat>(elem)) { arma::mat mat = Rcpp::as<arma::mat>(elem); // 处理稠密矩阵 Rcpp::Rcout << "Dense: " << mat.n_rows << "x" << mat.n_cols << "\n"; } else if (Rcpp::is<arma::sp_mat>(elem)) { arma::sp_mat spmat = Rcpp::as<arma::sp_mat>(elem); // 处理稀疏矩阵 Rcpp::Rcout << "Sparse: " << spmat.n_rows << "x" << spmat.n_cols << "\n"; } } }
内容的提问来源于stack exchange,提问作者mskb
相关产品推荐
相关产品推荐

