-- | Multi-head self-attention with causal masking.
module Circuit.LLM.Attention
  ( -- * Core attention
    scaledDotProductAttention,
    multiHeadAttention,

    -- * Components
    softmax,
    causalMask,

    -- * Dimensions
    splitHeads,
    mergeHeads,
  )
where

import Data.Foldable (maximum, sum)
import Data.List (foldl1')
import Data.Vector.Unboxed qualified as V
import Harpie.Array
  ( Array,
    array,
    concatenate,
    drop,
    expand,
    mult,
    reduces,
    reshape,
    shape,
    take,
    transpose,
    zipWith,
  )
import Harpie.Array qualified as H (imap)
import NumHask.Algebra.Additive (Additive)
import NumHask.Algebra.Multiplicative (Multiplicative)
import Prelude hiding (drop, maximum, sum, take, zipWith)

-- | Numerically stable softmax along dimension 0 (row-wise).
softmax :: (Floating a, Ord a, Additive a, Multiplicative a) => Array a -> Array a
softmax :: forall a.
(Floating a, Ord a, Additive a, Multiplicative a) =>
Array a -> Array a
softmax Array a
x =
  let rowMaxes :: Array a
rowMaxes = Dims -> (Array a -> a) -> Array a -> Array a
forall a b. Dims -> (Array a -> b) -> Array a -> Array b
reduces [Int
0] Array a -> a
forall a. Ord a => Array a -> a
forall (t :: * -> *) a. (Foldable t, Ord a) => t a -> a
maximum Array a
x
      expandedMax :: Array a
expandedMax = (a -> () -> a) -> Array a -> Array () -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
expand a -> () -> a
forall a b. a -> b -> a
const Array a
rowMaxes (Dims -> [()] -> Array ()
forall t a. FromVector t a => Dims -> t -> Array a
array [Int
nCols] [()]
unitVals)
      xShifted :: Array a
xShifted = (Dims -> a -> a) -> Array a -> Array a
forall a b. (Dims -> a -> b) -> Array a -> Array b
H.imap (\Dims
_ a
v -> a -> a
forall a. Floating a => a -> a
exp a
v) ((a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith (-) Array a
x Array a
expandedMax)
      rowSums :: Array a
rowSums = Dims -> (Array a -> a) -> Array a -> Array a
forall a b. Dims -> (Array a -> b) -> Array a -> Array b
reduces [Int
0] Array a -> a
forall a. Num a => Array a -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum Array a
xShifted
      expandedSums :: Array a
expandedSums = (a -> () -> a) -> Array a -> Array () -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
expand a -> () -> a
forall a b. a -> b -> a
const Array a
rowSums (Dims -> [()] -> Array ()
forall t a. FromVector t a => Dims -> t -> Array a
array [Int
nCols] [()]
unitVals)
   in (a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith a -> a -> a
forall a. Fractional a => a -> a -> a
(/) Array a
xShifted Array a
expandedSums
  where
    shapeList :: Dims
shapeList = Vector Int -> Dims
forall a. Unbox a => Vector a -> [a]
V.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
shape Array a
x)
    nCols :: Int
nCols =
      case Dims
shapeList of
        [Int
_, Int
n] -> Int
n
        Dims
_ ->
          -- Input is not two-dimensional; not a reachable state for valid use.
          [Char] -> Int
forall a. HasCallStack => [Char] -> a
error [Char]
"unreachable: softmax expects a 2-D array"
    unitVals :: [()]
unitVals = Int -> () -> [()]
forall a. Int -> a -> [a]
replicate Int
nCols ()

-- | Scaled dot-product attention.
scaledDotProductAttention ::
  (Floating a, Ord a, Additive a, Multiplicative a) =>
  Array a -> Array a -> Array a -> Array a -> Array a
scaledDotProductAttention :: forall a.
(Floating a, Ord a, Additive a, Multiplicative a) =>
Array a -> Array a -> Array a -> Array a -> Array a
scaledDotProductAttention Array a
q Array a
k Array a
v Array a
mask =
  let dk :: Double
dk = Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Dims -> Int
forall a. HasCallStack => [a] -> a
last (Vector Int -> Dims
forall a. Unbox a => Vector a -> [a]
V.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
shape Array a
k))) :: Double
      scores :: Array a
