PidokuInfra

Tensors and the Forward Pass

Foundations Beginner 1h Difficulty 2/5 Topic 04 of 11

Prerequisites 01, 02, 03


1. What is it?#

A tensor is an n-dimensional array of numbers. That is the whole definition. Physicists mean something more specific; in ML, “tensor” means “array with a shape.”

scalar   5                       shape ()          0-D
vector   [1, 2, 3]               shape (3,)        1-D
matrix   [[1,2],[3,4]]           shape (2,2)       2-D
3-tensor [[[...]]]               shape (2,3,4)     3-D

A forward pass is what happens when you push an input tensor through a model: a fixed sequence of tensor operations that ends with an output tensor.

That’s it. Inference is a forward pass. Everything else in this repository is about making forward passes fast.


2. Why does it exist?#

Two reasons, one mathematical and one brutally practical.

Mathematical: neural networks are built from operations that are naturally expressed on arrays — matrix multiplication, elementwise functions, reductions. Tensors are the right data structure.

Practical — and this is the one that matters for you: modern hardware is fast only when it does the same operation on many values at once. A CPU’s SIMD unit does 8-16 floats per instruction. A GPU does thousands. Expressing computation as tensor operations is what lets the hardware exploit that.

Scalar thinking:                  Tensor thinking:
for i := range y {                y = W · x     (one call, whole arrays)
    y[i] = w[i] * x[i]
}
one value per step                CPU SIMD: 8-16 values per instruction
                                  GPU:      thousands of values at once

Same math. 10,000x difference. The gap between those two lines of code is the gap between an ML hobbyist and an inference engineer.


3. Simple analogy#

A tensor is a spreadsheet with more than two dimensions.

  • 1-D: a column of numbers.
  • 2-D: a sheet.
  • 3-D: a workbook with several sheets.
  • 4-D: a filing cabinet of workbooks.

In LLM inference you constantly handle 3-D and 4-D tensors, and the dimensions almost always mean the same things:

(batch, sequence, hidden)              ← the main "residual stream" tensor
(batch, heads, sequence, head_dim)     ← attention tensors

If you can read a shape and say out loud what each axis means, you can debug 80% of inference bugs. Shape errors are the compile errors of ML.


4. Tiny example#

A complete forward pass, by hand, with shapes annotated:

Go
package main

import "fmt"

// Dense is one layer: out = W·x + b. W is rows×cols, stored row by row in one flat slice.
type Dense struct {
	W          []float64
	B          []float64
	Rows, Cols int
}

func (d Dense) Forward(x []float64) []float64 {
	out := make([]float64, d.Rows)
	for i := 0; i < d.Rows; i++ {
		sum := d.B[i]
		for j := 0; j < d.Cols; j++ {
			sum += d.W[i*d.Cols+j] * x[j]
		}
		out[i] = sum
	}
	return out
}

func relu(v []float64) []float64 {
	for i := range v {
		v[i] = max(0, v[i])
	}
	return v
}

var (
	l1 = Dense{W: []float64{0.5, 0.1, 0.2, 0.9, -0.3, 0.4}, B: []float64{0, 0.1, -0.1}, Rows: 3, Cols: 2} // (3, 2)
	l2 = Dense{W: []float64{1.0, -1.0, 0.5}, B: []float64{0}, Rows: 1, Cols: 3}                           // (1, 3)
)

func main() {
	x := []float64{1.0, 2.0} // (2,)  input

	h := l1.Forward(x) // (3,2)·(2,) -> (3,)   [0.7 2.1 0.4]
	a := relu(h)       // (3,)                 [0.7 2.1 0.4]
	y := l2.Forward(a) // (1,3)·(3,) -> (1,)   [-1.2]
	fmt.Printf("%.1f %.1f\n", a, y)
}

Do the first line by hand to make sure you believe it:

row 0:  0.5(1.0) + 0.1(2.0) = 0.7   + 0.0  = 0.7
row 1:  0.2(1.0) + 0.9(2.0) = 2.0   + 0.1  = 2.1
row 2: -0.3(1.0) + 0.4(2.0) = 0.5   - 0.1  = 0.4

Now batch it — this is the key move:

Go
// Three inputs at once: X has shape (3, 2). Add this to the program above.
X := [][]float64{
	{1.0, 2.0},
	{0.0, 1.0},
	{3.0, -1.0},
}
Y := make([][]float64, len(X)) // (3, 1)
for i, x := range X {
	Y[i] = l2.Forward(relu(l1.Forward(x))) // same weights, reused for every row
}
fmt.Printf("%.1f\n", Y)

Nothing changed except a dimension. W1 is read once and applied to all three inputs. This is the mechanical basis of batching, and it is why the @ operator (GEMM) is the single most important operation in this entire field.

Note the transpose: with a batch dimension first, we write X @ W.T rather than W @ x. Both compute the same thing; the layout differs. Layout matters enormously for performance (Section IV.07), and confusion about it is a top source of bugs.


5. Technical explanation#

