module Circuit.LLM.Attention
(
scaledDotProductAttention,
multiHeadAttention,
softmax,
causalMask,
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)
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
_ ->
[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 ()
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
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
_ ->
[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
_ ->
[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)
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
_ ->
[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
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))
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]
]
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
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