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

如何实现支持重载操作的std::reduce?求可编译代码示例

如何让std::reduce支持重载的二元操作函数

想要用std::reduce的并行策略替代std::accumulate,已经为二元操作实现了4个对应不同参数类型的重载版本,但直接传递重载函数名给std::reduce时无法编译通过——示例代码中定义了baz的四个重载,但std::reduce(v.begin(), v.end(), std::pair(0.0,1.0), baz)编译失败。

问题根源

std::reduce是模板函数,编译器无法从上下文推导出应该选择baz的哪个重载版本。模板参数推导需要明确的单一类型信息,而重载函数名本身不具备这种确定性。

解决方案

把重载的操作包装成多态函数对象,让编译器在调用时根据参数类型自动匹配对应的重载实现,以下是两种常用方案:

方案1:使用带重载operator()的结构体

定义一个结构体,在其中实现对应四种参数组合的operator():

struct BazOp {
    std::pair<double, double> operator()(std::pair<double, double> se, double vi) {
        se.first += vi;
        se.second *= 1 + vi;
        return se;
    }

    std::pair<double, double> operator()(double vi, std::pair<double, double> se) {
        return operator()(se, vi);
    }

    std::pair<double, double> operator()(double vj, double vi) {
        return std::pair(vi + vj, 1.0);
    }

    std::pair<double, double> operator()(std::pair<double, double> se1, std::pair<double, double> se2) {
        se1.first += se2.first;
        se1.second *= se2.second;
        return se1;
    }
};

在std::reduce中传递该结构体的实例即可:

auto bar = std::reduce(std::execution::par_unseq, v.begin(), v.end(), std::pair(0.0,1.0), BazOp{});

方案2:使用重载的lambda(C++17及以上)

C++17允许在lambda中定义多个重载的operator(),可以直接写出包含所有重载逻辑的lambda:

auto baz = [](auto&& a, auto&& b) -> std::pair<double, double> {
    using A = std::decay_t<decltype(a)>;
    using B = std::decay_t<decltype(b)>;

    if constexpr (std::is_same_v<A, std::pair<double, double>> && std::is_same_v<B, double>) {
        auto se = std::forward<decltype(a)>(a);
        auto vi = std::forward<decltype(b)>(b);
        se.first += vi;
        se.second *= 1 + vi;
        return se;
    } else if constexpr (std::is_same_v<A, double> && std::is_same_v<B, std::pair<double, double>>) {
        return baz(std::forward<decltype(b)>(b), std::forward<decltype(a)>(a));
    } else if constexpr (std::is_same_v<A, double> && std::is_same_v<B, double>) {
        return std::pair(std::forward<decltype(a)>(a) + std::forward<decltype(b)>(b), 1.0);
    } else if constexpr (std::is_same_v<A, std::pair<double, double>> && std::is_same_v<B, std::pair<double, double>>) {
        auto se1 = std::forward<decltype(a)>(a);
        auto se2 = std::forward<decltype(b)>(b);
        se1.first += se2.first;
        se1.second *= se2.second;
        return se1;
    } else {
        static_assert(std::disjunction_v<std::is_same<A, std::pair<double, double>>, std::is_same<A, double>>, "不支持的参数类型");
    }
};

直接传递该lambda给std::reduce即可:

auto bar = std::reduce(std::execution::par_unseq, v.begin(), v.end(), std::pair(0.0,1.0), baz);

完整可编译代码(方案1版本)

#include <algorithm>
#include <execution>
#include <iostream>
#include <random>
#include <utility>

#define N 100
#define seed 1

struct BazOp {
    std::pair<double, double> operator()(std::pair<double, double> se, double vi) {
        se.first += vi;
        se.second *= 1 + vi;
        return se;
    }

    std::pair<double, double> operator()(double vi, std::pair<double, double> se) {
        return operator()(se, vi);
    }

    std::pair<double, double> operator()(double vj, double vi) {
        return std::pair(vi + vj, 1.0);
    }

    std::pair<double, double> operator()(std::pair<double, double> se1, std::pair<double, double> se2) {
        se1.first += se2.first;
        se1.second *= se2.second;
        return se1;
    }
};

int main(){
    std::mt19937 rng(seed);
    std::uniform_real_distribution<double> dist(0,1);
    std::vector<double> v(N);

    std::for_each(std::execution::par_unseq, v.begin(), v.end(), [&dist,&rng](double& c){ c = dist(rng); });

    for(auto x : v)
        std::cout << x << ", ";
    std::cout << "\n";

    // 原accumulate代码
    auto foo = std::accumulate(v.begin(), v.end(), std::pair(0.0,1.0), 
        [](std::pair<double,double> se, double vi){
            se.first  += vi;
            se.second *= 1+vi;
            return se;
        });
    std::cout << foo.first << " " << foo.second << "\n";

    // 修改后的reduce代码
    auto bar = std::reduce(std::execution::par_unseq, v.begin(), v.end(), std::pair(0.0,1.0), BazOp{});
    std::cout << bar.first << " " << bar.second << "\n";
}

注意事项

  • 必须显式指定并行策略(如std::execution::par_unseq),否则std::reduce行为与std::accumulate类似,不保证并行执行。
  • 确保编译器支持C17或更高版本,std::execution库和重载lambda特性均依赖C17标准。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 13:45:15