自定义matrix数据结构:实现类似numpy的动态索引返回类型功能求助
解决方案
你的核心问题是同一个成员函数无法返回两种不同类型,要实现numpy式的索引行为,推荐使用代理类来适配两种返回场景,以下是具体实现方案:
实现思路
创建一个代理类,让operator()返回这个代理对象,代理类根据原matrix的维度(行数),支持隐式转换为标量引用(一维时)或一维matrix(二维时),同时支持赋值操作。
完整代码示例
#include <vector> #include <stdexcept> template<typename T> class matrix; // 代理类:适配一维/二维matrix的索引返回 template<typename T> class MatrixIndexProxy { private: matrix<T>& _mat; size_t _idx; public: MatrixIndexProxy(matrix<T>& mat, size_t idx) : _mat(mat), _idx(idx) {} // 一维matrix时,转换为标量引用 operator T&() { if (_mat.rows != 1) { throw std::runtime_error("无法将行代理转换为标量"); } if (_idx >= _mat.cols) { throw std::runtime_error("索引超出范围"); } return _mat._data[0][_idx]; } // 二维matrix时,转换为一维matrix(返回拷贝) operator matrix<T>() { if (_mat.rows == 1) { throw std::runtime_error("无法将标量代理转换为matrix"); } if (_idx >= _mat.rows) { throw std::runtime_error("索引超出范围"); } matrix<T> result(1, _mat.cols); for (size_t i = 0; i < _mat.cols; ++i) { result(0, i) = _mat._data[_idx][i]; } return result; } // 支持修改一维matrix的元素 MatrixIndexProxy& operator=(const T& val) { if (_mat.rows != 1) { throw std::runtime_error("无法为行代理赋值标量"); } if (_idx >= _mat.cols) { throw std::runtime_error("索引超出范围"); } _mat._data[0][_idx] = val; return *this; } // 支持将一维matrix赋值给二维matrix的某一行 MatrixIndexProxy& operator=(const matrix<T>& row) { if (_mat.rows == 1) { throw std::runtime_error("无法为标量代理赋值matrix"); } if (_idx >= _mat.rows || row.rows != 1 || row.cols != _mat.cols) { throw std::runtime_error("行赋值参数不合法"); } for (size_t i = 0; i < _mat.cols; ++i) { _mat._data[_idx][i] = row(0, i); } return *this; } }; template<typename T> class matrix { private: std::vector<std::vector<T>> _data; public: size_t rows; size_t cols; // 构造函数:创建rows行cols列的matrix matrix(size_t r, size_t c) : rows(r), cols(c), _data(r, std::vector<T>(c)) {} // 二维索引:访问指定行列的元素 T& operator()(size_t x, size_t y) { if (x >= rows || y >= cols) { throw std::runtime_error("索引超出范围"); } return _data[x][y]; } // 一维索引:返回代理对象 MatrixIndexProxy<T> operator()(size_t x) { if ((rows == 1 && x >= cols) || (rows != 1 && x >= rows)) { throw std::runtime_error("索引超出范围"); } return MatrixIndexProxy<T>(*this, x); } }; // 使用示例 int main() { // 一维matrix操作 matrix<int> mat1(1, 4); mat1(0,0) = 1; mat1(0,1) = 2; mat1(0,2) = 3; mat1(0,3) =4; int val = mat1(2); // 获取标量3 mat1(2) = 5; // 修改元素为5 // 二维matrix操作 matrix<int> mat2(2,3); mat2(0,0)=1; mat2(0,1)=2; mat2(0,2)=3; mat2(1,0)=4; mat2(1,1)=5; mat2(1,2)=6; matrix<int> row = mat2(0); // 获取第一行的一维matrix mat2(1) = row; // 将第一行赋值给第二行 return 0; }
关键说明
- 代理类
MatrixIndexProxy封装了原matrix的引用和索引,根据原matrix的行数自动适配转换逻辑 - 修复了你原代码中
result(1, i)的索引错误(matrix索引应从0开始) - 支持两种场景的赋值操作:修改一维matrix的元素,以及将一维matrix赋值给二维matrix的某一行
- 所有索引操作都会做越界检查,抛出明确的错误信息
内容的提问来源于stack exchange,提问作者Sayed Aulia
相关产品推荐
相关产品推荐

