C++中扩展多模式可变参数模板以支持Mat类
适配矩阵与向量的多模式apply函数扩展方案
问题背景
这是之前《C++中的多模式可变参数模板》问题的后续,原方案已经能处理Vec向量类的批量函数应用,但现在新增了mat<T, N, M>矩阵类(N行M列,Vec等价于M维列向量用于广播),需要扩展apply函数满足以下规则:
- 若参数列表包含N×M矩阵:所有矩阵必须是N×M,所有向量必须是M维,返回N×M矩阵
- 否则遵循原向量规则:若参数含N维向量,所有向量必须N维,返回N维向量
- 还要保证原
Vec的处理逻辑不受影响,同时为后续支持tensor<T, N...>预留扩展空间,解决新旧版本apply共存不冲突的问题
原解决方案代码
#include <array> #include <iostream> #include <type_traits> template <typename T, std::size_t N> class Vec; template <std::size_t, typename ...> struct dimVec; // ground case for no Vecs: unimplemented for SFINAE failure ! template <> struct dimVec<0U>; // ground case with one or more Vecs: size fixed template <std::size_t N> struct dimVec<N> : public std::integral_constant<std::size_t, N> { }; // first Vec: size detected template <std::size_t N, typename T, typename ... Ts> struct dimVec<0U, Vec<T, N>, Ts...> : public dimVec<N, Ts...> { }; // another Vec of same size: continue template <std::size_t N, typename T, typename ... Ts> struct dimVec<N, Vec<T, N>, Ts...> : public dimVec<N, Ts...> { }; // another Vec of different size: unimplemented for SFINAE failure ! template <std::size_t N1, std::size_t N2, typename T, typename ... Ts> struct dimVec<N1, Vec<T, N2>, Ts...>; // a not-Vec type: continue template <std::size_t N, typename T, typename ... Ts> struct dimVec<N, T, Ts...> : public dimVec<N, Ts...> { }; template <typename ... Args> static constexpr auto dimVecV { dimVec<0U, Args...>::value }; template <std::size_t I, typename T, std::size_t N> constexpr auto extrV (Vec<T, N> const & v) { return v[I]; } template <std::size_t I, typename T> constexpr auto extrV (T const & v) { return v; } template <typename T, std::size_t N> class Vec { private: std::array<T, N> d; public: template <typename ... Ts> Vec (Ts ... ts) : d{{ ts... }} { } T & operator[] (int i) { return d[i]; } T const & operator[] (int i) const { return d[i]; } }; template <std::size_t I, typename F, typename ... Args> auto applyH2 (F && f, Args ... as) { return f(extrV<I>(as)...); } template <std::size_t ... Is, typename F, typename ... Args> auto applyH1 (std::index_sequence<Is...> const &, F && f, Args ... as) -> Vec<decltype(applyH2<0U>(f, as...)), sizeof...(Is)> { return { applyH2<Is>(f, as...)... }; } template <typename F, typename ... Args, std::size_t N = dimVecV<Args...>> auto apply (F && f, Args ... as) { return applyH1(std::make_index_sequence<N>{}, f, as...); } long foo (int a, int b) { return a + b + 42; } int main () { Vec<int, 3U> v3; Vec<int, 2U> v2; auto r1 { apply(foo, v2, v2) }; auto r2 { apply(foo, v3, v3) }; auto r3 { apply(foo, v3, 0) }; static_assert( std::is_same<decltype(r1), Vec<long, 2U>>{}, "!" ); static_assert( std::is_same<decltype(r2), Vec<long, 3U>>{}, "!" ); static_assert( std::is_same<decltype(r3), Vec<long, 3U>>{}, "!" ); // apply(foo, v2, v3); // compilation error // apply(foo, 1, 2); // compilation error }
解决方案思路
咱们核心是用类型 trait 统一提取维度信息,再通过SFINAE区分矩阵/向量/标量场景:
- 定义
type_dimensionstrait,提取每个类型的维度序列(比如Vec<T,N>是std::index_sequence<N>,Mat<T,N,M>是std::index_sequence<N,M>,标量是std::index_sequence<>) - 定义
unify_dimensionstrait,检查所有参数的维度是否符合规则(矩阵存在时统一为N×M,向量存在时统一为N,标量无约束) - 实现通用的
extract函数,根据维度索引提取对应位置的元素(矩阵需要行+列索引,向量需要单索引,标量直接返回) - 用重载+
std::enable_if区分不同输出类型的apply函数,保证原Vec逻辑和新Mat逻辑共存
适配后的完整代码
#include <array> #include <iostream> #include <type_traits> #include <utility> // 前置声明 template <typename T, std::size_t N> class Vec; template <typename T, std::size_t N, std::size_t M> class Mat; // 提取类型的维度序列 trait template <typename T> struct type_dimensions; // 标量类型:空维度序列 template <typename T> requires (!std::is_same_v<T, Vec<typename T::value_type, T::size>> && !std::is_same_v<T, Mat<typename T::value_type, T::rows, T::cols>>) struct type_dimensions<T> { using type = std::index_sequence<>; }; // Vec类型:单维度序列 template <typename T, std::size_t N> struct type_dimensions<Vec<T, N>> { using type = std::index_sequence<N>; }; // Mat类型:二维维度序列 template <typename T, std::size_t N, std::size_t M> struct type_dimensions<Mat<T, N, M>> { using type = std::index_sequence<N, M>; }; template <typename T> using type_dimensions_t = typename type_dimensions<T>::type; // 统一维度序列的 trait:检查所有参数维度是否兼容 template <typename... Dims> struct unify_dimensions; // 基础情况:无参数,空维度 template <> struct unify_dimensions<> { using type = std::index_sequence<>; }; // 合并新维度与已有统一维度 template <typename UnifiedDims, typename NewDims> struct merge_dimensions; // 已有统一维度为空,直接用新维度 template <typename... NewIs> struct merge_dimensions<std::index_sequence<>, std::index_sequence<NewIs...>> { using type = std::index_sequence<NewIs...>; }; // 已有统一维度是向量(1维),新维度是标量或同维度向量 template <std::size_t U, typename... NewIs> struct merge_dimensions<std::index_sequence<U>, std::index_sequence<NewIs...>> { static_assert(sizeof...(NewIs) == 0 || (sizeof...(NewIs) == 1 && NewIs... == U), "Vector dimension mismatch!"); using type = std::index_sequence<U>; }; // 已有统一维度是矩阵(2维),新维度是标量、同列向量或同维度矩阵 template <std::size_t URow, std::size_t UCol, typename... NewIs> struct merge_dimensions<std::index_sequence<URow, UCol>, std::index_sequence<NewIs...>> { static_assert(sizeof...(NewIs) == 0 || (sizeof...(NewIs) == 1 && NewIs... == UCol) || (sizeof...(NewIs) == 2 && NewIs... == std::pair{URow, UCol}), "Matrix/Vector dimension mismatch!"); using type = std::index_sequence<URow, UCol>; }; // 递归统一所有参数的维度 template <typename UnifiedDims, typename First, typename... Rest> struct unify_dimensions<UnifiedDims, First, Rest...> { using merged = typename merge_dimensions<UnifiedDims, type_dimensions_t<First>>::type; using type = typename unify_dimensions<merged, Rest...>::type; }; template <typename... Args> using unify_dimensions_t = typename unify_dimensions<std::index_sequence<>, Args...>::type; // 元素提取函数 // 标量提取:直接返回 template <typename... Is, typename T> constexpr auto extract(std::index_sequence<Is...>, const T& val) { return val; } // Vec提取:按索引取元素 template <std::size_t I, typename T, std::size_t N> constexpr auto extract(std::index_sequence<I>, const Vec<T, N>& v) { return v[I]; } // Mat提取:按行+列索引取元素 template <std::size_t Row, std::size_t Col, typename T, std::size_t N, std::size_t M> constexpr auto extract(std::index_sequence<Row, Col>, const Mat<T, N, M>& mat) { return mat[Row][Col]; } // 生成所有索引组合的工具(用于矩阵的行+列遍历) template <typename DimSeq> struct generate_indices; // 向量维度:生成0..N-1的单索引 template <std::size_t N> struct generate_indices<std::index_sequence<N>> { using type = std::make_index_sequence<N>; }; // 矩阵维度:生成所有(Row, Col)对的索引序列 template <std::size_t N, std::size_t M> struct generate_indices<std::index_sequence<N, M>> { using row_indices = std::make_index_sequence<N>; using col_indices = std::make_index_sequence<M>; // 把(Row, Col)转为std::index_sequence<Row, Col> template <std::size_t Row, std::size_t Col> using index_pair = std::index_sequence<Row, Col>; // 生成所有行对应的列索引组合 template <std::size_t... Rows> static auto generate_row_pairs(std::index_sequence<Rows...>) { return std::tuple_cat(std::make_tuple(index_pair<Rows, Cols>())...); } using type = decltype(generate_row_pairs(row_indices{})); }; // Vec类实现(保持原逻辑) template <typename T, std::size_t N> class Vec { public: using value_type = T; static constexpr std::size_t size = N; std::array<T, N> d; template <typename... Ts> Vec(Ts... ts) : d{{ts...}} {} T& operator[](int i) { return d[i]; } const T& operator[](int i) const { return d[i]; } }; // Mat类实现 template <typename T, std::size_t N, std::size_t M> class Mat { public: using value_type = T; static constexpr std::size_t rows = N; static constexpr std::size_t cols = M; std::array<Vec<T, M>, N> d; template <typename... Rows> Mat(Rows... rows) : d{{rows...}} {} Vec<T, M>& operator[](int i) { return d[i]; } const Vec<T, M>& operator[](int i) const { return d[i]; } }; // apply的底层实现:处理单个索引组合 template <typename IndexSeq, typename F, typename... Args> auto apply_impl(F&& f, const Args&... args) { return f(extract(IndexSeq{}, args)...); } // apply重载1:返回Vec(当统一维度是1维时) template <typename F, typename... Args, std::size_t N> auto apply(F&& f, const Args&... args) -> Vec<decltype(apply_impl<std::index_sequence<0>>(f, args...)), N> requires std::is_same_v<unify_dimensions_t<Args...>, std::index_sequence<N>> { using Indices = typename generate_indices<std::index_sequence<N>>::type; auto build_vec = [&]<std::size_t... Is>(std::index_sequence<Is...>) { return Vec<decltype(apply_impl<std::index_sequence<0>>(f, args...)), N>{ apply_impl<std::index_sequence<Is>>(f, args...)... }; }; return build_vec(Indices{}); } // apply重载2:返回Mat(当统一维度是2维时) template <typename F, typename... Args, std::size_t N, std::size_t M> auto apply(F&& f, const Args&... args) -> Mat<decltype(apply_impl<std::index_sequence<0,0>>(f, args...)), N, M> requires std::is_same_v<unify_dimensions_t<Args...>, std::index_sequence<N, M>> { using IndexPairs = typename generate_indices<std::index_sequence<N, M>>::type; auto build_mat = [&]<std::size_t... Rows>(std::index_sequence<Rows...>) { auto build_row = [&]<std::size_t... Cols>(std::index_sequence<Cols...>) { return Vec<decltype(apply_impl<std::index_sequence<0,0>>(f, args...)), M>{ apply_impl<std::index_sequence<Rows, Cols>>(f, args...)... }; }; return Mat<decltype(apply_impl<std::index_sequence<0,0>>(f, args...)), N, M>{ build_row(std::make_index_sequence<M>{})... }; }; return build_mat(std::make_index_sequence<N>{}); } // 测试函数 long foo(int a, int b) { return a + b + 42; } int main() { // Vec测试(保持原逻辑) Vec<int, 3> v3{1,2,3}; Vec<int, 2> v2{10,20}; auto r1 = apply(foo, v2, v2); auto r2 = apply(foo, v3, 0); static_assert(std::is_same_v<decltype(r1), Vec<long, 2>>); static_assert(std::is_same_v<decltype(r2), Vec<long, 3>>); // Mat测试 Mat<int, 2, 3> mat{ Vec<int,3>{1,2,3}, Vec<int,3>{4,5,6} }; Vec<int,3> vec_col{10,20,30}; // 矩阵+矩阵 auto mat_mat = apply(foo, mat, mat); static_assert(std::is_same_v<decltype(mat_mat), Mat<long, 2, 3>>); // 矩阵+列向量(广播) auto mat_vec = apply(foo, mat, vec_col); static_assert(std::is_same_v<decltype(mat_vec), Mat<long, 2, 3>>); // 矩阵+标量(广播) auto mat_scalar = apply(foo, mat, 100); static_assert(std::is_same_v<decltype(mat_scalar), Mat<long, 2, 3>>); // 错误测试(编译报错,符合预期) // Mat<int,2,2> bad_mat; // apply(foo, mat, bad_mat); // 维度不匹配 // apply(foo, mat, v2); // 向量维度不匹配 return 0; }
扩展说明
这个方案天生支持后续的tensor<T, N...>类:只需要给tensor特化type_dimensions trait,再扩展merge_dimensions的特化逻辑,添加张量维度的兼容规则,就能无缝接入现有的apply函数,不需要修改核心逻辑。
内容的提问来源于stack exchange,提问作者Glen Fletcher
相关产品推荐
相关产品推荐

