Matrix类动态内存分配与运算符重载实现问题咨询
C++ Matrix类实现修复方案
原有代码问题汇总
- 索引计算错误:
getValue和setValue中二维转一维的索引公式错误,行优先存储的正确索引应为行号 * 总列数 + 列号,原有r*c + c逻辑仅当所有行号一致时才生效,完全不符合矩阵存储规则 - 内存管理漏洞:拷贝构造函数为浅拷贝,多个对象会共享同一块
mData内存,销毁时会出现重复释放问题;析构函数未释放动态申请的mData内存,存在内存泄漏 - 运算符重载逻辑完全不符合需求:未对结果矩阵分配内存就直接访问元素、未实现逐元素运算、返回值类型不匹配、维度校验规则错误
修复后完整代码
#include <stdint.h> #include <stdexcept> class Mat { private: uint16_t mRows; uint16_t mCols; double * mData; public: // 普通构造函数 Mat(uint16_t r, uint16_t c) : mRows(r), mCols(c) { if (r == 0 || c == 0) throw std::invalid_argument("矩阵行列数不能为0"); mData = new double[mRows * mCols](); // 括号初始化默认值为0 } // 行向量构造函数 Mat(uint16_t c) : mRows(1), mCols(c) { if (c == 0) throw std::invalid_argument("向量长度不能为0"); mData = new double[mCols](); } // 空构造函数 Mat() : mRows(0), mCols(0), mData(nullptr) {} // 深拷贝构造函数 Mat(const Mat &mat) : mRows(mat.mRows), mCols(mat.mCols) { if (mRows > 0 && mCols > 0) { mData = new double[mRows * mCols]; for (int i = 0; i < mRows * mCols; ++i) { mData[i] = mat.mData[i]; } } else { mData = nullptr; } } // 赋值运算符重载(三五法则补充) Mat& operator=(const Mat &mat) { if (this == &mat) return *this; delete[] mData; mRows = mat.mRows; mCols = mat.mCols; if (mRows > 0 && mCols > 0) { mData = new double[mRows * mCols]; for (int i = 0; i < mRows * mCols; ++i) { mData[i] = mat.mData[i]; } } else { mData = nullptr; } return *this; } // 元素读取 double getValue(uint16_t r, uint16_t c) const { if (r >= mRows || c >= mCols) throw std::out_of_range("矩阵下标越界"); return mData[r * mCols + c]; } // 元素写入 void setValue(uint16_t r, uint16_t c, double value) { if (r >= mRows || c >= mCols) throw std::out_of_range("矩阵下标越界"); mData[r * mCols + c] = value; } // 析构函数 ~Mat() { delete[] mData; } // 矩阵逐元素加法 Mat operator + (const Mat &mat) const { if (mat.mRows != mRows || mat.mCols != mCols) { throw std::invalid_argument("逐元素加法要求两个矩阵行列数完全一致"); } Mat result(mRows, mCols); for (int i = 0; i < mRows * mCols; ++i) { result.mData[i] = mData[i] + mat.mData[i]; } return result; } // 矩阵逐元素乘法 Mat operator * (const Mat &mat) const { if (mat.mRows != mRows || mat.mCols != mCols) { throw std::invalid_argument("逐元素乘法要求两个矩阵行列数完全一致"); } Mat result(mRows, mCols); for (int i = 0; i < mRows * mCols; ++i) { result.mData[i] = mData[i] * mat.mData[i]; } return result; } // 标准矩阵乘法 Mat operator ^ (const Mat &mat) const { if (mCols != mat.mRows) { throw std::invalid_argument("矩阵乘法要求左矩阵列数等于右矩阵行数"); } Mat result(mRows, mat.mCols); for (uint16_t i = 0; i < mRows; ++i) { for (uint16_t k = 0; k < mCols; ++k) { double a = mData[i * mCols + k]; if (a == 0) continue; for (uint16_t j = 0; j < mat.mCols; ++j) { result.mData[i * mat.mCols + j] += a * mat.mData[k * mat.mCols + j]; } } } return result; } };
实现说明
- 下标默认从0开始计数,如果需要从1开始计数,可以在
getValue和setValue中对入参r、c各减1后再计算索引 - 新增了参数校验逻辑,非法操作会抛出标准异常方便调试,也可根据需求改成返回错误码
- 补充了赋值运算符重载,符合C++三五法则要求,避免浅拷贝导致的内存问题
内容的提问来源于stack exchange,提问作者Stefan Daniel
相关产品推荐
相关产品推荐

