PidokuInfra

Project 01 — A From-Scratch Inference Engine

Beginner 6h Difficulty 2/5 Topic 01 of 15

Prerequisites Section I, Section III (01-07)

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#


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 everywhere

5. Milestones#

  1. Export. Train a 784→256→128→10 MLP on MNIST in PyTorch (2 epochs is enough). Save state_dict as .npz. Note that PyTorch stores Linear.weight as (out, in).
  2. Operators. Implement each op as a pure function. Write a unit test per op comparing against the PyTorch equivalent on random input.
  3. 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).
  4. Match. Run 1,000 test images through both. Require max|diff| < 1e-5 on logits and identical argmax.
  5. Batch sweep. Time batch sizes 1, 2, 4 … 1024. Plot latency per batch and images/sec.

6. Starter skeleton#

Go
// 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#

MeasurementExpectation to write down first
Latency at batch 1 vs batch 256Not 256× — why?
Images/sec vs batch sizeRises, then flattens. Where, and what limits it?
Bytes of weights vs .npz size on diskparams × 4
FLOPs per image (2 × Σ in×out) vs measured GFLOP/sCompare with Project 02’s peak
OMP_NUM_THREADS=1 vs defaultBLAS 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 Conv2d via im2col + GEMM (Section IV.04) and run a small CNN.
  • Pre-allocate every intermediate buffer so a forward pass does zero allocations.
  • Fuse linear + relu into 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#

  1. What does a model file actually contain?
  2. Why is one batch-256 call far faster than 256 batch-1 calls? Where does the gain come from?
  3. Why subtract the max before exp in softmax?
  4. What is the FLOP count of a linear layer, and how do you verify it empirically?

12. Next#

Project 02 — CPU matmul benchmark

↑↓ navigate↵ openesc close