如何实现类似Kokkos View的多维数组?模板参数处理问询
实现类似Kokkos View的多维数组组件:模板参数处理方案
要实现能区分编译期固定维度和运行时可变维度的多维数组组件,核心是通过模板元编程解析类型参数里的星号(*)和方括号([])信息,具体可按以下步骤处理:
1. 递归解析类型,统计维度信息
通过模板特化递归拆解类型参数,分别统计运行时维度数(星号数量)、编译期维度数(方括号数量),同时提取编译期维度的大小:
核心模板解析结构示例
// 基础模板:匹配待解析的复合类型 template <typename T> struct ViewTraits; // 特化:处理运行时维度(星号) template <typename T> struct ViewTraits<T*> { using BaseTraits = ViewTraits<T>; static constexpr size_t RuntimeDims = BaseTraits::RuntimeDims + 1; static constexpr size_t CompileTimeDims = BaseTraits::CompileTimeDims; using CompileTimeSizes = typename BaseTraits::CompileTimeSizes; using ValueType = typename BaseTraits::ValueType; }; // 特化:处理编译期维度(方括号) template <typename T, size_t N> struct ViewTraits<T[N]> { using BaseTraits = ViewTraits<T>; static constexpr size_t RuntimeDims = BaseTraits::RuntimeDims; static constexpr size_t CompileTimeDims = BaseTraits::CompileTimeDims + 1; // 将当前编译期维度大小N加入尺寸列表 using CompileTimeSizes = std::integer_sequence<size_t, N, typename BaseTraits::CompileTimeSizes::value...>; using ValueType = typename BaseTraits::ValueType; }; // 终止特化:匹配最底层的元素类型(如double) template <typename T> struct ViewTraits { static constexpr size_t RuntimeDims = 0; static constexpr size_t CompileTimeDims = 0; using CompileTimeSizes = std::integer_sequence<size_t>; using ValueType = T; };
2. 基于解析结果构建View类
利用ViewTraits解析出的维度数据,在View类中处理构造逻辑:
- 运行时维度大小通过构造函数参数传入
- 编译期维度大小直接从
CompileTimeSizes中获取
View类简化实现示例
template <typename T> class View { private: using Traits = ViewTraits<T>; std::array<size_t, Traits::CompileTimeDims + Traits::RuntimeDims> m_dims; typename Traits::ValueType* m_data; public: // 构造函数:仅接收与运行时维度数匹配的参数 template <typename... Args, typename = std::enable_if_t<sizeof...(Args) == Traits::RuntimeDims>> View(Args... runtime_sizes) { // 填充编译期维度大小 size_t idx = 0; []<size_t... Cs>(std::index_sequence<Cs...>, auto& self) { ((self.m_dims[self.CompileTimeDims - Cs - 1] = Cs), ...); }(typename Traits::CompileTimeSizes{}, *this); // 填充运行时维度大小 ((m_dims[Traits::CompileTimeDims + idx++] = runtime_sizes), ...); // 计算总元素数并分配内存 size_t total = 1; for (auto d : m_dims) total *= d; m_data = new typename Traits::ValueType[total]; } ~View() { delete[] m_data; } // 可添加元素访问、维度查询等成员函数 };
3. 关键细节说明
- 递归解析顺序:类型中的星号和方括号从右到左逐层拆解,比如
double*[N1][N2]会先解析[N2],再解析[N1],最后解析*,需注意调整编译期尺寸的存储顺序以匹配用户的维度逻辑。 - 编译期校验:用
std::enable_if_t确保构造函数接收的参数数量与运行时维度数一致,避免传参错误。 - 扩展性:无论星号和方括号的组合顺序如何(如
double*[N][M]***),递归模板都能正确解析所有维度信息。
内容的提问来源于stack exchange,提问作者aaronfu
相关产品推荐
相关产品推荐

