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

如何利用模板类从花括号初始化列表创建任意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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 17:25:30