Skip to content

Iterator & Generator Protocol

Overview

Iterators and generators are fundamental to ML pipelines: - DataLoaders use iterators to stream data - Generators are memory-efficient (don't load entire dataset) - Yield enables lazy evaluation for large datasets - JAX transformations leverage generator composition

This is why PyTorch DataLoaders don't load all data into memory upfront.


The Iterator Protocol

Core Concept

An iterator is an object that implements two methods:

class Iterator:
    def __iter__(self):
        """Return self (iterator object)."""
        return self

    def __next__(self):
        """Return next value, raise StopIteration when done."""
        if self.has_next():
            return self.get_next()
        raise StopIteration

Simple Example

# Built-in iterator
data = [1, 2, 3]
iterator = iter(data)  # Calls data.__iter__()

print(next(iterator))  # 1, calls iterator.__next__()
print(next(iterator))  # 2
print(next(iterator))  # 3
print(next(iterator))  # StopIteration exception

# For loop uses iterator protocol
for value in data:  # Implicitly calls iter() and __next__()
    print(value)

Custom Iterator

class RangeIterator:
    """Custom iterator (like range)."""
    def __init__(self, start, end):
        self.current = start
        self.end = end

    def __iter__(self):
        return self

    def __next__(self):
        if self.current < self.end:
            value = self.current
            self.current += 1
            return value
        raise StopIteration

# Usage
for value in RangeIterator(0, 5):
    print(value)  # 0, 1, 2, 3, 4

Generators: Simplified Iterators

Core Concept

A generator is a function that yields values instead of returning:

# Generator function (uses yield)
def count_to_n(n):
    """Generator that yields 0 to n-1."""
    i = 0
    while i < n:
        yield i  # Pause here, return value, resume when next() called
        i += 1

# Generator object (not executed yet)
gen = count_to_n(3)
print(type(gen))  # <class 'generator'>

# Iteration
print(next(gen))  # 0 (function resumes after first yield)
print(next(gen))  # 1
print(next(gen))  # 2
print(next(gen))  # StopIteration

Why Generators Matter for ML

# WITHOUT generators: Load entire dataset into memory
def load_all_data():
    data = []
    for line in open("huge_dataset.txt"):
        data.append(process(line))  # Stores everything in RAM
    return data

all_data = load_all_data()  # Uses gigabytes of RAM!

# WITH generators: Lazy loading
def load_data_lazy():
    for line in open("huge_dataset.txt"):
        yield process(line)  # Only one sample in RAM at a time

for sample in load_data_lazy():
    # Process one sample at a time
    train_model(sample)

Memory Comparison:

All-at-once:  Load 1B samples × 1KB each = 1TB RAM (impossible!)
Generator:    Load 1 sample at a time = 1KB RAM (feasible!)


Generator Expressions

Lazy versions of list comprehensions:

# List comprehension: Computes all values immediately
squares_list = [x**2 for x in range(1000000)]  # Allocates huge list

# Generator expression: Computes lazily
squares_gen = (x**2 for x in range(1000000))  # Tiny object

# Iterate through generator
for square in squares_gen:
    print(square)

# Or convert to list if needed (but then defeats purpose)
squares_list = list(squares_gen)  # Now it's big

Chaining Generators

# Multiple generators compose efficiently
data = range(1000000)

# Don't do this (allocates intermediate lists)
filtered = [x for x in data if x % 2 == 0]
squared = [x**2 for x in filtered]
top_100 = squared[:100]

# Do this (composable, lazy)
filtered = (x for x in data if x % 2 == 0)
squared = (x**2 for x in filtered)
top_100 = list(itertools.islice(squared, 100))

PyTorch DataLoader: Iterator Pattern

How DataLoader Uses Iterators

import torch
from torch.utils.data import DataLoader, Dataset

class MyDataset(Dataset):
    """Custom dataset."""
    def __len__(self):
        return 1000

    def __getitem__(self, idx):
        # Called by DataLoader to fetch individual samples
        return torch.randn(3, 224, 224), idx % 10

# Create iterator
dataset = MyDataset()
loader = DataLoader(
    dataset,
    batch_size=32,
    shuffle=True,
    num_workers=4  # Load batches in parallel
)

# DataLoader is iterable
for batch_images, batch_labels in loader:  # Uses iterator protocol
    print(batch_images.shape)  # (32, 3, 224, 224)
    # Process batch

Multiple Epochs

model = torch.nn.Linear(10, 5)
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)

for epoch in range(3):
    # Each epoch creates new iterator over same dataset
    for batch_x, batch_y in loader:
        # Forward pass
        pred = model(batch_x)
        loss = torch.nn.functional.mse_loss(pred, batch_y)

        # Backward pass
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

    print(f"Epoch {epoch} done")

Key insight: DataLoader creates a fresh iterator each epoch, respecting shuffle order.


Advanced: Stateful Generators

Generators can maintain state:

# Generator with state
def batch_generator(data, batch_size):
    """Yield batches of data."""
    for i in range(0, len(data), batch_size):
        batch = data[i:i+batch_size]
        yield batch
        # State (i) is preserved between yields

# Usage
batches = batch_generator([1, 2, 3, 4, 5], batch_size=2)
print(next(batches))  # [1, 2]
print(next(batches))  # [3, 4]
print(next(batches))  # [5]
print(next(batches))  # StopIteration

Generator with Initialization