The four properties of a tensor#

Go
// A tensor is a flat buffer plus the metadata that says how to read it.
type Tensor struct {
	Data    []float32 // the bytes: one contiguous block
	Shape   []int     // [2 3 4]     the logical dimensions
	Strides []int     // [12 4 1]    how far to jump in Data to move one step along each axis
	Device  string    // "cpu" / "cuda:0"   where the bytes physically are
}

// At returns the element at index (i, j, k) — pure arithmetic on the strides.
func (t Tensor) At(idx ...int) float32 {
	off := 0
	for axis, i := range idx {
		off += i * t.Strides[axis]
	}
	return t.Data[off]
}

Stride is the one people skip and shouldn’t. Memory is one-dimensional. A tensor’s shape is a fiction imposed on a flat buffer; strides are the translation. For shape (2,3,4) stored contiguously:

index (i,j,k)  ->  flat offset  =  i*12 + j*4 + k*1

Why you care: transpose() in PyTorch does not move data — it just swaps strides, producing a non-contiguous view. Some kernels require contiguous memory and will silently trigger a copy (costing bandwidth), or refuse outright. Understanding stride is how you predict when a .T costs nothing and when it costs a full tensor copy.

The shapes of LLM inference#

Committing this to memory pays off constantly:

input_ids           (B, S)                       token integers
embeddings          (B, S, d)                    after lookup
                    ─────── the residual stream ───────
q                   (B, S, h, d_head)  →  (B, h, S, d_head)
k, v                (B, S, h_kv, d_head)
attn scores         (B, h, S, S)                 ← quadratic in S. The problem child.
attn out            (B, h, S, d_head) → (B, S, d)
ffn intermediate    (B, S, d_ff)
logits              (B, S, V)                    ← V is ~128k. Large!

Two of these are the memory villains of Section V and VII:

  • (B, h, S, S) — attention scores. At B=8, h=32, S=8192, FP16: 32 GB for one layer. FlashAttention exists to never materialize this.
  • (B, S, V) — logits. At B=8, S=8192, V=128k, FP16: 16 GB. Which is why you only compute logits for the last position during prefill.

The forward pass of one transformer layer#

x  ──┬────────────────────────────────────────┐
     │                                        │
     ▼                                        │
  RMSNorm                                     │
     │                                        │
     ▼                                        │
  q,k,v = x@Wq, x@Wk, x@Wv                    │
     │                                        │
     ▼                                        │
  attention(q, k, v)      ← reads KV cache    │
     │                                        │
     ▼                                        │
  @ Wo                                        │
     │                                        │
     ▼                                        │
    (+)  ◄──────────────────────────────────┘   residual add
     │
     ├────────────────────────────────────────┐
     ▼                                        │
  RMSNorm                                     │
     │                                        │
     ▼                                        │
  FFN:  down( silu(gate(x)) * up(x) )         │
     │                                        │
     ▼                                        │
    (+)  ◄──────────────────────────────────┘   residual add
     │
     ▼
   output x'   (same shape as input — this is why layers stack)

Repeat L times. Then a final norm, then @ W_output to get logits over the vocabulary, then sample. That is a complete LLM forward pass. Every optimization in Sections VII and XIII is a modification to some box in that diagram.


6. Under the hood#

When you write Y = X @ W, the following happens:

  1. The framework checks shapes and dtypes, and picks a kernel.
  2. If on GPU, it calls into cuBLAS (or a custom kernel), which selects a tiling strategy: how to break the big matrices into tiles that fit in shared memory and registers.
  3. A kernel launch is queued on a CUDA stream. This costs ~5-10 µs of CPU time regardless of how big the matmul is — which is why decode, with its many tiny kernels, is launch-bound until you use CUDA graphs (Section VI.07).
  4. Thousands of threads each load a tile from HBM into shared memory, do a burst of multiply-accumulates (on tensor cores if precision permits), and write results back.
  5. The GPU is asynchronous: your Python line returns immediately. The work happens later. This is why naive timing code measures nothing (Section X.07).
Python
# WRONG — measures kernel launch, not execution
t0 = time.time(); y = model(x); print(time.time() - t0)   # ~0.0001s. Meaningless.

# RIGHT
torch.cuda.synchronize(); t0 = time.time()
y = model(x)
torch.cuda.synchronize(); print(time.time() - t0)

Memorize that. You will otherwise publish a benchmark showing your model runs in 100 µs and embarrass yourself.


7. Performance implications#

Big tensor ops good, many small tensor ops bad. Each op has fixed overhead (launch, memory round-trip). One (4096, 4096) @ (4096, 4096) matmul is vastly more efficient than 4096 (4096,) @ (4096, 4096) matrix-vector products, even though the FLOP counts differ by 4096x in the other direction. This is exactly the prefill-vs-decode distinction in embryo.