scores = (Dims -> a -> a) -> Array a -> Array a
forall a b. (Dims -> a -> b) -> Array a -> Array b
H.imap (\Dims
_ a
s -> a
s a -> a -> a
forall a. Fractional a => a -> a -> a
/ Double -> a
forall a b. (Real a, Fractional b) => a -> b
realToFrac (Double -> Double
forall a. Floating a => a -> a
sqrt Double
dk)) (Array a -> Array a -> Array a
forall a.
(Additive a, Multiplicative a) =>
Array a -> Array a -> Array a
mult Array a
q (Array a -> Array a
forall a. Array a -> Array a
transpose Array a
k))
      masked :: Array a
masked = (a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith a -> a -> a
forall a. Num a => a -> a -> a
(+) Array a
scores Array a
mask
      attn :: Array a
attn = Array a -> Array a
forall a.
(Floating a, Ord a, Additive a, Multiplicative a) =>
Array a -> Array a
softmax Array a
masked
   in Array a -> Array a -> Array a
forall a.
(Additive a, Multiplicative a) =>
Array a -> Array a -> Array a
mult Array a
attn Array a
v

-- | Causal mask.
causalMask :: Int -> Array Double
causalMask :: Int -> Array Double
causalMask Int
n =
  Dims -> [Double] -> Array Double
forall t a. FromVector t a => Dims -> t -> Array a
array
    [Int
n, Int
n]
    [ if Int
j Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
i then Double
0 else -(Double
1.0 Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
0.0)
    | Int
i <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1],
      Int
j <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
    ]

splitHeads :: Int -> Array a -> Array a
splitHeads :: forall a. Int -> Array a -> Array a
splitHeads Int
nHead Array a
x =
  let shapeList :: Dims
shapeList = Vector Int -> Dims
forall a. Unbox a => Vector a -> [a]
V.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
shape Array a
x)
      (Int
seqLen, Int
nEmbd) =
        case Dims
shapeList of
          [Int
s, Int
e] -> (Int
s, Int
e)
          Dims
_ ->
            -- Input is not two-dimensional; not a reachable state for valid use.
            [Char] -> (Int, Int)
forall a. HasCallStack => [Char] -> a
error [Char]
"unreachable: splitHeads expects a 2-D array"
      headDim :: Int
headDim = Int
nEmbd Int -> Int -> Int
forall a. Integral a => a -> a -> a
`div` Int
nHead
   in Dims -> Array a -> Array a
forall a. Dims -> Array a -> Array a
reshape [Int
nHead, Int
seqLen, Int
headDim] (Dims -> Array a -> Array a
forall a. Dims -> Array a -> Array a
reshape [Int
seqLen, Int
nHead, Int
headDim] Array a
x)

mergeHeads :: Array a -> Array a
mergeHeads :: forall a. Array a -> Array a
mergeHeads Array a
x =
  let shapeList :: Dims
shapeList = Vector Int -> Dims
forall a. Unbox a => Vector a -> [a]
V.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
shape Array a
x)
      (Int
nHead, Int
seqLen, Int
headDim) =
        case Dims
shapeList of
          [Int
h, Int
s, Int
d] -> (Int
h, Int
s, Int
d)
          Dims
_ ->
            -- Input is not three-dimensional; not a reachable state for valid use.
            [Char] -> (Int, Int, Int)
forall a. HasCallStack => [Char] -> a
error [Char]
"unreachable: mergeHeads expects a 3-D array"
   in Dims -> Array a -> Array a
forall a. Dims -> Array a -> Array a
reshape [Int
seqLen, Int
nHead Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
headDim] (Dims -> Array a -> Array a
forall a. Dims -> Array a -> Array a
reshape [Int
seqLen, Int
nHead, Int
headDim] Array a
x)

-- | Multi-head self-attention. Projects x directly with per-head weight slices
-- to avoid deep backpermute chains.
multiHeadAttention ::
  (Floating a, Ord a, Additive a, Multiplicative a) =>
  Int -> Array a -> Array a -> Array a -> Array a -> Array a -> Array a -> Array a
