C++如何实现支持普通矩阵与对称矩阵输入的通用乘法函数
方案实现
首先修正你代码里的笔误:当前基类定义名为MatrixAbstract,但子类继承时写的是MatrixSuperclass,需要先统一名称。
步骤1:修改基类暴露模板参数
在MatrixAbstract的public区域添加静态常量,方便后续模板萃取矩阵维度:
template<unsigned int ROWS, unsigned int COLS, unsigned int NUMEL> class MatrixAbstract { private: std::array<double, NUMEL> matVals; public: static constexpr unsigned int rows = ROWS; static constexpr unsigned int cols = COLS; static constexpr unsigned int numel = NUMEL; MatrixAbstract(){} virtual unsigned int index_from_rc(const unsigned int& row, const unsigned int& col) const = 0; // get/set方法保持不变 double get_value(const int& row, const int& col) const { return this->matVals[this->index_from_rc(row,col)]; } void set_value(const int& row, const int& col, double value) { this->matVals[this->index_from_rc(row,col)] = value; } };
步骤2:实现类型萃取工具
引入头文件<type_traits>,添加两个工具元函数,用于判断矩阵类型、推导乘法返回值:
// 判断是否为对称矩阵类型 template <typename T> struct is_matrix_sym : std::false_type {}; template <unsigned int ROWS, unsigned int COLS, unsigned int NUMEL> struct is_matrix_sym<MatrixSym<ROWS, COLS, NUMEL>> : std::true_type {}; template <typename T> inline constexpr bool is_matrix_sym_v = is_matrix_sym<T>::value; // 推导乘法结果类型:两个输入都是对称矩阵则返回对称矩阵,否则返回普通矩阵 template <typename MatA, typename MatB> struct multiply_result { static constexpr unsigned int out_rows = MatA::rows; static constexpr unsigned int out_cols = MatB::cols; using type = std::conditional_t< is_matrix_sym_v<MatA> && is_matrix_sym_v<MatB>, MatrixSym<out_rows, out_cols>, Matrix<out_rows, out_cols> >; }; template <typename MatA, typename MatB> using multiply_result_t = typename multiply_result<MatA, MatB>::type;
步骤3:实现通用乘法函数
删除原来4个重复的multiply函数,替换为以下通用实现:
template <typename MatA, typename MatB> multiply_result_t<MatA, MatB> multiply(const MatA& a, const MatB& b) { static_assert(MatA::cols == MatB::rows, "矩阵乘法要求第一个矩阵的列数等于第二个矩阵的行数"); using OutMat = multiply_result_t<MatA, MatB>; OutMat out; constexpr unsigned int ROWS = MatA::rows; constexpr unsigned int COLS = MatB::cols; constexpr unsigned int INNER = MatA::cols; for (unsigned int r = 0; r < ROWS; r++) { for (unsigned int c = 0; c < COLS; c++) { double val = 0.0; for (unsigned int rc = 0; rc < INNER; rc++) { val += a.get_value(r, rc) * b.get_value(rc, c); } out.set_value(r, c, val); } } return out; }
方案说明
- 完全符合你的需求:所有内存都是静态分配,底层存储仍为
std::array,没有重复逻辑,后续加法、数乘等其他运算都可以用相同的模式实现。 - 不需要修改原有的
main函数,原有调用逻辑完全兼容。 - 你之前模板报错的原因是:
Matrix、MatrixSym本身是模板类而非具体类型,普通class类型参数只能接收具体类型,要接收模板类需要用模板模板参数语法,上面的实现通过类型萃取自动推导参数,比手动指定模板模板参数更简洁易用。
内容的提问来源于stack exchange,提问作者MountainDrew
相关产品推荐
相关产品推荐

