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

表达式模板实现矩阵乘法:(A+B)*(C+D)编译失败求助

问题:支持表达式模板中两个AddExp对象的矩阵乘法

我正在用表达式模板实现矩阵代数,通过懒求值处理逐元素操作(如+、-、+=等),矩阵乘法则采用立即求值策略。目前代码支持A * B、A * (B+C)、(A+B)*C这类表达式,但无法支持(A+B)*(C+D)——即左右操作数均为AddExp对象的情况。尽管已添加相关方法定义,编译器仍无法匹配到正确的函数签名,求指点需要补充什么定义来支持该表达式。


问题根源

  1. AddExp类中的operator*是非const成员函数,而(C+C)这类表达式会生成临时对象,临时对象只能调用const成员函数,导致无法匹配。
  2. 原operator*的维度模板参数逻辑错误,没有正确推导结果矩阵的行列数。
  3. 缺少通用的表达式间乘法重载,无法覆盖所有表达式组合的场景。

具体修改步骤

1. 为AddExp的operator*添加const修饰,并修正维度逻辑

临时对象只能调用const成员函数,同时要根据左右操作数的维度正确推导结果矩阵的尺寸:

template <typename LHS, typename RHS, typename T>
class AddExp {
public:
    // ... 其他成员不变

    // 修正:添加const修饰,同时正确推导结果维度
    template <typename L2, typename R2>
    Matrix<T, LHS::ROWS, R2::COLS> operator*(const AddExp<L2, R2, T>& exp2) const {
        static_assert(LHS::COLS == L2::ROWS, "Matrix multiplication dimension mismatch: left cols != right rows");
        Matrix<T, LHS::ROWS, R2::COLS> result;
        
        int kMax = LHS::COLS;
        for (int i{}; i < LHS::ROWS; ++i)
        {
            for (int j{}; j < R2::COLS; ++j)
            {
                T sum = T{};
                for (int k{}; k < kMax; ++k)
                {
                    sum += (*this)(i, k) * exp2(k, j);
                }
                result(i, j) = sum;
            }
        }
        return result;
    }

    // 同样给Matrix乘法的operator*添加const修饰
    template <int P>
    Matrix<T, LHS::ROWS, P> operator*(const Matrix<T, LHS::COLS, P>& exp) const {
        static_assert(LHS::COLS == Matrix<T, LHS::COLS, P>::COLS, "Dimension mismatch");
        Matrix<T, LHS::ROWS, P> result;
        
        int kMax = LHS::COLS;
        for (int i{}; i < LHS::ROWS; ++i)
        {
            for (int j{}; j < P; ++j)
            {
                T sum = T{};
                for (int k{}; k < kMax; ++k)
                {
                    sum += (*this)(i, k) * exp(k, j);
                }
                result(i, j) = sum;
            }
        }
        return result;
    }
};

2. 给Matrix类添加静态维度常量

为了让AddExp能直接获取行列数的编译时常量,需要给Matrix添加静态成员:

template <typename T = double, int ROWS = 3, int COLS = 3>
class Matrix {
public:
    static constexpr int ROWS = ROWS;
    static constexpr int COLS = COLS;
    // ... 其他成员不变
};

3. 添加全局通用乘法重载(可选,增强扩展性)

为了支持任意表达式类型的乘法(如AddExp * SubExp、SubExp * AddExp等),可以添加全局模板:

template <typename LExp, typename RExp, typename T>
Matrix<T, LExp::ROWS, RExp::COLS> operator*(const LExp& lhs, const RExp& rhs) {
    static_assert(LExp::COLS == RExp::ROWS, "Matrix multiplication dimension mismatch");
    Matrix<T, LExp::ROWS, RExp::COLS> result;
    for (int i = 0; i < LExp::ROWS; ++i) {
        for (int j = 0; j < RExp::COLS; ++j) {
            T sum = T{};
            for (int k = 0; k < LExp::COLS; ++k) {
                sum += lhs(i, k) * rhs(k, j);
            }
            result(i, j) = sum;
        }
    }
    return result;
}

4. 修正Matrix类中operator*的返回类型错误

原Matrix类中operator*(const Matrix<T, COLS, P>& exp)的返回类型错误,应该返回Matrix<T, ROWS, P>:

template <int P>
Matrix<T, ROWS, P> operator*(const Matrix<T, COLS, P>& exp) const {
    Matrix<T, ROWS, P> result;
    assert(COLS == exp.getRows());
    int kMax = COLS;
    for (int i{}; i < ROWS; ++i)
    {
        for (int j{}; j < P; ++j)
        {
            T sum = T{};
            for (int k{}; k < kMax; ++k)
            {
                sum += (*this)(i, k) * exp(k, j);
            }
            result(i, j) = sum;
        }
    }
    return result;
}

测试验证

修改完成后,(C+C)*(A+B)即可正常编译运行,临时AddExp对象能正确调用const版本的乘法运算符,维度推导也符合矩阵乘法规则。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 20:44:57