如何为C++自定义Matrix类实现类似NumPy的多维度索引?
实现NumPy风格的C++ matrix类:支持一维单括号、二维双括号及线性索引
要实现你需要的索引行为,核心问题是解决二维矩阵同时支持双括号[i][j]和单括号线性索引[k]的矛盾——因为C++的operator[]只能接受一个参数,直接返回行引用会导致线性索引越界,直接返回元素又会丢失双括号索引能力。这里用代理类来解决这个问题,以下是具体修改方案:
完整修改代码
#include <iostream> #include <vector> #include <stdexcept> namespace np { template <typename T> class matrix; // 代理类:处理两种索引逻辑 template <typename T> class matrix_proxy { private: matrix<T>& parent; int idx; public: matrix_proxy(matrix<T>& mat, int index) : parent(mat), idx(index) {} // 双括号索引:代理对象后接[col],返回行列对应元素 T& operator[](int col) { if (parent.get_rows() == 1) { throw std::runtime_error("1D matrix does not support double indexing."); } if (idx < 0 || idx >= parent.get_rows()) { throw std::out_of_range("Row index out of bounds."); } if (col < 0 || col >= parent.get_cols()) { throw std::out_of_range("Column index out of bounds."); } return parent.get_mat()[idx][col]; } // 隐式转换为T&:直接使用代理对象时,视为线性索引 operator T&() { if (parent.get_rows() == 1) { if (idx < 0 || idx >= parent.get_cols()) { throw std::out_of_range("Index out of bounds."); } return parent.get_mat()[0][idx]; } else { int total = parent.get_rows() * parent.get_cols(); if (idx < 0 || idx >= total) { throw std::out_of_range("Linear index out of bounds."); } int row = idx / parent.get_cols(); int col = idx % parent.get_cols(); return parent.get_mat()[row][col]; } } // 常量版本 const T& operator[](int col) const { if (parent.get_rows() == 1) { throw std::runtime_error("1D matrix does not support double indexing."); } if (idx < 0 || idx >= parent.get_rows()) { throw std::out_of_range("Row index out of bounds."); } if (col < 0 || col >= parent.get_cols()) { throw std::out_of_range("Column index out of bounds."); } return parent.get_mat()[idx][col]; } operator const T&() const { if (parent.get_rows() == 1) { if (idx < 0 || idx >= parent.get_cols()) { throw std::out_of_range("Index out of bounds."); } return parent.get_mat()[0][idx]; } else { int total = parent.get_rows() * parent.get_cols(); if (idx < 0 || idx >= total) { throw std::out_of_range("Linear index out of bounds."); } int row = idx / parent.get_cols(); int col = idx % parent.get_cols(); return parent.get_mat()[row][col]; } } }; template <typename T> class matrix { private: std::vector<std::vector<T>> mat; int rows; int cols; // 给代理类提供安全访问接口 std::vector<std::vector<T>>& get_mat() { return mat; } const std::vector<std::vector<T>>& get_mat() const { return mat; } int get_rows() const { return rows; } int get_cols() const { return cols; } friend class matrix_proxy<T>; public: matrix(int numRows, int numCols) : rows(numRows), cols(numCols) { if (numRows <= 0 || numCols <= 0) { throw std::invalid_argument("Matrix dimensions must be positive."); } mat.resize(rows, std::vector<T>(cols, T())); } matrix(std::initializer_list<T> initlist) : rows(1), cols(initlist.size()) { if (cols == 0) { throw std::invalid_argument("1D matrix cannot be empty."); } mat.emplace_back(initlist); } matrix(std::initializer_list<std::initializer_list<T>> initlist) : rows(initlist.size()) { if (rows == 0) { throw std::invalid_argument("2D matrix cannot be empty."); } cols = initlist.begin()->size(); for (const auto& row : initlist) { if (row.size() != cols) { throw std::runtime_error("Initializer list is not rectangular."); } mat.emplace_back(row); } } // 返回代理类,统一处理索引逻辑 matrix_proxy<T> operator[](int index) { return matrix_proxy<T>(*this, index); } const matrix_proxy<T> operator[](int index) const { return matrix_proxy<T>(const_cast<matrix<T>&>(*this), index); } // 可选:支持多参数索引(NumPy风格的mat(row, col)) T& operator()(int row, int col) { if (row < 0 || row >= rows || col < 0 || col >= cols) { throw std::out_of_range("Index out of bounds."); } return mat[row][col]; } const T& operator()(int row, int col) const { if (row < 0 || row >= rows || col < 0 || col >= cols) { throw std::out_of_range("Index out of bounds."); } return mat[row][col]; } }; }; int main() { try { np::matrix<int> mat1 = {{1, 2, 3}, {4, 5, 6}}; np::matrix<int> mat2 = {1, 2, 3, 4, 5, 6}; // 二维矩阵双括号索引 std::cout << mat1[1][2] << std::endl; // 输出6 // 二维矩阵线性索引(行优先) std::cout << mat1[4] << std::endl; // 输出5 // 一维矩阵单括号索引 std::cout << mat2[3] << std::endl; // 输出4 // 可选:多参数索引 std::cout << mat1(1,1) << std::endl; // 输出5 } catch (const std::exception& e) { std::cerr << "Error: " << e.what() << std::endl; return 1; } return 0; }
核心修改说明
1. 代理类的作用
matrix_proxy作为中间层,承接matrix::operator[]的返回值,处理两种索引场景:
- 当用户写
mat1[1][2]时,第一个[1]返回代理对象,第二个[2]调用代理类的operator[],返回第1行第2列的元素 - 当用户写
mat1[4]时,代理对象被隐式转换为T&,此时按行优先计算线性索引对应的行列(4 = 1*3 + 1,即第1行第1列,值为5)
2. 封装与安全
- 通过私有getter函数让代理类访问matrix的内部数据,避免直接暴露私有成员
- 所有索引操作添加了越界检查,抛出标准异常,避免未定义行为
- 构造函数增加合法性校验,禁止空矩阵、非矩形初始化列表等非法输入
3. 扩展能力
额外添加了operator()重载,支持mat(row, col)的多参数索引方式,这也是NumPy中常用的索引写法,进一步贴近目标行为。
测试效果
运行main函数后,会输出:
6 5 4 5
完全符合你预期的索引行为。
内容的提问来源于stack exchange,提问作者Sayed Aulia
相关产品推荐
相关产品推荐

