{-# LANGUAGE NamedFieldPuns #-}
{-# LANGUAGE PatternSynonyms #-}
module Circuit.PCA.Diff
(
scoresD,
reconstructD,
projectRowsD,
reconstructionLossD,
)
where
import Circuit.Diff.Circuit (Diff, Diff', pattern Diff)
import Circuit.PCA (PCAModel (..), projectRows)
import Circuit.PCA.Lin (fromMatrix, multM, toMatrix, transpose2)
import Control.Category
import Harpie.Array.Storable (Array)
import Numeric.LinearAlgebra qualified as LA
import Prelude hiding (id, (.))
subtractMeanD :: Array Double -> Diff' (Array Double) (Array Double)
subtractMeanD :: Array Double -> Diff' (Array Double) (Array Double)
subtractMeanD Array Double
mean =
(Array Double -> (Array Double, Array Double -> Array Double))
-> Diff' (Array Double) (Array Double)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((Array Double -> (Array Double, Array Double -> Array Double))
-> Diff' (Array Double) (Array Double))
-> (Array Double -> (Array Double, Array Double -> Array Double))
-> Diff' (Array Double) (Array Double)
forall a b. (a -> b) -> a -> b
$ \Array Double
x ->
case (Array Double -> Maybe (Matrix Double)
toMatrix Array Double
x, Array Double -> Maybe (Matrix Double)
toMatrix Array Double
mean) of
(Just Matrix Double
m, Just Matrix Double
muM) ->
let muFlat :: Vector Double
muFlat = Matrix Double -> Vector Double
forall t. Element t => Matrix t -> Vector t
LA.flatten Matrix Double
muM
xc :: Array Double
xc = Matrix Double -> Array Double
fromMatrix (Matrix Double
m Matrix Double -> Matrix Double -> Matrix Double
forall a. Num a => a -> a -> a
- Vector Double -> Matrix Double
forall a. Storable a => Vector a -> Matrix a
LA.asRow Vector Double
muFlat)
in (Array Double
xc, Array Double -> Array Double
forall a. a -> a
forall {k} (cat :: k -> k -> *) (a :: k). Category cat => cat a a
id)
(Maybe (Matrix Double), Maybe (Matrix Double))
_ -> [Char] -> (Array Double, Array Double -> Array Double)
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.PCA.Diff.subtractMeanD: rank-2 input and rank-1 mean expected"
addMeanD :: Array Double -> Diff' (Array Double) (Array Double)
addMeanD :: Array Double -> Diff' (Array Double) (Array Double)
addMeanD Array Double
mean =
(Array Double -> (Array Double, Array Double -> Array Double))
-> Diff' (Array Double) (Array Double)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((Array Double -> (Array Double, Array Double -> Array Double))
-> Diff' (Array Double) (Array Double))
-> (Array Double -> (Array Double, Array Double -> Array Double))
-> Diff' (Array Double) (Array Double)
forall a b. (a -> b) -> a -> b
$ \Array Double
x ->
case (Array Double -> Maybe (Matrix Double)
toMatrix Array Double
x, Array Double -> Maybe (Matrix Double)
toMatrix Array Double
mean) of
(Just Matrix Double
m, Just Matrix Double
muM) ->
let muFlat :: Vector Double
muFlat = Matrix Double -> Vector Double
forall t. Element t => Matrix t -> Vector t
LA.flatten Matrix Double
muM
xbar :: Array Double
xbar = Matrix Double -> Array Double
fromMatrix (Matrix Double
m Matrix Double -> Matrix Double -> Matrix Double
forall a. Num a => a -> a -> a
+ Vector Double -> Matrix Double
forall a. Storable a => Vector a -> Matrix a
LA.asRow Vector Double
muFlat)
in (Array Double
xbar, Array Double -> Array Double
forall a. a -> a
forall {k} (cat :: k -> k -> *) (a :: k). Category cat => cat a a
id)
(Maybe (Matrix Double), Maybe (Matrix Double))
_ -> [Char] -> (Array Double, Array Double -> Array Double)
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.PCA.Diff.addMeanD: rank-2 input and rank-1 mean expected"
multMD :: Array Double -> Diff' (Array Double) (Array Double)
multMD :: Array Double -> Diff' (Array Double) (Array Double)
multMD Array Double
w =
(Array Double -> (Array Double, Array Double -> Array Double))
-> Diff' (Array Double) (Array Double)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((Array Double -> (Array Double, Array Double -> Array Double))
-> Diff' (Array Double) (Array Double))
-> (Array Double -> (Array Double, Array Double -> Array Double))
-> Diff' (Array Double) (Array Double)
forall a b. (a -> b) -> a -> b
$ \Array Double
x ->
let y :: Array Double
y = Array Double -> Array Double -> Array Double
multM Array Double
x Array Double
w
in (Array Double
y, \Array Double
dy -> Array Double -> Array Double -> Array Double
multM Array Double
dy (Array Double -> Array Double
transpose2 Array Double
w))
scoresD :: PCAModel -> Diff' (Array Double) (Array Double)
scoresD :: PCAModel -> Diff' (Array Double) (Array Double)
scoresD PCAModel {Array Double
pcaMean :: Array Double
pcaMean :: PCAModel -> Array Double
pcaMean, Array Double
pcaComponents :: Array Double
pcaComponents :: PCAModel -> Array Double
pcaComponents} =
Array Double -> Diff' (Array Double) (Array Double)
multMD Array Double
pcaComponents Diff' (Array Double) (Array Double)
-> Diff' (Array Double) (Array Double)
-> Diff' (Array Double) (Array Double)
forall b c a. Diff () b c -> Diff () a b -> Diff () a c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. Array Double -> Diff' (Array Double) (Array Double)
subtractMeanD Array Double
pcaMean
reconstructD :: PCAModel -> Diff' (Array Double) (Array Double)
reconstructD :: PCAModel -> Diff' (Array Double) (Array Double)
reconstructD PCAModel {Array Double
pcaMean :: PCAModel -> Array Double
pcaMean :: Array Double
pcaMean, Array Double
pcaComponents :: PCAModel -> Array Double
pcaComponents :: Array Double
pcaComponents} =
Array Double -> Diff' (Array Double) (Array Double)
addMeanD Array Double
pcaMean Diff' (Array Double) (Array Double)
-> Diff' (Array Double) (Array Double)
-> Diff' (Array Double) (Array Double)
forall b c a. Diff () b c -> Diff () a b -> Diff () a c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. Array Double -> Diff' (Array Double) (Array Double)
multMD (Array Double -> Array Double
transpose2 Array Double
pcaComponents)
projectRowsD :: PCAModel -> Diff' (Array Double) (Array Double)
projectRowsD :: PCAModel -> Diff' (Array Double) (Array Double)
projectRowsD PCAModel
model = PCAModel -> Diff' (Array Double) (Array Double)
reconstructD PCAModel
model Diff' (Array Double) (Array Double)
-> Diff' (Array Double) (Array Double)
-> Diff' (Array Double) (Array Double)
forall b c a. Diff () b c -> Diff () a b -> Diff () a c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. PCAModel -> Diff' (Array Double) (Array Double)
scoresD PCAModel
model
reconstructionLossD :: PCAModel -> Diff' (Array Double) Double
reconstructionLossD :: PCAModel -> Diff' (Array Double) Double
reconstructionLossD PCAModel
model =
(Array Double -> (Double, Double -> Array Double))
-> Diff' (Array Double) Double
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((Array Double -> (Double, Double -> Array Double))
-> Diff' (Array Double) Double)
-> (Array Double -> (Double, Double -> Array Double))
-> Diff' (Array Double) Double
forall a b. (a -> b) -> a -> b
$ \Array Double
x ->
case Array Double -> Maybe (Matrix Double)
toMatrix Array Double
x of
Maybe (Matrix Double)
Nothing -> [Char] -> (Double, Double -> Array Double)
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.PCA.Diff.reconstructionLossD: rank-2 input expected"
Just Matrix Double
m ->
let xhM :: Matrix Double
xhM = case Array Double -> Maybe (Matrix Double)
toMatrix (PCAModel -> Array Double -> Array Double
projectRows PCAModel
model Array Double
x) of Just Matrix Double
m' -> Matrix Double
m'; Maybe (Matrix Double)
Nothing -> [Char] -> Matrix Double
forall a. HasCallStack => [Char] -> a
error [Char]
"projectRows not rank-2"
rM :: Matrix Double
rM = Matrix Double
m Matrix Double -> Matrix Double -> Matrix Double
forall a. Num a => a -> a -> a
- Matrix Double
xhM
n :: Double
n = Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Matrix Double -> Int
forall t. Matrix t -> Int
LA.rows Matrix Double
m Int -> Int -> Int
forall a. Num a => a -> a -> a
* Matrix Double -> Int
forall t. Matrix t -> Int
LA.cols Matrix Double
m) :: Double
loss :: Double
loss = Matrix Double -> Double
forall (c :: * -> *) e. Container c e => c e -> e
LA.sumElements ((Double -> Double) -> Matrix Double -> Matrix Double
forall b (c :: * -> *) e.
(Element b, Container c e) =>
(e -> b) -> c e -> c b
LA.cmap (\Double
e -> Double
e Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
e) Matrix Double
rM) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
n
in ( Double
loss,
\Double
dl -> Matrix Double -> Array Double
fromMatrix ((Double -> Double) -> Matrix Double -> Matrix Double
forall b (c :: * -> *) e.
(Element b, Container c e) =>
(e -> b) -> c e -> c b
LA.cmap (\Double
e -> Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
dl Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
n Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
e) Matrix Double
rM)
)