~/systems/iterators

Iterators and generators

Hand out items one at a time from nested data, merged streams or files, and save the position so a fresh iterator can resume it.

what

Keep the position as a few plain variables, move it to the next real item in one helper, and expose has_next, next, and a get_state/set_state pair that saves just that position.

use when

Flattening nested or jagged data, merging sorted streams lazily, filtering or interleaving feeds, or pausing and resuming a long walk over data.

time

O(1) amortized per item

space

O(1) beyond the data (O(k) to merge k streams)

You’ll recognise it when

  • The API is has_next() and next(), or Python’s __iter__ and __next__.
  • The data is nested or jagged: a list of lists with empty rows, a tree, blocks of rows.
  • Items come from several sources that must be merged, interleaved or filtered as they’re read.
  • You must pause and resume: save the position, maybe in another process, and carry on exactly where you left off.
  • Reading everything up front is too slow or too big, so the work must be lazy.

It’s often confused with plain traversal. Flattening everything into a list first passes small tests, but it breaks the laziness the question is testing, and it can’t resume a walk over a file or a stream.

The idea

A bookmark doesn’t hold a copy of the book. It holds a page number. Closing the book and opening it tomorrow, even a different copy of the same edition, takes you to the same page.

An iterator is a bookmark into data. Its whole job is to keep the position as explicit state: a few integers (row and col, a byte offset, an index per stream) that say where the next item is. next() reads the item at the position and moves it. has_next() answers by moving the position forward past anything that isn’t an item, and never past an item. And because the position is plain data, saving it is trivial: get_state() returns those integers and set_state() puts them back.

How it works

Take a jagged 2D list, [[1, 2], [], [3], [], []], and build an iterator over it that can be saved and resumed. The position is the pair (row, col): the next item is rows[row][col].

  1. Start at row = 0, col = 0. The position may point at nothing yet (row 0 could be empty). That’s allowed until someone asks.
  2. Settle the position in one helper: while col is past the end of the current row, move to the start of the next row. Empty rows are skipped by the same loop, wherever they are. Afterwards the position is either a real item or the end.
  3. has_next() settles, then checks row < len(rows). Settling never skips a real item, so calling has_next() ten times in a row is the same as calling it once.
  4. next() calls has_next(), raises if there’s nothing left, returns rows[row][col] and adds 1 to col. It doesn’t settle afterwards: the next call will.
  5. get_state() returns {"row": row, "col": col}. set_state(state) on a fresh iterator over the same rows copies them back. The two iterators then hand out exactly the same items.

On the example, next() gives 1 and the saved state is {row: 0, col: 1}. Then 2, then has_next() settles past the empty row to (2, 0), and next() gives 3. Now has_next() walks past the two trailing empty rows and returns false. Restore the saved state into a new iterator, and it hands out 2 and 3 again.

Why it’s correct: every item is at exactly one (row, col), positions only move forward, and the settle loop only steps over positions that aren’t items. So each item is returned once, in order, and the state at any moment is enough to rebuild the iterator.

