如何设计支持行列正反迭代的STL风格C++矩阵迭代器?
Got it, let's work through this problem step by step. The core goal here is to create an iterator whose increment behavior adapts to whether it's traversing a row, a forward column, or a backward column. Here's a complete, STL-compliant implementation that fits your requirements:
First, we'll define the iterator class, then create "view" classes to wrap row/column access, and finally update the Matrix class to expose these views.
1. The Iter Class (Core Iterator)
This class handles all iteration logic, with a configurable step size that controls how ++/-- operations move through the matrix data.
#include <vector> #include <iterator> class Iter { public: // STL iterator requirements (enables compatibility with STL algorithms) using iterator_category = std::random_access_iterator_tag; using value_type = int; using difference_type = std::ptrdiff_t; using pointer = int*; using reference = int&; // Constructor: ties iterator to matrix data, sets starting index and step Iter(std::vector<int>& data, size_t idx, difference_type step) : data_(data), idx_(idx), step_(step) {} // Dereference operators reference operator*() { return data_[idx_]; } pointer operator->() { return &data_[idx_]; } // Increment (prefix/postfix) Iter& operator++() { idx_ += step_; return *this; } Iter operator++(int) { Iter temp = *this; idx_ += step_; return temp; } // Decrement (prefix/postfix) Iter& operator--() { idx_ -= step_; return *this; } Iter operator--(int) { Iter temp = *this; idx_ -= step_; return temp; } // Random access operations Iter& operator+=(difference_type n) { idx_ += step_ * n; return *this; } Iter operator+(difference_type n) const { Iter temp = *this; temp += n; return temp; } Iter& operator-=(difference_type n) { idx_ -= step_ * n; return *this; } Iter operator-(difference_type n) const { Iter temp = *this; temp -= n; return temp; } difference_type operator-(const Iter& other) const { // Assumes iterators belong to the same range (same step size) return (idx_ - other.idx_) / step_; } reference operator[](difference_type n) { return data_[idx_ + step_ * n]; } // Comparison operators bool operator==(const Iter& other) const { return &data_ == &other.data_ && idx_ == other.idx_; } bool operator!=(const Iter& other) const { return !(*this == other); } bool operator<(const Iter& other) const { // Adjust comparison based on step direction (forward vs reverse) return step_ > 0 ? idx_ < other.idx_ : idx_ > other.idx_; } bool operator>(const Iter& other) const { return other < *this; } bool operator<=(const Iter& other) const { return !(*this > other); } bool operator>=(const Iter& other) const { return !(*this < other); } private: std::vector<int>& data_; size_t idx_; difference_type step_; // Controls movement: +1 for rows, +nCol for forward cols, -nCol for reverse cols };
2. View Classes (RowView & ColView)
These act as proxy objects to expose the correct iterators for rows and columns. They encapsulate the starting/ending indices and step size configuration.
RowView (For Row Access)
class RowView { public: RowView(std::vector<int>& data, size_t rowID, size_t nCol) : data_(data), start_idx_(rowID * nCol), end_idx_(start_idx_ + nCol) {} // Forward row iteration (step = 1) Iter begin() { return Iter(data_, start_idx_, 1); } Iter end() { return Iter(data_, end_idx_, 1); } private: std::vector<int>& data_; size_t start_idx_; size_t end_idx_; };
ColView (For Column Access, Including Reverse)
class ColView { public: ColView(std::vector<int>& data, size_t colID, size_t nRow, size_t nCol) : data_(data), col_id_(colID), n_row_(nRow), n_col_(nCol) {} // Forward column iteration (step = +nCol) Iter begin() { return Iter(data_, col_id_, static_cast<std::ptrdiff_t>(n_col_)); } Iter end() { return Iter(data_, col_id_ + n_row_ * n_col_, static_cast<std::ptrdiff_t>(n_col_)); } // Reverse column iteration (step = -nCol) Iter rbegin() { return Iter(data_, col_id_ + (n_row_ - 1) * n_col_, -static_cast<std::ptrdiff_t>(n_col_)); } Iter rend() { return Iter(data_, col_id_ - n_col_, -static_cast<std::ptrdiff_t>(n_col_)); } private: std::vector<int>& data_; size_t col_id_; size_t n_row_; size_t n_col_; };
3. Updated Matrix Class
We modify the Matrix to return our view classes instead of raw iterators, enabling the A.row(3).begin() syntax you want.
class Matrix { public: RowView row(size_t rowID) { // Optional: add bounds checking here (e.g., if rowID >= nRow, throw or handle) return RowView(data_, rowID, nCol); } ColView col(size_t colID) { // Optional: add bounds checking here return ColView(data_, colID, nRow, nCol); } private: std::vector<int> data_{1, 2, 3, 4, 5, 6}; size_t nRow{3}; size_t nCol{2}; };
- Step Size is the Magic: The
step_member inIterdictates how the iterator moves:- Row iteration uses
step = 1(moves to the next element in the same row) - Forward column iteration uses
step = nCol(jumps down to the same column in the next row) - Reverse column iteration uses
step = -nCol(jumps up to the same column in the previous row)
This is exactly what makes++A.row(0).begin()and++A.col(0).rbegin()behave differently.
- Row iteration uses
- STL Compliance: The iterator implements all requirements for a random-access iterator, so you can use it with STL algorithms like
std::for_each,std::copy, or range-based for loops. - View Abstraction: The view classes hide the low-level index calculations from the user, making the Matrix interface clean and intuitive.
Here's how you'd use this implementation in practice:
#include <iostream> #include <algorithm> int main() { Matrix mat; // Traverse row 0 std::cout << "Row 0: "; for (auto it = mat.row(0).begin(); it != mat.row(0).end(); ++it) { std::cout << *it << " "; } std::cout << "\n"; // Traverse column 0 forward std::cout << "Column 0 (forward): "; for (auto it = mat.col(0).begin(); it != mat.col(0).end(); ++it) { std::cout << *it << " "; } std::cout << "\n"; // Traverse column 0 backward std::cout << "Column 0 (reverse): "; for (auto it = mat.col(0).rbegin(); it != mat.col(0).rend(); ++it) { std::cout << *it << " "; } std::cout << "\n"; // Use STL algorithm to modify row 1 std::cout << "Row 1 (doubled): "; auto row1 = mat.row(1); std::for_each(row1.begin(), row1.end(), [](int& x) { x *= 2; }); for (int val : row1) { // Range-based for works! std::cout << val << " "; } std::cout << "\n"; return 0; }
Output:
Row 0: 1 2 Column 0 (forward): 1 3 5 Column 0 (reverse): 5 3 1 Row 1 (doubled): 6 8
内容的提问来源于stack exchange,提问作者DrBombe

