Flash Attention in CUDA
Derive the core equations slowly, attach every symbol to code, and verify the result with small numerical examples before scaling the implementation.
How to study the mathematics
Read each chapter in four passes: intuition, symbols, derivation, and implementation. Recalculate the worked example by hand. Then change one number and predict the direction of the result before running code.
Tiling
Intuition before notation
Tiling is a transformation whose meaning comes from its domain, codomain, objective, and invariants. The equation is useful only when every symbol maps to a concrete tensor, state, or measurement.
Symbol dictionary
- x,y
- input and target
- f_theta
- parameterized transformation
- ell
- data objective
- Omega
- regularizer
- lambda
- regularization weight
Derive it one move at a time
- 1
Specify the input representation.
- 2
Define the parameterized transformation.
- 3
Choose a loss connected to desired behavior.
- 4
Average over the training evidence.
- 5
Add explicit inductive bias or constraints.
Worked numerical example
For a one-parameter predictor f(x)=theta*x with x=2, y=6, squared loss is (2theta-6)^2. The minimum without regularization is theta=3.
Translate the derivation into code
- Give every axis a semantic name.
- Implement a scalar reference first.
- Compare optimized output and gradients against the reference.
Online softmax
Intuition before notation
Softmax converts relative logits into positive normalized probabilities while temperature controls how strongly differences are expressed.
Symbol dictionary
- z_i
- logit for outcome i
- m
- maximum logit
- tau
- temperature
- p_i
- normalized probability
Derive it one move at a time
- 1
Divide logits by temperature.
- 2
Find the maximum scaled logit.
- 3
Subtract it without changing probability ratios.
- 4
Exponentiate the shifted values.
- 5
Divide by their sum.
Worked numerical example
Logits [1000,999] overflow naively. Subtracting 1000 gives [0,-1], whose probabilities are approximately [0.731,0.269].
Translate the derivation into code
- Reduce maximum with keepdims.
- Use the same axis for maximum and sum.
- Test translation invariance by adding a constant to every logit.
Memory hierarchy
Intuition before notation
Memory hierarchy is a transformation whose meaning comes from its domain, codomain, objective, and invariants. The equation is useful only when every symbol maps to a concrete tensor, state, or measurement.
Symbol dictionary
- x,y
- input and target
- f_theta
- parameterized transformation
- ell
- data objective
- Omega
- regularizer
- lambda
- regularization weight
Derive it one move at a time
- 1
Specify the input representation.
- 2
Define the parameterized transformation.
- 3
Choose a loss connected to desired behavior.
- 4
Average over the training evidence.
- 5
Add explicit inductive bias or constraints.
Worked numerical example
For a one-parameter predictor f(x)=theta*x with x=2, y=6, squared loss is (2theta-6)^2. The minimum without regularization is theta=3.
Translate the derivation into code
- Give every axis a semantic name.
- Implement a scalar reference first.
- Compare optimized output and gradients against the reference.
Optimization and parameter updates
Intuition before notation
objective and gradient update is a transformation whose meaning comes from its domain, codomain, objective, and invariants. The equation is useful only when every symbol maps to a concrete tensor, state, or measurement.
Symbol dictionary
- x,y
- input and target
- f_theta
- parameterized transformation
- ell
- data objective
- Omega
- regularizer
- lambda
- regularization weight
Derive it one move at a time
- 1
Specify the input representation.
- 2
Define the parameterized transformation.
- 3
Choose a loss connected to desired behavior.
- 4
Average over the training evidence.
- 5
Add explicit inductive bias or constraints.
Worked numerical example
For a one-parameter predictor f(x)=theta*x with x=2, y=6, squared loss is (2theta-6)^2. The minimum without regularization is theta=3.
Translate the derivation into code
- Give every axis a semantic name.
- Implement a scalar reference first.
- Compare optimized output and gradients against the reference.
Probability, normalization, and calibration
Intuition before notation
Softmax converts relative logits into positive normalized probabilities while temperature controls how strongly differences are expressed.
Symbol dictionary
- z_i
- logit for outcome i
- m
- maximum logit
- tau
- temperature
- p_i
- normalized probability
Derive it one move at a time
- 1
Divide logits by temperature.
- 2
Find the maximum scaled logit.
- 3
Subtract it without changing probability ratios.
- 4
Exponentiate the shifted values.
- 5
Divide by their sum.
Worked numerical example
Logits [1000,999] overflow naively. Subtracting 1000 gives [0,-1], whose probabilities are approximately [0.731,0.269].
Translate the derivation into code
- Reduce maximum with keepdims.
- Use the same axis for maximum and sum.
- Test translation invariance by adding a constant to every logit.
Evaluation uncertainty and error bars
Intuition before notation
A reported metric is an estimate from finite evidence. Uncertainty separates stable improvement from sampling noise.
Symbol dictionary
- x_i
- per-example or per-run measurement
- x-bar
- sample mean
- s
- sample standard deviation
- n
- independent observations
Derive it one move at a time
- 1
Choose the independent unit.
- 2
Compute one measurement per unit.
- 3
Estimate mean and sample variance.
- 4
Convert variation into standard error.
- 5
Report an interval with assumptions or use bootstrap resampling.
Worked numerical example
Five seeded scores with mean 0.80 and standard deviation 0.04 have SE about 0.0179, giving a rough 95% interval 0.765 to 0.835.
Translate the derivation into code
- Store per-example and per-seed values, not only the average.
- Use stratified or paired intervals when the design requires them.
- Never treat correlated tokens or timesteps as independent runs.