Run a trained network with nothing but Go’s standard library, and match PyTorch to five decimal places.
1. What you build#
A small inference engine: a set of operators (Linear, ReLU, GELU, LayerNorm, Softmax,
Embedding), a way to load trained weights from a file, and a forward() that executes them
in order. You train a tiny MNIST MLP in PyTorch (or download one), export its weights, and run
it in your engine.
No autograd, no framework. Just arrays in, arrays out.
2. Why it matters#
Everything later in the curriculum is an optimization of this loop. If the forward pass is a black box to you, then KV caches, fusion, and quantization are tricks to memorize. If you have written it, they are obvious edits to code you own.
It also teaches the first production lesson: inference needs only the weights and the arithmetic. Everything else in a training framework is dead weight at serving time (Section I.02).
3. Read first#
- I.04 — Tensors and the forward pass
- III.03 — Tensors, shapes, broadcasting
- III.04 — Neural networks from scratch
- III.07 — Softmax and stability
4. Spec#
Input: weights.safetensors (name -> array) + model.json (ordered list of ops)
Engine: Load(path) (*Engine, error) ; (*Engine).Forward(x Tensor) Tensor
Ops: linear, relu, gelu, layernorm, softmax, embedding, add (residual)
Batch: every op must accept a leading batch dimension
Dtype: float32 everywhere5. Milestones#
- Export. Train a 784→256→128→10 MLP on MNIST in PyTorch (2 epochs is enough). Save
state_dictas.npz. Note that PyTorch storesLinear.weightas(out, in). - Operators. Implement each op as a pure function. Write a unit test per op comparing against the PyTorch equivalent on random input.
- Graph. Describe the model as a list of
{op, inputs, params}and execute it in order. This is a computational graph in its simplest form (Section IV.02 formalizes it). - Match. Run 1,000 test images through both. Require
max|diff| < 1e-5on logits and identical argmax. - Batch sweep. Time batch sizes 1, 2, 4 … 1024. Plot latency per batch and images/sec.
6. Starter skeleton#
// Tensor is a batch of vectors: Rows × Cols, stored row by row. Always float32.
type Tensor struct {
Rows, Cols int
Data []float32
}
// Each operator takes a tensor and its parameters and returns a tensor.
type Op func(x Tensor, params ...Tensor) Tensor
func linear(x Tensor, params ...Tensor) Tensor { // W is (out, in) as PyTorch stores it
W, b := params[0], params[1]
y := Tensor{x.Rows, W.Rows, make([]float32, x.Rows*W.Rows)}
for r := 0; r < x.Rows; r++ {
in := x.Data[r*x.Cols : (r+1)*x.Cols]
for o := 0; o < W.Rows; o++ {
sum := b.Data[o]
for i, w := range W.Data[o*W.Cols : (o+1)*W.Cols] {
sum += w * in[i]
}
y.Data[r*W.Rows+o] = sum
}
}
return y
}
func relu(x Tensor, _ ...Tensor) Tensor {
for i, v := range x.Data {
x.Data[i] = max(v, 0) // in place: no new allocation
}
return x
}
func softmax(x Tensor, _ ...Tensor) Tensor {
for r := 0; r < x.Rows; r++ {
row := x.Data[r*x.Cols : (r+1)*x.Cols]
m, sum := slices.Max(row), float32(0) // subtract the max for stability (III.07)
for i, v := range row {
row[i] = float32(math.Exp(float64(v - m)))
sum += row[i]
}
for i := range row {
row[i] /= sum
}
}
return x
}
var ops = map[string]Op{"linear": linear, "relu": relu, "softmax": softmax /* layernorm, gelu, embedding: yours */}
// Node is one step of the model: an operator and the names of its weights.
type Node struct {
Op string `json:"op"`
Params []string `json:"params"`
}
type Engine struct {
Graph []Node
Weights map[string]Tensor
}
func (e *Engine) Forward(x Tensor) Tensor {
for _, n := range e.Graph {
params := make([]Tensor, len(n.Params))
for i, name := range n.Params {
params[i] = e.Weights[name]
}
x = ops[n.Op](x, params...)
}
return x
}7. What to measure#
| Measurement | Expectation to write down first |
|---|---|
| Latency at batch 1 vs batch 256 | Not 256× — why? |
| Images/sec vs batch size | Rises, then flattens. Where, and what limits it? |
Bytes of weights vs .npz size on disk | params × 4 |
FLOPs per image (2 × Σ in×out) vs measured GFLOP/s | Compare with Project 02’s peak |
OMP_NUM_THREADS=1 vs default | BLAS threading is doing the work |
Record all of it in numbers.md.
8. Done when#
- All ops pass a per-op test against PyTorch.
- End-to-end logits match within
1e-5; argmax identical on 1,000 images. - You have a throughput-vs-batch plot and can explain its shape in two sentences.
- You can state the model’s parameter count, weight bytes, and FLOPs per image from memory.
9. Common pitfalls#
Float64 leaking in. Go’s math package works in float64; convert at the call and store
float32, or every tensor doubles in size and halves in speed.
Forgetting the transpose. (out, in) vs (in, out). Pre-transpose once at load time
rather than on every call.
Timing the first call. It includes page faults and BLAS thread spin-up. Warm up.
Comparing against PyTorch in training mode. Call model.eval().
10. Stretch goals#
- Add
Conv2dvia im2col + GEMM (Section IV.04) and run a small CNN. - Pre-allocate every intermediate buffer so a forward pass does zero allocations.
- Fuse
linear + reluinto one function and measure the difference (preview of IV.09). - Load weights with
np.load(..., mmap_mode="r")and measure cold vs warm load time (II.07).
11. Interview questions this project answers#
- What does a model file actually contain?
- Why is one batch-256 call far faster than 256 batch-1 calls? Where does the gain come from?
- Why subtract the max before
expin softmax? - What is the FLOP count of a linear layer, and how do you verify it empirically?