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

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;
}

方案说明

  1. 完全符合你的需求:所有内存都是静态分配,底层存储仍为std::array,没有重复逻辑,后续加法、数乘等其他运算都可以用相同的模式实现。
  2. 不需要修改原有的main函数,原有调用逻辑完全兼容。
  3. 你之前模板报错的原因是:Matrix、MatrixSym本身是模板类而非具体类型,普通class类型参数只能接收具体类型,要接收模板类需要用模板模板参数语法,上面的实现通过类型萃取自动推导参数,比手动指定模板模板参数更简洁易用。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 16:36:03