Elementwise ops are pure bandwidth. a + b, relu(x), x * scale each read and write the whole tensor and do ~1 FLOP per element. Arithmetic intensity ≈ 0.08 FLOP/byte in FP32. Catastrophically memory-bound. This is why kernel fusion (Section IV.09) — doing five elementwise ops in one pass over memory — is such a large, easy win.

Shape matters for hardware. Tensor cores want dimensions that are multiples of 8 (FP16) or 16. A hidden size of 4096 runs at full speed; 4095 can fall off a cliff. Vocabulary sizes are padded for this reason.


8. Production implications#

  • Log shapes, not just latencies. When a request is slow, the first question is “what shape was it?” A p99 latency spike is usually a p99 sequence length.
  • Guard against pathological shapes. A single 128k-token request will materialize activations that OOM your server. Enforce max input length at the gateway (Section XII.08).
  • Contiguity bugs are silent performance bugs. A stray .transpose() before a kernel that needs contiguous input can insert a multi-GB copy per layer. Profilers show it as a mysterious copy_ kernel eating 30% of your time.
  • Dtype consistency matters. A single FP32 tensor sneaking into an FP16 model forces upcasting and disables tensor cores for that op.

9. Common mistakes#

Thinking in loops instead of tensors. If you call the model once per batch element in a for loop, you have probably given up 100x. The exception: genuinely sequential dependencies, like the decode loop.

Ignoring the batch dimension’s position. (B, S, d) vs (S, B, d) — both exist in the wild. Get it wrong and you get plausible-looking garbage, not an error.

Timing without synchronize(). See above. The most common benchmarking error in ML.

Materializing the attention matrix. Writing attention “textbook style” as softmax([email protected]/sqrt(d)) @ V allocates (B,h,S,S). Fine for S=128, fatal for S=8192. Use a fused attention kernel.

Computing logits for all positions during prefill. You only need the last one. Computing all of them costs S × V × d FLOPs and B×S×V memory for nothing. Frameworks handle this, but custom code often doesn’t.


10. Hands-on exercise#

A. Shapes by hand. Given d=512, h=8, d_head=64, B=4, S=128, V=32000, write down the shape of every tensor in the layer diagram in section 5, and its FP16 size in MB. Which is largest?

B. Strides. Build the tensor metadata yourself:

Go
package main

import "fmt"

type Tensor struct {
	Data           []float32
	Shape, Strides []int
}

func Arange(shape ...int) Tensor {
	n, strides := 1, make([]int, len(shape))
	for i := len(shape) - 1; i >= 0; i-- {
		strides[i] = n
		n *= shape[i]
	}
	data := make([]float32, n)
	for i := range data {
		data[i] = float32(i)
	}
	return Tensor{data, shape, strides}
}

// Transpose swaps two axes WITHOUT touching Data: only the metadata changes.
func (t Tensor) Transpose(a, b int) Tensor {
	shape, strides := append([]int{}, t.Shape...), append([]int{}, t.Strides...)
	shape[a], shape[b] = shape[b], shape[a]
	strides[a], strides[b] = strides[b], strides[a]
	return Tensor{t.Data, shape, strides}
}

// IsContiguous: does walking the last axis fastest visit Data in order?
func (t Tensor) IsContiguous() bool {
	want := 1
	for i := len(t.Shape) - 1; i >= 0; i-- {
		if t.Strides[i] != want {
			return false
		}
		want *= t.Shape[i]
	}
	return true
}

func main() {
	t := Arange(2, 3, 4)
	fmt.Println(t.Strides, t.IsContiguous()) // [12 4 1] true
	u := t.Transpose(1, 2)
	fmt.Println(u.Shape, u.Strides, u.IsContiguous()) // [2 4 3] [12 1 4] false
}

Explain each output. Then add a Sum() method that walks the tensor in logical order using Strides, and time it on a large contiguous tensor vs its transpose. Why is one slower?

C. The batching effect. Time X @ W on GPU for W of shape (4096,4096) and X of shape (N,4096) for N in [1, 2, 4, 8, …, 512]. Plot latency vs N. You should see latency almost flat for small N and then linear. Explain the flat region. (This plot is the single most important empirical fact in LLM serving; you will re-derive it analytically in file 08.)

D. The wrong benchmark. Time a matmul without synchronize() and with it. Report both. Never make that mistake again.


11. Interview questions#

  1. What is a stride, and when does transpose() cost you memory bandwidth?
  2. Why is (B, h, S, S) a problem, and what does FlashAttention do about it?
  3. Given a batch of 1 vs 64 through the same linear layer, how do FLOPs and bytes-read change?
  4. Why must you call torch.cuda.synchronize() when benchmarking?
  5. Why are hidden sizes always round numbers like 4096, 5120, 8192?

12. Further reading#

  • [REFERENCE] PyTorch tensor internals blog post (ezyang) — the best explanation of strides
  • [FUNDAMENTAL] 3Blue1Brown, Essence of Linear Algebra, videos 3-4
  • [REFERENCE] NumPy broadcasting rules documentation
  • Next: 05 — Latency, throughput, metrics

↑↓ navigate↵ openesc close