{-# LANGUAGE NamedFieldPuns #-}
{-# LANGUAGE PatternSynonyms #-}

-- | Differentiable PCA operations.
--
-- The fitted 'PCAModel' is treated as constant: this module differentiates
-- through the /use/ of a PCA (centering, projection, reconstruction, loss),
-- not yet through the spectral fit itself.
--
-- Each operation is packaged as a 'Circuit.Diff.Diff' arrow, so it composes
-- with the rest of the circuits-ad stack.
module Circuit.PCA.Diff
  ( -- * Differentiable linear maps
    scoresD,
    reconstructD,
    projectRowsD,

    -- * Differentiable loss
    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, (.))

-- | Subtract a constant mean row.  The mean is a parameter, not an input.
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"

-- | Add a constant mean row.
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"

-- | Multiply on the right by a fixed matrix.
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))

-- | Differentiable scores: center, then project onto principal axes.
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

-- | Differentiable reconstruction from scores.
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)

-- | Full differentiable project-through-model: @reconstruct . scores@.
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

-- | Mean-squared reconstruction error, differentiable with respect to the
-- input array.
--
-- Forward: @loss = (1/n) * ||x - projectRows model x||²@
-- Pullback: @dloss/dx = (2/n) * (x - projectRows model x)@
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)
            )