def infinite_loader(data):
    """Cycle through data infinitely (for validation)."""
    while True:
        for sample in data:
            yield sample

# Usage in validation loop
gen = infinite_loader(val_dataset)
for _ in range(num_validation_batches):
    batch = next(gen)
    validate(model, batch)

Important Generator Features

send() Method: Bidirectional Communication

def echo_generator():
    """Receive values via send()."""
    while True:
        value = yield  # Wait for value to be sent
        print(f"Received: {value}")

gen = echo_generator()
next(gen)  # Prime the generator (move to first yield)
gen.send("Hello")  # Send value into generator
gen.send("World")
# Output:
# Received: Hello
# Received: World

ML Application: Feedback-driven data generation

def curriculum_generator(all_data):
    """Adjust difficulty based on model performance."""
    easy_data = [d for d in all_data if d['difficulty'] < 5]
    hard_data = [d for d in all_data if d['difficulty'] >= 5]

    current = easy_data

    while True:
        # Yield a batch
        batch = yield current

        # Receive feedback on model performance
        if batch.get('avg_loss') < 0.1:
            # Model is doing well, increase difficulty
            current = hard_data
        else:
            # Keep easy examples
            current = easy_data

Nested Generators with yield from

# yield from (Python 3.3+)
def flatten(nested_list):
    """Flatten nested lists using yield from."""
    for item in nested_list:
        if isinstance(item, list):
            yield from flatten(item)  # Delegate to nested generator
        else:
            yield item

# Usage
data = [1, [2, [3, 4]], 5]
for value in flatten(data):
    print(value)  # 1, 2, 3, 4, 5

# Practical example: Flattening dataset hierarchies
def iter_dataset(root_dir):
    """Recursively iterate all samples."""
    for item in os.listdir(root_dir):
        path = os.path.join(root_dir, item)
        if os.path.isdir(path):
            yield from iter_dataset(path)  # Recurse
        else:
            yield path  # Leaf node

Itertools: Functional Iterator Tools

Common Patterns

import itertools

# chain: Combine multiple iterables
data1 = [1, 2, 3]
data2 = [4, 5, 6]
combined = itertools.chain(data1, data2)

# islice: Take first N elements
first_100 = itertools.islice(all_data, 100)

# cycle: Repeat iterator infinitely
repeated = itertools.cycle([1, 2, 3])

# zip_longest: Zip with padding
a = [1, 2]
b = [3, 4, 5]
zipped = itertools.zip_longest(a, b, fillvalue=0)  # [(1,3), (2,4), (0,5)]

# groupby: Group consecutive equal elements
sorted_data = [1, 1, 1, 2, 2, 3, 3, 3, 3]
for key, group in itertools.groupby(sorted_data):
    print(f"{key}: {list(group)}")

Practical DataLoader Example

from itertools import chain, islice

# Combine train and validation iterators
train_loader = DataLoader(train_dataset, batch_size=32)
val_loader = DataLoader(val_dataset, batch_size=32)

# First train 100 batches, then validate
combined = chain(islice(train_loader, 100), val_loader)

for batch in combined:
    if batch in train_loader:
        train_step(batch)
    else:
        val_step(batch)

Performance: Lists vs Generators

import sys

# List: All in memory
list_comp = [x**2 for x in range(1000000)]
print(sys.getsizeof(list_comp))  # ~8MB

# Generator: Tiny object
gen_expr = (x**2 for x in range(1000000))
print(sys.getsizeof(gen_expr))  # ~136 bytes

# Time to create
import time

start = time.time()
list_comp = [x**2 for x in range(10000000)]
print(f"List: {time.time() - start:.3f}s")  # ~0.5s

start = time.time()
gen_expr = (x**2 for x in range(10000000))
print(f"Generator: {time.time() - start:.3f}s")  # ~0.0s (just creates object)

For ML: Use generators for datasets, lists only when needed.


The Iterable vs Iterator Distinction

# Iterable: Has __iter__() (returns an iterator)
data = [1, 2, 3]
iterator1 = iter(data)
iterator2 = iter(data)
# iterator1 and iterator2 are different objects!
# Each starts from beginning

# Iterator: Has __iter__() and __next__()
iterator = iter([1, 2, 3])
print(iter(iterator) is iterator)  # True! Iterator's __iter__ returns self
print(next(iterator))  # 1
print(next(iterator))  # 2

Common Pitfalls

Pitfall 1: Exhausted Generators

gen = (x**2 for x in range(5))

# First iteration
for x in gen:
    print(x)

# Second iteration doesn't work!
for x in gen:
    print(x)  # Prints nothing - generator is exhausted

# Solution: Recreate generator
gen = (x**2 for x in range(5))

Pitfall 2: Closing Generator

def my_gen():
    try:
        yield 1
        yield 2
    finally:
        print("Cleaning up...")

gen = my_gen()
next(gen)  # 1
gen.close()  # Raises GeneratorExit, triggers finally block
# Output: Cleaning up...

# Practical: Ensure resources are freed
def file_line_generator(filename):
    with open(filename) as f:
        try:
            for line in f:
                yield line.strip()
        finally:
            print("File closed")  # Always runs

Summary

  • Iterators = Objects with __iter__() and __next__()
  • Generators = Functions with yield (simplified iterators)
  • Generator expressions = Lazy list comprehensions
  • DataLoader = Iterator pattern for batching
  • Memory efficiency = Process one sample at a time
  • Composition = Chain multiple generators efficiently