C++自定义mv类实现标量双向乘法的最优方案咨询
双向标量乘法的最优实现方法
问题描述
我正在完成大学练习,编写了一个模仿MATLAB向量行为的mv类,代码如下:
#include <iostream> #include <vector> #include <tuple> using namespace std; double const pi = 3.141592; class mv{ public: mv(size_t size = 0, double init = 0) {for(size_t i = 0; i < size; i++) a.push_back(init);}; size_t size() const {return a.size();}; // 注:原代码中`a.vector<double>::size()`写法错误,已修正为`a.size()` void print() const {for(auto &el : a) std::cout << el << " ";}; double get(size_t n) const {return a[n];}; double& operator[](size_t n) {return a[n];}; private: vector<double> a; };
我希望为该类添加与标量的乘法功能,且支持双向运算(即a * 3和3 * a均能正常工作),主函数示例如下:
int main() { mv a(5, 1); a[3] = 2; mv b(5); b = a * pi; b.print(); cout << endl; mv c(5); c = pi * a; c.print(); cout << endl; return 0; }
最直接的实现方式是编写两个operator*重载函数,但存在明显的代码重复:
mv operator*(const mv& a, const double& scalar){ mv res(a.size()); for(size_t i=0; i<a.size(); i++) res[i] = a.get(i) * scalar; return res; }; mv operator*(const double& scalar, const mv& a){ mv res(a.size()); for(size_t i=0; i<a.size(); i++) res[i] = a.get(i) * scalar; return res; };
我尝试用模板方式减少重复,但实现笨拙且易出错:
template<typename T, typename V> mv operator*(const T& lhs, const V& rhs){ auto tuple = std::tie(lhs, rhs); return mvTimes(std::get<const mv&>(tuple), std::get<const double&>(tuple)); }; mv mvTimes(const mv& a, const double& scalar){ mv res(a.size()); for(size_t i=0; i<a.size(); i++) res[i] = a.get(i) * scalar; return res; };
请问实现该双向标量乘法的最优方法是什么?
最优解法
不需要复杂的模板技巧,只需要实现一次核心乘法逻辑,通过函数复用避免代码重复,以下两种简洁方式都能满足需求:
方式一:成员函数+非成员转发函数
- 先在
mv类中实现成员版的operator*,处理mv * scalar的场景:
class mv{ // 保留原有成员 public: // 添加成员乘法运算符 mv operator*(double scalar) const { mv res(size()); for(size_t i = 0; i < size(); ++i) { res[i] = a[i] * scalar; // 直接访问私有成员a,比调用get()更高效 } return res; } };
- 再实现非成员的
operator*,处理scalar * mv的场景,直接复用成员函数的逻辑:
mv operator*(double scalar, const mv& vec) { // 利用乘法交换律,直接返回vec * scalar return vec * scalar; }
这种方式逻辑清晰,无代码冗余,且直接访问私有成员的效率更高。
方式二:带类型约束的模板(C++11及以上)
如果想更严谨地限制只有mv和double的组合才能匹配乘法运算符,避免模板被其他无关类型误匹配,可以结合std::enable_if做类型约束:
#include <type_traits> // 处理 mv * scalar template<typename T, typename = std::enable_if_t<std::is_same_v<T, double>>> mv operator*(const mv& vec, T scalar) { mv res(vec.size()); for(size_t i = 0; i < vec.size(); ++i) { res[i] = vec.get(i) * scalar; } return res; } // 处理 scalar * mv template<typename T, typename = std::enable_if_t<std::is_same_v<T, double>>> mv operator*(T scalar, const mv& vec) { return vec * scalar; }
这种方式同样避免了代码重复,同时通过类型约束提升了代码的安全性。
内容的提问来源于stack exchange,提问作者Francesco Derme
相关产品推荐
相关产品推荐

