| Safe Haskell | None |
|---|---|
| Language | GHC2024 |
Circuit.LLM.Attention
Description
Multi-head self-attention with causal masking.
Synopsis
- scaledDotProductAttention :: (Floating a, Ord a, Additive a, Multiplicative a) => Array a -> Array a -> Array a -> Array a -> Array a
- multiHeadAttention :: (Floating a, Ord a, Additive a, Multiplicative a) => Int -> Array a -> Array a -> Array a -> Array a -> Array a -> Array a -> Array a
- softmax :: (Floating a, Ord a, Additive a, Multiplicative a) => Array a -> Array a
- causalMask :: Int -> Array Double
- splitHeads :: Int -> Array a -> Array a
- mergeHeads :: Array a -> Array a
Core attention
scaledDotProductAttention :: (Floating a, Ord a, Additive a, Multiplicative a) => Array a -> Array a -> Array a -> Array a -> Array a Source #
Scaled dot-product attention.
multiHeadAttention :: (Floating a, Ord a, Additive a, Multiplicative a) => Int -> Array a -> Array a -> Array a -> Array a -> Array a -> Array a -> Array a Source #
Multi-head self-attention. Projects x directly with per-head weight slices to avoid deep backpermute chains.
Components
softmax :: (Floating a, Ord a, Additive a, Multiplicative a) => Array a -> Array a Source #
Numerically stable softmax along dimension 0 (row-wise).
Dimensions
mergeHeads :: Array a -> Array a Source #