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

如何为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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 08:19:53