表达式模板实现矩阵乘法:(A+B)*(C+D)编译失败求助
问题:支持表达式模板中两个
AddExp对象的矩阵乘法 我正在用表达式模板实现矩阵代数,通过懒求值处理逐元素操作(如+、-、+=等),矩阵乘法则采用立即求值策略。目前代码支持A * B、A * (B+C)、(A+B)*C这类表达式,但无法支持(A+B)*(C+D)——即左右操作数均为AddExp对象的情况。尽管已添加相关方法定义,编译器仍无法匹配到正确的函数签名,求指点需要补充什么定义来支持该表达式。
问题根源
AddExp类中的operator*是非const成员函数,而(C+C)这类表达式会生成临时对象,临时对象只能调用const成员函数,导致无法匹配。- 原
operator*的维度模板参数逻辑错误,没有正确推导结果矩阵的行列数。 - 缺少通用的表达式间乘法重载,无法覆盖所有表达式组合的场景。
具体修改步骤
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
相关产品推荐
相关产品推荐

