# How NanoChat's Custom Linear Layer Improves Numerical Precision

> Discover how NanoChat's custom Linear layer uses float32 master weights for high-precision updates and fast low-precision matrix multiplication.

- Repository: [Andrej/nanochat](https://github.com/karpathy/nanochat)
- Tags: deep-dive
- Published: 2026-03-10

---

**NanoChat's custom Linear layer maintains master weights in float32 while dynamically casting them to the activation dtype during forward passes, delivering high-precision optimizer updates without sacrificing the speed benefits of low-precision matrix multiplication.**

The `karpathy/nanochat` repository implements a precision-aware training technique that solves the numerical stability issues common in mixed-precision transformers. This article examines how nanochat's custom Linear layer improves precision by separating weight storage dtype from computation dtype, a design pattern implemented in [`nanochat/gpt.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/gpt.py) that eliminates quantization drift during training.

## How NanoChat's Custom Linear Layer Improves Precision

Standard `torch.nn.Linear` layers in mixed-precision training typically store weights in the same low-precision dtype as activations (e.g., `bfloat16`). This forces the optimizer to update weights using rounded gradients, accumulating numerical errors across deep transformer stacks. NanoChat's approach resolves this by storing the master weight tensor in full precision while casting only during the forward computation.

### Source Code in nanochat/gpt.py

The custom implementation resides in [`nanochat/gpt.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/gpt.py) at lines 45-51, subclassing `nn.Linear` to override the forward pass with explicit dtype handling:

```python
class Linear(nn.Linear):
    """nn.Linear that casts weights to match input dtype in forward.
    Replaces autocast: master weights stay fp32 for optimizer precision,
    but matmuls run in the activation dtype (typically bf16 from embeddings)."""
    def forward(self, x):
        return F.linear(x, self.weight.to(dtype=x.dtype))

```

This minimal override replaces PyTorch's automatic casting behavior with explicit control over when and how weight dtype conversion occurs.

### The On-the-Fly Casting Mechanism

During the forward pass, the layer executes `self.weight.to(dtype=x.dtype)` to cast the master weight to match the input activation's dtype—typically `bfloat16` as defined in [`nanochat/common.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/common.py) via the global `COMPUTE_DTYPE` setting. The matrix multiplication executes in the activation's low-precision format via `F.linear`, but the underlying weight tensor referenced by the optimizer remains in `float32`. After the forward pass, gradients flow back to the full-precision master copy, allowing the optimizer to apply updates without low-precision rounding errors.

## Key Benefits of the Precision-Preserving Design

This architecture delivers two critical advantages that standard mixed-precision approaches cannot achieve simultaneously.

### High-Precision Optimizer Updates

By retaining weights in `float32`, the optimizer computes weight updates using full-precision gradients. **Standard low-precision training** stores weights in `bfloat16`, which limits the smallest representable weight change and causes quantization drift over thousands of training steps. According to the nanochat source code, this custom layer eliminates drift by ensuring the master copy maintains 32-bit precision throughout training, improving convergence stability.

### Efficient Low-Precision Computation

The matrix multiplication itself runs in the activation dtype (usually `bfloat16`), preserving memory bandwidth and tensor core utilization. **Dynamic casting** occurs only during the forward pass as a view operation, adding negligible computational overhead while allowing the model to use fast low-precision kernels. This design appears throughout the transformer architecture in [`nanochat/engine.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/engine.py), which constructs the model using these custom layers exclusively.

## Practical Usage Example

The following example demonstrates the dtype behavior in practice:

```python
import torch
from nanochat.gpt import Linear

# Create a custom Linear layer (weights stay in float32)

layer = Linear(in_features=768, out_features=3072, bias=False)

# Input activations are in bf16 (common for NanoChat)

x = torch.randn(4, 768, dtype=torch.bfloat16, device='cuda')

# Forward pass – weight is cast to bf16 on-the-fly

y = layer(x)                     # y is also bf16

print(y.dtype)                   # torch.bfloat16

# During training, the underlying weight remains float32

print(layer.weight.dtype)        # torch.float32

```

This pattern ensures that `layer.weight` receives full-precision gradient updates during backpropagation while the forward computation benefits from efficient `bfloat16` operations.

## Summary

- **NanoChat's custom Linear layer**, defined in [`nanochat/gpt.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/gpt.py) lines 45-51, subclasses `nn.Linear` to implement dtype-aware weight casting.
- **Master weights remain in float32** to provide high-precision optimizer updates and prevent quantization drift.
- **Dynamic casting to activation dtype** (typically `bfloat16` from [`nanochat/common.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/common.py)) occurs during forward passes, enabling efficient matrix multiplication without manual autocast management.
- This approach balances **memory efficiency** during inference with **numerical stability** during training, solving the precision-speed trade-off inherent in standard mixed-precision training.

## Frequently Asked Questions

### Why does NanoChat use a custom Linear layer instead of PyTorch's autocast?

PyTorch's `autocast` converts both inputs and weights to low-precision for each operation, but the master weight copy remains in that low-precision dtype. NanoChat's custom layer keeps the master weights in full-precision `float32` permanently, only casting a temporary view during the forward pass. This ensures the optimizer updates weights using accurate gradients rather than rounded low-precision values.

### Does storing weights in float32 increase GPU memory usage?

The master weight tensor occupies twice the memory of a `bfloat16` tensor. However, during the forward pass, only the cast view exists temporarily, and activations remain in low-precision. The trade-off accepts slightly higher parameter memory for significantly improved training stability and convergence speed, which typically reduces the total number of training steps required.

### What determines the activation dtype in NanoChat?

The global `COMPUTE_DTYPE` variable defined in [`nanochat/common.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/common.py) controls the activation precision, typically set to `torch.bfloat16`. The custom `Linear` layer automatically detects the input dtype via `x.dtype` and casts weights accordingly, making the implementation flexible across different precision configurations.

### Is this technique beneficial for inference only, or training as well?

The technique primarily benefits **training**. During inference, weights could remain in low-precision without significant accuracy loss. However, during training, the custom layer prevents the gradual degradation of weight precision caused by repeated low-precision optimizer updates. According to the `karpathy/nanochat` source code, this design ensures stable convergence when training deep transformer models from scratch.