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

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区分矩阵/向量/标量场景:

  1. 定义type_dimensions trait,提取每个类型的维度序列(比如Vec<T,N>是std::index_sequence<N>,Mat<T,N,M>是std::index_sequence<N,M>,标量是std::index_sequence<>)
  2. 定义unify_dimensions trait,检查所有参数的维度是否符合规则(矩阵存在时统一为N×M,向量存在时统一为N,标量无约束)
  3. 实现通用的extract函数,根据维度索引提取对应位置的元素(矩阵需要行+列索引,向量需要单索引,标量直接返回)
  4. 用重载+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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.27 09:49:36