class Flat2D:
"""Walks a list of rows, item by item, and can save and restore its place."""
def __init__(self, rows):
self.rows = rows
self.row = 0 # the next item is rows[row][col]
self.col = 0
def _settle(self):
# Skip finished and empty rows, so (row, col) is a real item or the end.
while self.row < len(self.rows) and self.col >= len(self.rows[self.row]):
self.row += 1
self.col = 0
def has_next(self):
self._settle() # safe to call any number of times
return self.row < len(self.rows)
def next(self):
if not self.has_next():
raise StopIteration
item = self.rows[self.row][self.col]
self.col += 1
return item
def get_state(self):
return {"row": self.row, "col": self.col} # plain data: easy to save
def set_state(self, state):
self.row, self.col = state["row"], state["col"]
# The Python protocol, so a for loop works too.
def __iter__(self):
return self
def __next__(self):
return self.next()
#include <stdexcept>
#include <vector>
using namespace std;
struct State { int row, col; }; // plain data: easy to save
// Walks a list of rows, item by item, and can save and restore its place.
class Flat2D {
const vector<vector<int>>& rows;
int row = 0, col = 0; // the next item is rows[row][col]
// Skip finished and empty rows, so (row, col) is a real item or the end.
void settle() {
while (row < (int)rows.size() && col >= (int)rows[row].size()) {
row++;
col = 0;
}
}
public:
explicit Flat2D(const vector<vector<int>>& rows) : rows(rows) {}
bool has_next() {
settle(); // safe to call any number of times
return row < (int)rows.size();
}
int next() {
if (!has_next()) throw out_of_range("no more items");
return rows[row][col++];
}
State get_state() const { return {row, col}; }
void set_state(State s) {
row = s.row;
col = s.col;
}
};
import java.util.*;
// Walks a list of rows, item by item, and can save and restore its place.
class Flat2D implements Iterator<Integer> {
record State(int row, int col) {} // plain data: easy to save
private final List<List<Integer>> rows;
private int row = 0, col = 0; // the next item is rows.get(row).get(col)
Flat2D(List<List<Integer>> rows) { this.rows = rows; }
// Skip finished and empty rows, so (row, col) is a real item or the end.
private void settle() {
while (row < rows.size() && col >= rows.get(row).size()) {
row++;
col = 0;
}
}
public boolean hasNext() {
settle(); // safe to call any number of times
return row < rows.size();
}
public Integer next() {
if (!hasNext()) throw new NoSuchElementException();
return rows.get(row).get(col++);
}
State getState() { return new State(row, col); }
void setState(State s) {
row = s.row();
col = s.col();
}
}

Generators: the same thing, written for you

In Python, a function with yield is an iterator whose position is the paused function itself:

def flatten(rows):
for r in rows:
for x in r:
yield x

It’s the shortest way to write a lazy iterator, and fine whenever you don’t need to save the position. When you do, write the class: a live generator can’t be pickled or copied to another process, and its position is hidden in its frame. Raising StopIteration inside a generator doesn’t end it quietly either: since Python 3.7 it turns into a RuntimeError. Use return.

Merging sorted streams lazily

To merge k sorted streams, keep a min-heap with the current head of each stream. Pop the smallest, yield it, and push the next item from the same stream. The stream’s index goes in the tuple to break ties, so the heap never compares two iterators.

import heapq
_DONE = object()
def merge_streams(streams):
heap = []
for sid, s in enumerate(streams):
it = iter(s)
first = next(it, _DONE)
if first is not _DONE:
heap.append((first, sid, it))
heapq.heapify(heap)
while heap:
value, sid, it = heapq.heappop(heap)
yield value
nxt = next(it, _DONE)
if nxt is not _DONE:
heapq.heappush(heap, (nxt, sid, it))
print(list(merge_streams([[1, 4, 9], [2, 3], [], [4, 8]]))) # [1, 2, 3, 4, 4, 8, 9]

The standard library’s heapq.merge does the same. A stream that exposes peek() makes “look at the next value without taking it” easy; on an iterator without one, cache a single lookahead item.

Why it’s O(1) amortized

next() does O(1) work plus the settle loop. The settle loop can step over many empty rows in one call, but each row is stepped over at most once in the iterator’s whole life. Over a full walk the total is O(items + rows), so each call is O(1) amortized. Saving and restoring the state copies two integers: O(1).

Merging k streams costs O(log k) per item for the heap, and the heap holds at most k entries. The iterator itself stores no copy of the data: O(1) extra space for the 2D walk.

Common mistakes

has_next() that moves past an item

If has_next() consumes something (reads the next item to see if it exists and drops it), calling it twice loses an item. Settling should skip only non-items.

def has_next(self): self.col += 1; ... # ✗ two calls skip an item
def has_next(self): self._settle(); ... # ✓ only steps over empty space

Forgetting empty rows at the end

Advancing once after next() handles an empty row in the middle by luck, but trailing empty rows make has_next() say true when nothing is left. Loop until the position is a real item or the end.

