如何实现支持重载操作的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
相关产品推荐
相关产品推荐

