如何利用模板类从花括号初始化列表创建任意n维数组(矩阵)
实现支持任意维度的C++矩阵类(类似MATLAB风格)
问题描述
现有如下C++代码,仅支持一维数组,希望能像MATLAB那样通过mat a = {{1, 2, 3}, {4, 5, 6}};这类嵌套初始化的方式创建任意n维数组,请问如何实现?
原代码:
#include<iostream> #include<tuple> using namespace std; template<typename T, std::size_t N> class mat { private: T ptr[N]; public: mat() = default; template <typename... Args> mat(Args&&... args) : ptr{std::forward<Args>(args)...} {} void print(){ for (int i=0; i<N; i++){ std::cout << (ptr+i) << ' ' << *(ptr+i) <<'\t'; } } T operator[](int j){ return *(ptr+j-1); } }; // deduction guide template<typename... Args> mat(Args&&...) -> mat<typename std::tuple_element<0, std::tuple<Args...>>::type,sizeof...(Args)>; // need to do -> multi dimensional array; int main() { mat a = {1, 2, 3, 4, 5, 6}; a.print(); }
解决方案
核心思路是通过递归模板嵌套实现多维矩阵:高维矩阵的元素是低一维的矩阵,直到最后一维退化为基础数据类型的数组。以下是完整实现:
完整代码
#include <iostream> #include <tuple> #include <type_traits> using namespace std; // 前向声明:用于递归的矩阵模板 template<typename T, std::size_t... Dims> class mat; // 终止递归:一维矩阵的特化 template<typename T, std::size_t N> class mat<T, N> { private: T ptr[N]; public: // 默认构造 mat() = default; // 接收基础类型参数的构造函数 template<typename... Args, typename = std::enable_if_t<(std::is_convertible_v<Args, T> && ...)>> mat(Args&&... args) : ptr{std::forward<Args>(args)...} {} // 打印一维矩阵 void print() const { for (std::size_t i = 0; i < N; ++i) { std::cout << ptr[i] << " "; } } // 下标运算符:返回元素值 T& operator[](std::size_t idx) { return ptr[idx]; } const T& operator[](std::size_t idx) const { return ptr[idx]; } }; // 递归定义:多维矩阵(维度数>1) template<typename T, std::size_t FirstDim, std::size_t... RestDims> class mat<T, FirstDim, RestDims...> { private: mat<T, RestDims...> sub_mat[FirstDim]; public: // 默认构造 mat() = default; // 接收低维矩阵参数的构造函数 template<typename... Args, typename = std::enable_if_t<std::is_same_v<std::decay_t<Args>, mat<T, RestDims...>> && ...>> mat(Args&&... args) : sub_mat{std::forward<Args>(args)...} {} // 递归打印多维矩阵:先打印每个子矩阵,换行分隔 void print() const { for (std::size_t i = 0; i < FirstDim; ++i) { sub_mat[i].print(); std::cout << "\n"; } } // 下标运算符:返回低维矩阵的引用 mat<T, RestDims...>& operator[](std::size_t idx) { return sub_mat[idx]; } const mat<T, RestDims...>& operator[](std::size_t idx) const { return sub_mat[idx]; } }; // 推导指南1:处理嵌套初始化列表(二维及以上) template<typename T, typename... Args> mat(mat<T, Args...>...) -> mat<T, sizeof...(Args)+1, Args...>; // 推导指南2:处理一维初始化列表 template<typename... Args> mat(Args&&...) -> mat<std::common_type_t<Args...>, sizeof...(Args)>; int main() { // 一维矩阵 mat a1 = {1, 2, 3, 4}; cout << "一维矩阵:\n"; a1.print(); cout << "\n\n"; // 二维矩阵 mat a2 = {{1, 2, 3}, {4, 5, 6}, {7, 8, 9}}; cout << "二维矩阵:\n"; a2.print(); cout << "\n\n"; // 三维矩阵 mat a3 = {{{1,2}, {3,4}}, {{5,6}, {7,8}}}; cout << "三维矩阵:\n"; a3.print(); // 下标访问示例 cout << "\n\n访问二维矩阵a2[1][2]:" << a2[1][2] << endl; return 0; }
关键实现说明
- 递归模板结构:用特化的
mat<T, N>作为一维终止版本,mat<T, FirstDim, RestDims...>作为多维递归版本,每个多维矩阵的元素是低一维的mat对象。 - 构造函数约束:用
std::enable_if确保构造函数只接收匹配类型的参数,避免歧义。 - 推导指南:两个推导指南分别处理一维和多维的初始化列表,让编译器能自动推导元素类型和各维度尺寸。
- 递归打印:多维矩阵的
print会递归调用子矩阵的print,自动添加换行区分不同维度的层级。 - 下标运算符:多维场景下返回低维矩阵的引用,支持链式访问(比如
a2[1][2])。
内容的提问来源于stack exchange,提问作者rupsagar
相关产品推荐
相关产品推荐