if self.col == len(self.rows[self.row]): self.row += 1 # ✗ one step only
while self.row < len(self.rows) and self.col >= len(self.rows[self.row]): ... # ✓

Saving the live object instead of the position

A saved state should be small plain data that survives a restart. Pickling a generator fails, and keeping a reference to the iterator object isn’t a checkpoint: it keeps moving after you “save” it.

self.saved = self.iterator # ✗ the same object, still moving
self.saved = {"row": self.row, "col": self.col} # ✓ a copy of the position

Walking an iterator twice

An iterator is used up once it’s been read to the end. Passing the same iterator to two loops, or calling sum(it) and then len(list(it)), silently gives an empty second pass.

total, count = sum(it), len(list(it)) # ✗ the second pass sees nothing
items = list(it); total, count = sum(items), len(items) # ✓ materialize once

Variations

  • Peeking. Wrap any iterator and cache one item: peek() fills the cache if empty and returns it, next() returns the cache first.
  • Filtering and interleaving. A filter iterator settles by skipping items the test rejects. Round-robin interleaving keeps a queue of iterators and puts each one back after taking an item, dropping it once it’s empty.
  • Deeper nesting and trees. For 3D lists, keep one index per level. For arbitrary nesting or a tree walk, keep a stack of iterators and settle by pushing into sublists and popping finished ones.
  • Files. The position is a byte offset: f.tell() to save and f.seek(offset) to restore. Open the file in binary mode so offsets are exact.
  • Snapshot iterators. An iterator that must see a set as it was when the walk started, even while it changes, can copy on creation, or tag each element with versions and skip those added later or deleted earlier.

Climb the ladder

Our iterator problems in ladder order.

  1. An iterator class that yields chunks: the protocol by hand.
  2. Collapse repeated sensor readings: one item of lookahead.
  3. Iterator with Filter: settle by skipping rejected items.
  4. Range Iterator: a position that’s a number, with negative steps.
  5. Interleave Iterator with Cycle: a queue of iterators.
  6. Flatten a 2D array, with remove: this page’s iterator, plus removals.
  7. Merge K sorted streams: a heap of stream heads.
  8. Weighted data batcher with checkpointing: state that resumes a training run exactly.
  9. Resumable iterators (list, 2D, 3D, file): get_state and set_state at every depth.
  10. Snapshot Set Iterator: frozen views of a set that keeps changing.

Check yourself

5 quick questions. Pick an answer to see why it's right or wrong.

  1. 1

    This 2D iterator moves to the next row at most once per has_next() call. What does this print?

    class Flat:
    def __init__(self, rows):
    self.rows, self.r, self.c = rows, 0, 0
    def has_next(self):
    if self.r < len(self.rows) and self.c == len(self.rows[self.r]):
    self.r, self.c = self.r + 1, 0
    return self.r < len(self.rows)
    def next(self):
    self.has_next()
    x = self.rows[self.r][self.c]
    self.c += 1
    return x
    it = Flat([[1], [], []])
    out = []
    try:
    while it.has_next():
    out.append(it.next())
    except IndexError:
    out.append("IndexError")
    print(out)
  2. 2

    You must merge 1,000 sorted log files, each too big for memory, into one sorted output. What’s the right approach?

  3. 3

    What does this print on Python 3.7 or later?

    def first_of_each(rows):
    for r in rows:
    yield next(iter(r))
    try:
    print(list(first_of_each([[1, 2], [], [3]])))
    except RuntimeError:
    print("RuntimeError")
  4. 4

    A resumable iterator walks the lines of a 50 GB file and must survive a process restart. What should get_state() return?

  5. 5

    A 2D iterator settles past empty rows inside has_next(). The data has n rows, most of them empty, and m items in total. What does a full walk with has_next() and next() cost?

Practice problems

Solve these right here, in Python, C++ or Java. Tests run as you go.

All 10 problems on this topic

Further reading

esc