C++模板推导失败:std::enable_if添加等值校验时无匹配函数
问题背景
我正在对不同的求和实现方式开展基准测试,期望使用如下形式的调用接口:
avx2_sum<sum_algorithm::normal>(container.begin(), container.end());
当前编写的实现代码如下:
enum class sum_algorithm: char{ normal, kahan, twofold_fast }; template<sum_algorithm algorithm_t, typename iterator_t, typename sum_t = typename std::iterator_traits<iterator_t>::value_type, std::enable_if_t<std::is_same<sum_t, double>::value && (algorithm_t == sum_algorithm::normal)> = true> sum_t avx2_sum(const iterator_t begin, const iterator_t end) noexcept { // SIMD并行求和阶段 auto running_sums = _mm256_set1_pd(0); auto iterator_skip = 256/sizeof(sum_t); for (iterator_t it = begin; it + iterator_skip < end; it += iterator_skip){ //TODO: 切换为双加载归约实现 running_sums = _mm256_add_pd(_mm256_load_pd(it), running_sums); } // 串行求和收尾 running_sums = _mm256_hadd_pd(running_sums, running_sums); running_sums = _mm256_hadd_pd(running_sums, running_sums); return _mm256_cvtsd_f64(running_sums); }
编译时产生如下报错:
error: no matching function for call to 'avx2_sum' std::cout << "avx2<float, normal>: " << accumulators::avx2_sum<algo::normal>(float_arr.begin(), float_arr.end()) <<"\n"; ^~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ /home/rlfactory/dev/thommmj1/cppbenchmarks/cpp_utils/algorithms/cpu_accumulators.hpp:55:11: note: candidate template ignored: requirement 'std::is_same<float, double>::value' was not satisfied [with algorithm_t = accumulators::sum_algorithm::normal, iterator_t = __gnu_cxx::__normal_iterator<float *, std::vector<float, std::allocator<float> > >, sum_t = float] sum_t avx2_sum(const iterator_t begin, const iterator_t end) noexcept { ^ /home/rlfactory/dev/thommmj1/cppbenchmarks/cpp_utils/algorithms/cpu_accumulators.hpp:71:11: note: candidate template ignored: substitution failure [with algorithm_t = accumulators::sum_algorithm::normal, iterator_t = __gnu_cxx::__normal_iterator<float *, std::vector<float, std::allocator<float> > >, sum_t = float]: a non-type template parameter cannot have type 'std::enable_if_t<std::is_same<float, float>::value && ((sum_algorithm)'\x00' == sum_algorithm::normal)>' (aka 'void') sum_t avx2_sum(const iterator_t begin, const iterator_t end) noexcept { ^ /home/rlfactory/dev/thommmj1/cppbenchmarks/cpp_utils/algorithms/cpu_accumulators.hpp:87:11: note: candidate template ignored: requirement 'std::is_same<float, double>::value' was not satisfied [with algorithm_t = accumulators::sum_algorithm::normal, iterator_t = __gnu_cxx::__normal_iterator<float *, std::vector<float, std::allocator<float> > >, sum_t = float] sum_t avx2_sum(const iterator_t begin, const iterator_t end) noexcept { ^ /home/rlfactory/dev/thommmj1/cppbenchmarks/cpp_utils/algorithms/cpu_accumulators.hpp:109:11: note: candidate template ignored: requirement 'std::is_same<float, float>::value && ((accumulators::sum_algorithm)'\x00' == sum_algorithm::kahan)' was not satisfied [with algorithm_t = accumulators::sum_algorithm::normal, iterator_t = __gnu_cxx::__normal_iterator<float *, std::vector<float, std::allocator<float> > >, sum_t = float] sum_t avx2_sum(const iterator_t begin, const iterator_t end) noexcept {
移除algorithm_t模板参数及对应校验逻辑时,代码可正常编译运行,不确定该问题是否与using algo = accumulators::sum_algorithm;的别名声明有关。
问题根因
该问题和algo别名声明完全无关,报错来自两个代码错误:
- SFINAE写法非法:你写的
std::enable_if_t<...> = true存在语法问题,enable_if_t在条件成立时默认返回void类型,你试图给void类型的非类型模板参数赋值true,直接触发替换失败,对应报错里的a non-type template parameter cannot have type 'void'。 - 缺少对应重载:当前实现只编写了
sum_t为double类型的normal算法版本,调用时传入的是float数组,没有匹配的重载版本,自然报无匹配函数错误。
修复方案
- 修正SFINAE的写法:给
enable_if_t传入第二个参数指定合法的非类型参数类型,通用写法是指定为int,默认值设为0即可。 - 补全
float类型、其他算法的对应重载,匹配不同的调用场景。
修正后的模板声明参考:
// double类型normal算法版本,修正SFINAE写法 template<sum_algorithm algorithm_t, typename iterator_t, typename sum_t = typename std::iterator_traits<iterator_t>::value_type, // enable_if_t第二个参数指定为int,避免void类型错误 std::enable_if_t<std::is_same_v<sum_t, double> && (algorithm_t == sum_algorithm::normal), int> = 0> sum_t avx2_sum(const iterator_t begin, const iterator_t end) noexcept { auto running_sums = _mm256_set1_pd(0); // 步长是编译期常量,加constexpr提升性能 constexpr auto iterator_skip = 256/sizeof(sum_t); for (iterator_t it = begin; it + iterator_skip < end; it += iterator_skip){ running_sums = _mm256_add_pd(_mm256_load_pd(it), running_sums); } running_sums = _mm256_hadd_pd(running_sums, running_sums); running_sums = _mm256_hadd_pd(running_sums, running_sums); return _mm256_cvtsd_f64(running_sums); } // 新增float类型normal算法版本,匹配float数组调用 template<sum_algorithm algorithm_t, typename iterator_t, typename sum_t = typename std::iterator_traits<iterator_t>::value_type, std::enable_if_t<std::is_same_v<sum_t, float> && (algorithm_t == sum_algorithm::normal), int> = 0> sum_t avx2_sum(const iterator_t begin, const iterator_t end) noexcept { // 单精度浮点使用_mm256_*_ps系列指令 auto running_sums = _mm256_set1_ps(0); constexpr auto iterator_skip = 256/sizeof(sum_t); for (iterator_t it = begin; it + iterator_skip < end; it += iterator_skip){ running_sums = _mm256_add_ps(_mm256_load_ps(it), running_sums); } running_sums = _mm256_hadd_ps(running_sums, running_sums); running_sums = _mm256_hadd_ps(running_sums, running_sums); return _mm256_cvtss_f32(running_sums); }
额外注意两个实现bug:
- 当前循环边界判断
it + iterator_skip < end会漏掉末尾不足一个向量长度的元素,SIMD循环结束后需要补充标量累加逻辑处理剩余元素,否则计算结果错误。 _mm256_load_pd/ps要求传入指针32字节对齐,如果传入的容器内存不保证对齐,要替换为_mm256_loadu_pd/ps非对齐加载指令,否则会触发内存访问错误。
内容的提问来源于stack exchange,提问作者M46f988b814
相关产品推荐
相关产品推荐

