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

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

请问实现该双向标量乘法的最优方法是什么?


最优解法

不需要复杂的模板技巧,只需要实现一次核心乘法逻辑,通过函数复用避免代码重复,以下两种简洁方式都能满足需求:

方式一:成员函数+非成员转发函数

  1. 先在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;
        }
};
  1. 再实现非成员的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 08:09:56