multiHeadAttention :: forall a.
(Floating a, Ord a, Additive a, Multiplicative a) =>
Int
-> Array a
-> Array a
-> Array a
-> Array a
-> Array a
-> Array a
-> Array a
multiHeadAttention Int
nHead Array a
x Array a
wQ Array a
wK Array a
wV Array a
wO Array a
mask =
  let shapeList :: Dims
shapeList = Vector Int -> Dims
forall a. Unbox a => Vector a -> [a]
V.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
shape Array a
x)
      (Int
seqLen, Int
nEmbd) =
        case Dims
shapeList of
          [Int
s, Int
e] -> (Int
s, Int
e)
          Dims
_ ->
            -- Input is not two-dimensional; not a reachable state for valid use.
            [Char] -> (Int, Int)
forall a. HasCallStack => [Char] -> a
error [Char]
"unreachable: multiHeadAttention expects a 2-D array"
      headDim :: Int
headDim = Int
nEmbd Int -> Int -> Int
forall a. Integral a => a -> a -> a
`div` Int
nHead

      -- Project per head: mult x with column slice [n_embd, head_dim] of weights
      headQ :: Int -> Array a
headQ Int
h = Array a -> Array a -> Array a
forall a.
(Additive a, Multiplicative a) =>
Array a -> Array a -> Array a
mult Array a
x (Int -> Int -> Array a -> Array a
forall a. Int -> Int -> Array a -> Array a
take Int
1 Int
headDim (Int -> Int -> Array a -> Array a
forall a. Int -> Int -> Array a -> Array a
drop Int
1 (Int
h Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
headDim) Array a
wQ))
      headK :: Int -> Array a
headK Int
h = Array a -> Array a -> Array a
forall a.
(Additive a, Multiplicative a) =>
Array a -> Array a -> Array a
mult Array a
x (Int -> Int -> Array a -> Array a
forall a. Int -> Int -> Array a -> Array a
take Int
1 Int
headDim (Int -> Int -> Array a -> Array a
forall a. Int -> Int -> Array a -> Array a
drop Int
1 (Int
h Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
headDim) Array a
wK))
      headV :: Int -> Array a
headV Int
h = Array a -> Array a -> Array a
forall a.
(Additive a, Multiplicative a) =>
Array a -> Array a -> Array a
mult Array a
x (Int -> Int -> Array a -> Array a
forall a. Int -> Int -> Array a -> Array a
take Int
1 Int
headDim (Int -> Int -> Array a -> Array a
forall a. Int -> Int -> Array a -> Array a
drop Int
1 (Int
h Int -> Int -> Int
forall a. Num a => a -> a -> a
* Int
headDim) Array a
wV))

      -- Per-head attention
      attnHeads :: [Array a]
attnHeads =
        [ Dims -> Array a -> Array a
forall a. Dims -> Array a -> Array a
reshape [Int
1, Int
seqLen, Int
headDim] (Array a -> Array a) -> Array a -> Array a
forall a b. (a -> b) -> a -> b
$
            Array a -> Array a -> Array a -> Array a -> Array a
forall a.
(Floating a, Ord a, Additive a, Multiplicative a) =>
Array a -> Array a -> Array a -> Array a -> Array a
scaledDotProductAttention (Int -> Array a
headQ Int
h) (Int -> Array a
headK Int
h) (Int -> Array a
headV Int
h) Array a
mask
        | Int
h <- [Int
0 .. Int
nHead Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]
        ]

      -- Stack: [n_head, seq_len, head_dim]
      attnH :: Array a
attnH = (Array a -> Array a -> Array a) -> [Array a] -> Array a
forall a. HasCallStack => (a -> a -> a) -> [a] -> a
foldl1' (Int -> Array a -> Array a -> Array a
forall a. Int -> Array a -> Array a -> Array a
concatenate Int
0) [Array a]
attnHeads
   in -- Merge: [seq_len, n_embd] and output projection
      Array a -> Array a -> Array a
forall a.
(Additive a, Multiplicative a) =>
Array a -> Array a -> Array a
mult (Array a -> Array a
forall a. Array a -> Array a
mergeHeads Array a
attnH) Array a
wO