Building makemore Part 5: Building a WaveNet
Watch on YouTube →
Overview
Andrej Karpathy refactors the 'makemore' character-level language model to resemble a WaveNet architecture, increasing input context from 3 to 8 characters and implementing a hierarchical fusion of information. He introduces custom 'flatten' and 'sequential' modules, refines the Batch Norm 1D layer for multi-dimensional inputs, and demonstrates how this deeper, more structured approach, despite similar parameter counts, shows potential for improved performance, reaching a validation loss of 1.993 with larger embeddings.
Key takeaways
- Hierarchical information fusion, inspired by WaveNet, improves character-level language models by processing context progressively.
- PyTorch's `nn.Linear` layer's ability to handle multi-dimensional inputs is crucial for parallelizing feature group processing (e.g., bigrams).
- Custom `FlattenConsecutive` and `Sequential` modules streamline model architecture design and management.
- Batch Norm 1D requires careful implementation to correctly handle multi-dimensional inputs, averaging statistics over the appropriate dimensions (batch and time, not just batch).
- Scaling model size (embeddings, hidden units) and context length significantly impacts performance, but requires robust experimental setups for effective hyperparameter tuning.
- Convolutions in models like WaveNet are primarily an efficiency optimization for sliding the model structure over sequences, not a fundamental architectural change.
Chapters
- Goal: Move beyond a simple MLP to a deeper architecture that progressively fuses information.
- Inspiration: WaveNet paper (2016) for its hierarchical approach to predicting sequences.
- Current model: 3 previous characters predict the 4th using a simple MLP.
- New approach: Process more characters and use a deeper, hierarchical structure.
- Starter code based on Part 3, with minimal changes.
- Data generation: 182,000 examples of 3 characters predicting the 4th.
- Building blocks: Custom 'Linear' and 'BatchNorm1D' layers mimicking PyTorch APIs.
- BatchNorm1D complexity: Handles training/evaluation modes and uses exponential moving average for running stats.
- Problem: Loss curve is 'daggers in my eyes' due to raw float list plotting.
- Solution: Convert loss list to a PyTorch tensor, reshape into rows of 1000 elements.
- Method: Use tensor `view` to create a 200x1000 shape, then compute row-wise mean.
- Result: A much smoother and more interpretable loss curve plot.
- Goal: Consolidate model logic into a list of layers, removing special cases.
- New modules: 'Embedding' (for character lookup) and 'Flatten' (for reshaping tensor).
- Implementation: These modules mimic `torch.nn.Embedding` and `torch.nn.Flatten`.
- Benefit: Simplifies the forward pass by treating embedding and flattening as standard layers.
- Problem: Managing layers in a naked list is cumbersome.
- Solution: Implement a 'Sequential' module, similar to `torch.nn.Sequential`.
- Sequential functionality: Passes input through a list of layers in order.
- Benefit: Organizes the model into a single module, simplifying parameter management and forward pass calls.
- Issue: Model re-initialization led to gibberish output due to Batch Norm in training mode.
- Cause: Passing a single example to Batch Norm in training mode results in NaN variance.
- Explanation: Variance of a single number is undefined.
- Fix: Ensure the model is in evaluation mode (`model.eval()`) during sampling/evaluation.
- Current state: Training loss 2.05, validation loss 2.10.
- Strategy: Increase context length from 3 to 8 characters.
- Dataset change: Block size increased to 8.
- Result: Validation loss improved to 2.02, indicating benefit from more context.
- Debugging setup: Temporarily use a batch size of 4 for shape inspection.
- Input shape: 4x8 (batch size x block size).
- Embedding output: 4x8x10 (batch, block, embedding dim).
- Flatten output: 4x80 (concatenated embeddings).
- Linear layer: Takes 80 input features, outputs 200 channels (4x200 shape).
- Surprising feature: PyTorch's matrix multiplication (`nn.Linear`) handles higher-dimensional inputs.
- Mechanism: Matrix multiplication operates on the last two dimensions, broadcasting over preceding dimensions.
- Example: Input 4x5x80 becomes 4x5x200.
- Application: Enables processing groups of features (e.g., bigrams) in parallel within the linear layer.
- Goal: Fuse pairs of characters, then pairs of bigrams, etc., hierarchically.
- New layer: 'FlattenConsecutive(n)' concatenates 'n' consecutive elements along the last dimension.
- Example: Input 4x8x10 becomes 4x4x20 (fusing pairs of 10-dim embeddings).
- Architecture: Sequential application of FlattenConsecutive(2), Linear, BatchNorm, Tanh, repeated.
Summary, takeaways, and chapters were generated by AI from the video's transcript and may contain errors. The video belongs to its creator, Andrej Karpathy.