causal_attention
std.nn.causal_attention · Level L2Causal (decoder) self-attention: row i attends only to rows j ≤ i. The mask compares positions from Iota (j > i is the future); future scores are replaced by the row minimum before the stable softmax, so nothing can overflow, and their weights are set to zero after it.
softmax over j ≤ i of (Q Kᵀ/√d)ᵢⱼ, then · V
Signature
causal_attention(Q: f64[n, d], K: f64[n, d], V: f64[n, d]) → f64[n, d]
Structure
The function as NOVA stores it: one box per input, operation and output, and arrows that carry values. A double border marks another library function this one runs — called once, or by Scan once per element; select it to open that function.
- input
- operation
- constant
- call
- output
Verification
- Signature proven by NOVA’s shape solver, for every size.
- Agrees with the reference
for each row i: softmax((Q @ K.T)[i, :i+1] / sqrt(d)) @ V[:i+1]to 80 digits (100-digit arithmetic), on all 40 test cases. - All 643 float64 results inside the running error bound; the closest uses 7% of it.
- Interpreter and NumPy backend return bit-identical results.
Accuracy in detail
- correctly rounded (the float64 nearest the exact value)
- 56%
- bit-equal to the NumPy formula in float64
- 77%
- largest error, in units in the last place
- 87
Large ulp counts appear only where cancellation drives a result toward zero; the absolute error is still inside the bound.
Identity
Calls
—
Called by
—
sha256:59d015d39695c7167bcbbecf68f41331d9cee191db5703ea63bd3a3c3e9bce82The semantic hash of the graph. It changes when the program changes, and never when only its documentation does.