{-# LANGUAGE DataKinds #-}
{-# LANGUAGE UndecidableInstances #-}
{-# OPTIONS_GHC -Wno-orphans #-}

-- | Bridging parameterised reverse-mode AD ('DiffP') and polynomial monomial
-- lenses.
--
-- A 'DiffP p a b' is a /parametric/ lens: for every parameter value @p@ it
-- gives an ordinary lens @Mono a a -> Mono b b@, and it additionally produces
-- a parameter gradient @dp@. In 'Poly' this is naturally expressed via the
-- copower 'Depend': a parameter-indexed family of lenses.
--
-- This gives the "shared parameter" reading of 'DiffP' a home in 'Poly': the
-- parameter is the index of a lens family, not an extra layer. The
-- paired-parameter reading ('Circuit.Diff.Param.splitP', 'joinP') remains
-- available for independent-layer composition.
module Circuit.Poly.DiffP
  ( -- * DiffP as a parameter-indexed lens family
    diffPAt,
    diffPAsFamily,
    diffPParamGrad,

    -- * Recover a DiffP from its Poly decomposition
    diffPFromFamily,

    -- * Star-based feedback trace
    traceDiffPFrom,
    traceDiffPD,
    traceDiffPMatrix,
  )
where

import Circuit.Bimonoid (MergeZero, Zero (..))
import Circuit.Category (Category (..))
import Circuit.Channel (Channel (..), Strength (..))
import Circuit.Diff.Param (DiffP (..))
import Circuit.Mat.Dense (Matrix, fromLists, matVec, starMatrix, toLists)
import Circuit.Poly (Mono, Morphism (..), Poly (..), applyLens, lens)
import NumHask.Algebra.Additive qualified as NHA
import NumHask.Algebra.Multiplicative qualified as NHM
import NumHask.Algebra.Ring qualified as NHR
import NumHask.Free.Carriers (FieldStar (..))
import Prelude hiding (id, (.))

-- | The lens at a fixed parameter value.
--
-- Forward: @a -> b@. Backward: @a -> db -> da@.
diffPAt :: DiffP p a b -> p -> Morphism (Mono a a) (Mono b b)
diffPAt :: forall p a b. DiffP p a b -> p -> Morphism (Mono a a) (Mono b b)
diffPAt (DiffP p -> a -> (b, b -> (a, p))
f) p
p = (a -> b) -> (a -> b -> a) -> Morphism (Mono a a) (Mono b b)
forall a b db da.
(a -> b) -> (a -> db -> da) -> Morphism (Mono da a) (Mono db b)
lens a -> b
get a -> b -> a
put
  where
    get :: a -> b
get a
a = (b, b -> (a, p)) -> b
forall a b. (a, b) -> a
fst (p -> a -> (b, b -> (a, p))
f p
p a
a)
    put :: a -> b -> a
put a
a b
db = (a, p) -> a
forall a b. (a, b) -> a
fst ((b, b -> (a, p)) -> b -> (a, p)
forall a b. (a, b) -> b
snd (p -> a -> (b, b -> (a, p))
f p
p a
a) b
db)

-- | The parameter gradient extracted from a 'DiffP'.
--
-- For a parameter @p@, input @a@ and output cotangent @db@, return @dp@.
diffPParamGrad :: DiffP p a b -> p -> a -> b -> p
diffPParamGrad :: forall p a b. DiffP p a b -> p -> a -> b -> p
diffPParamGrad (DiffP p -> a -> (b, b -> (a, p))
f) p
p a
a b
db = (a, p) -> p
forall a b. (a, b) -> b
snd ((b, b -> (a, p)) -> b -> (a, p)
forall a b. (a, b) -> b
snd (p -> a -> (b, b -> (a, p))
f p
p a
a) b
db)

-- | View a 'DiffP' as a 'Poly' morphism @Const p * Mono a a -> Mono b b@.
--
-- The parameter is carried by the constant functor; choosing a position @p@
-- selects one lens from the family.
diffPAsFamily :: DiffP p a b -> Morphism ('Prod ('Const p) (Mono a a)) (Mono b b)
diffPAsFamily :: forall p a b.
DiffP p a b -> Morphism ('Prod ('Const p) (Mono a a)) (Mono b b)
diffPAsFamily DiffP p a b
d = (p -> Morphism (Mono a a) (Mono b b))
-> Morphism ('Prod ('Const p) (Mono a a)) (Mono b b)
forall a (p1 :: Poly) (q :: Poly).
(a -> Morphism p1 q) -> Morphism ('Prod ('Const a) p1) q
Depend (DiffP p a b -> p -> Morphism (Mono a a) (Mono b b)
forall p a b. DiffP p a b -> p -> Morphism (Mono a a) (Mono b b)
diffPAt DiffP p a b
d)

-- | Recover a 'DiffP' from its fixed-parameter lens family and parameter
-- gradient. This is the inverse of the 'diffPAt'/'diffPParamGrad' split.
diffPFromFamily ::
  (p -> Morphism (Mono a a) (Mono b b)) ->
  (p -> a -> b -> p) ->
  DiffP p a b
diffPFromFamily :: forall p a b.
(p -> Morphism (Mono a a) (Mono b b))
-> (p -> a -> b -> p) -> DiffP p a b
diffPFromFamily p -> Morphism (Mono a a) (Mono b b)
family p -> a -> b -> p
grad = (p -> a -> (b, b -> (a, p))) -> DiffP p a b
forall p a b. (p -> a -> (b, b -> (a, p))) -> DiffP p a b
DiffP ((p -> a -> (b, b -> (a, p))) -> DiffP p a b)
-> (p -> a -> (b, b -> (a, p))) -> DiffP p a b
forall a b. (a -> b) -> a -> b
$ \p
p a
a ->
  let (b
b, b -> a
put) = Morphism (Mono a a) (Mono b b) -> a -> (b, b -> a)
forall da a db b.
Morphism (Mono da a) (Mono db b) -> a -> (b, db -> da)
applyLens (p -> Morphism (Mono a a) (Mono b b)
family p
p) a
a
   in (b
b, \b
db -> (b -> a
put b
db, p -> a -> b -> p
grad p
p a
a b
db))

-- | Cartesian strength for 'DiffP'.
--
-- Threads a plain morphism through the feedback channel, copying the
-- parameter gradient unchanged.
instance (MergeZero (->) p) => Strength (,) (DiffP p) where
  strength :: forall b c a. DiffP p b c -> DiffP p (a, b) (a, c)
strength (DiffP p -> b -> (c, c -> (b, p))
f) = (p -> (a, b) -> ((a, c), (a, c) -> ((a, b), p)))
-> DiffP p (a, b) (a, c)
forall p a b. (p -> a -> (b, b -> (a, p))) -> DiffP p a b
DiffP ((p -> (a, b) -> ((a, c), (a, c) -> ((a, b), p)))
 -> DiffP p (a, b) (a, c))
-> (p -> (a, b) -> ((a, c), (a, c) -> ((a, b), p)))
-> DiffP p (a, b) (a, c)
forall a b. (a -> b) -> a -> b
$ \p
p0 (a
a, b
b) ->
    let (c
c, c -> (b, p)
back) = p -> b -> (c, c -> (b, p))
f p
p0 b
b
     in ( (a
a, c
c),
          \(a
da, c
dc) ->
            let (b
db, p
dp) = c -> (b, p)
back c
dc
             in ((a
da, b
db), p
dp)
        )

-- | Star-based trace for 'DiffP'.
--
-- The forward pass iterates the state channel to a fixed point; the backward
-- pass solves the feedback adjoint using the Kleene star of the channel
-- self-coupling.  This is the Schur-complement view of backpropagation
-- through feedback: for a body @s' = f(s, i), o = g(s, i)@ linearised at
-- the fixed point, the closed gradient is
--
-- > do/di = D + C · star(A) · B
--
-- where @A = ∂s'/∂s@, @B = ∂s'/∂i@, @C = ∂o/∂s@, @D = ∂o/∂i@.
-- The same star appears in the parameter gradient.
traceDiffPFrom ::
  (NHR.StarSemiring j, MergeZero (->) o) =>
  -- | forward seed for the state channel
  j ->
  -- | number of forward iterations
  Int ->
  DiffP p (j, i) (j, o) ->
  DiffP p i o
traceDiffPFrom :: forall j o p i.
(StarSemiring j, MergeZero (->) o) =>
j -> Int -> DiffP p (j, i) (j, o) -> DiffP p i o
traceDiffPFrom j
j0 Int
n (DiffP p -> (j, i) -> ((j, o), (j, o) -> ((j, i), p))
body) = (p -> i -> (o, o -> (i, p))) -> DiffP p i o
forall p a b. (p -> a -> (b, b -> (a, p))) -> DiffP p a b
DiffP ((p -> i -> (o, o -> (i, p))) -> DiffP p i o)
-> (p -> i -> (o, o -> (i, p))) -> DiffP p i o
forall a b. (a -> b) -> a -> b
$ \p
p0 i
i ->
  let stepFwd :: j -> j
stepFwd j
j = let ((j
j', o
_), (j, o) -> ((j, i), p)
_) = p -> (j, i) -> ((j, o), (j, o) -> ((j, i), p))
body p
p0 (j
j, i
i) in j
j'
      a :: j
a = (j -> j) -> j -> [j]
forall a. (a -> a) -> a -> [a]
iterate j -> j
stepFwd j
j0 [j] -> Int -> j
forall a. HasCallStack => [a] -> Int -> a
!! Int
n
      ((j
_, o
o), (j, o) -> ((j, i), p)
backward) = p -> (j, i) -> ((j, o), (j, o) -> ((j, i), p))
body p
p0 (j
a, i
i)
      -- Probe the feedback self-coupling @A@ with a unit feedback cotangent.
      ((j
aJ, i
_), p
_) = (j, o) -> ((j, i), p)
backward (j
forall a. Multiplicative a => a
NHM.one, () -> o
forall {k} (arr :: * -> k -> *) (a :: k). Zero arr a => arr () a
zero ())
      aStar :: j
aStar = j -> j
forall a. StarSemiring a => a -> a
NHR.star j
aJ
      pullback :: o -> (i, p)
pullback o
do_ =
        let -- Cross-coupling @B · do_@ for this particular output cotangent.
            cdc :: j
cdc = (j, i) -> j
forall a b. (a, b) -> a
fst (((j, i), p) -> (j, i)
forall a b. (a, b) -> a
fst ((j, o) -> ((j, i), p)
backward (j
forall a. Additive a => a
NHA.zero, o
do_)))
            dj :: j
dj = j
aStar j -> j -> j
forall a. Multiplicative a => a -> a -> a
NHM.* j
cdc
            ((j
_, i
di), p
dp) = (j, o) -> ((j, i), p)
backward (j
dj, o
do_)
         in (i
di, p
dp)
   in (o
o, o -> (i, p)
pullback)

-- | 'traceDiffPFrom' specialised to a scalar 'Double' state channel.
--
-- The primal is iterated until @|s' - s| <= tol@ or @maxIter@ is reached.
-- The feedback Jacobian @A@ is then probed and guarded: @|A| >= 1@ is
-- rejected with an error, because outside the contractive regime the star
-- @1/(1-A)@ either diverges or inverts the sign.  This closes the §7
-- silent-failure gap for the scalar case.
traceDiffPD ::
  (MergeZero (->) o) =>
  -- | forward seed for the state channel
  Double ->
  -- | residual tolerance for the primal fixed point
  Double ->
  -- | maximum number of forward iterations
  Int ->
  DiffP p (Double, i) (Double, o) ->
  DiffP p i o
traceDiffPD :: forall o p i.
MergeZero (->) o =>
Double
-> Double -> Int -> DiffP p (Double, i) (Double, o) -> DiffP p i o
traceDiffPD Double
j0 Double
tol Int
maxIter (DiffP p -> (Double, i) -> ((Double, o), (Double, o) -> ((Double, i), p))
body) = (p -> i -> (o, o -> (i, p))) -> DiffP p i o
forall p a b. (p -> a -> (b, b -> (a, p))) -> DiffP p a b
DiffP ((p -> i -> (o, o -> (i, p))) -> DiffP p i o)
-> (p -> i -> (o, o -> (i, p))) -> DiffP p i o
forall a b. (a -> b) -> a -> b
$ \p
p0 i
i ->
  let stepFwd :: Double -> Double
stepFwd Double
j = let ((Double
j', o
_), (Double, o) -> ((Double, i), p)
_) = p -> (Double, i) -> ((Double, o), (Double, o) -> ((Double, i), p))
body p
p0 (Double
j, i
i) in Double
j'
      go :: Int -> Double -> Double
go Int
0 Double
j = Double
j
      go Int
n Double
j =
        let j' :: Double
j' = Double -> Double
stepFwd Double
j
         in if Double -> Double
forall a. Num a => a -> a
abs (Double
j' Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
j) Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
tol then Double
j' else Int -> Double -> Double
go (Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) Double
j'
      a :: Double
a = Int -> Double -> Double
go Int
maxIter Double
j0
      aNext :: Double
aNext = Double -> Double
stepFwd Double
a
      ((Double
_, o
o), (Double, o) -> ((Double, i), p)
backward) = p -> (Double, i) -> ((Double, o), (Double, o) -> ((Double, i), p))
body p
p0 (Double
a, i
i)
      -- Probe the feedback self-coupling @A@ with a unit feedback cotangent.
      ((Double
aJ, i
_), p
_) = (Double, o) -> ((Double, i), p)
backward (Double
1.0, () -> o
forall {k} (arr :: * -> k -> *) (a :: k). Zero arr a => arr () a
zero ())
      aStar :: Double
aStar = Double -> Double
forall a. Fractional a => a -> a
recip (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
aJ)
      pullback :: o -> (i, p)
pullback o
do_ =
        let -- Cross-coupling @B · do_@ for this particular output cotangent.
            cdc :: Double
cdc = (Double, i) -> Double
forall a b. (a, b) -> a
fst (((Double, i), p) -> (Double, i)
forall a b. (a, b) -> a
fst ((Double, o) -> ((Double, i), p)
backward (Double
0.0, o
do_)))
            dj :: Double
dj = Double
aStar Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
cdc
            ((Double
_, i
di), p
dp) = (Double, o) -> ((Double, i), p)
backward (Double
dj, o
do_)
         in (i
di, p
dp)
   in if Double -> Double
forall a. Num a => a -> a
abs (Double
aNext Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
a) Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
tol
        then [Char] -> (o, o -> (i, p))
forall a. HasCallStack => [Char] -> a
error ([Char]
"traceDiffPD: primal fixed point did not converge within " [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ Int -> [Char]
forall a. Show a => a -> [Char]
show Int
maxIter [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ [Char]
" iterations")
        else
          if Double -> Double
forall a. Num a => a -> a
abs Double
aJ Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
>= Double
1.0
            then [Char] -> (o, o -> (i, p))
forall a. HasCallStack => [Char] -> a
error ([Char]
"traceDiffPD: feedback Jacobian |" [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ Double -> [Char]
forall a. Show a => a -> [Char]
show Double
aJ [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ [Char]
"| >= 1 is outside the contractive regime")
            else (o
o, o -> (i, p)
pullback)

-- | Star-based trace for a vector-channel 'DiffP'.
--
-- The state channel is a list @[[Double]]@ of fixed dimension.  The forward pass
-- iterates to a fixed point; the backward pass probes the feedback Jacobian
-- column by column, builds a 'Matrix', and solves the adjoint with
-- 'starMatrix'.  Each column is wrapped in 'FieldStar' so the matrix star is
-- honest @(I − A)⁻¹@.
--
-- This is the multi-agent extension of 'traceDiffPD': instead of a scalar
-- self-coupling @a@, the feedback Jacobian is a matrix @A@, and the star is
-- the Neumann series @(I − A)⁻¹@.
traceDiffPMatrix ::
  (MergeZero (->) o) =>
  -- | forward seed for the state channel (its length is the channel dimension)
  [Double] ->
  -- | residual tolerance for the primal fixed point
  Double ->
  -- | maximum number of forward iterations
  Int ->
  DiffP p ([Double], i) ([Double], o) ->
  DiffP p i o
traceDiffPMatrix :: forall o p i.
MergeZero (->) o =>
[Double]
-> Double
-> Int
-> DiffP p ([Double], i) ([Double], o)
-> DiffP p i o
traceDiffPMatrix [Double]
x0 Double
tol Int
maxIter (DiffP p
-> ([Double], i)
-> (([Double], o), ([Double], o) -> (([Double], i), p))
body) = (p -> i -> (o, o -> (i, p))) -> DiffP p i o
forall p a b. (p -> a -> (b, b -> (a, p))) -> DiffP p a b
DiffP ((p -> i -> (o, o -> (i, p))) -> DiffP p i o)
-> (p -> i -> (o, o -> (i, p))) -> DiffP p i o
forall a b. (a -> b) -> a -> b
$ \p
p0 i
i ->
  let dim :: Int
dim = [Double] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Double]
x0
      stepFwd :: [Double] -> [Double]
stepFwd [Double]
xs = let (([Double]
xs', o
_), ([Double], o) -> (([Double], i), p)
_) = p
-> ([Double], i)
-> (([Double], o), ([Double], o) -> (([Double], i), p))
body p
p0 ([Double]
xs, i
i) in [Double]
xs'
      zeroV :: [Double]
zeroV = Int -> Double -> [Double]
forall a. Int -> a -> [a]
replicate Int
dim Double
0.0
      oneV :: Int -> [Double]
oneV Int
k = [if Int
j Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
k then Double
1.0 else Double
0.0 | Int
j <- [Int
0 .. Int
dim Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
      diffV :: [Double] -> [Double] -> [Double]
diffV = (Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith (-)
      normInf :: [a] -> a
normInf [a]
v = [a] -> a
forall a. Ord a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Ord a) => t a -> a
maximum ((a -> a) -> [a] -> [a]
forall a b. (a -> b) -> [a] -> [b]
map a -> a
forall a. Num a => a -> a
abs [a]
v)
      go :: Int -> [Double] -> [Double]
go Int
0 [Double]
xs = [Double]
xs
      go Int
n [Double]
xs =
        let xs' :: [Double]
xs' = [Double] -> [Double]
stepFwd [Double]
xs
         in if [Double] -> Double
forall {a}. (Ord a, Num a) => [a] -> a
normInf ([Double] -> [Double] -> [Double]
diffV [Double]
xs' [Double]
xs) Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
<= Double
tol then [Double]
xs' else Int -> [Double] -> [Double]
go (Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) [Double]
xs'
      a :: [Double]
a = Int -> [Double] -> [Double]
go Int
maxIter [Double]
x0
      aNext :: [Double]
aNext = [Double] -> [Double]
stepFwd [Double]
a
      (([Double]
_, o
o), ([Double], o) -> (([Double], i), p)
backward) = p
-> ([Double], i)
-> (([Double], o), ([Double], o) -> (([Double], i), p))
body p
p0 ([Double]
a, i
i)
      -- Probe the feedback self-coupling @A@: column k is the feedback
      -- cotangent vector produced by a unit cotangent on channel k.
      cols :: [[Double]]
cols = [([Double], i) -> [Double]
forall a b. (a, b) -> a
fst ((([Double], i), p) -> ([Double], i)
forall a b. (a, b) -> a
fst (([Double], o) -> (([Double], i), p)
backward (Int -> [Double]
oneV Int
k, () -> o
forall {k} (arr :: * -> k -> *) (a :: k). Zero arr a => arr () a
zero ()))) | Int
k <- [Int
0 .. Int
dim Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
      -- Assemble rows: row i is [col_0 !! i, ..., col_{dim-1} !! i].
      aMat :: Matrix FieldStar
aMat = [[FieldStar]] -> Matrix FieldStar
forall a. [[a]] -> Matrix a
fromLists [[Double -> FieldStar
FieldStar ([[Double]]
cols [[Double]] -> Int -> [Double]
forall a. HasCallStack => [a] -> Int -> a
!! Int
j [Double] -> Int -> Double
forall a. HasCallStack => [a] -> Int -> a
!! Int
k) | Int
j <- [Int
0 .. Int
dim Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]] | Int
k <- [Int
0 .. Int
dim Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
      aStar :: Matrix Double
aStar = [[Double]] -> Matrix Double
forall a. [[a]] -> Matrix a
fromLists (([FieldStar] -> [Double]) -> [[FieldStar]] -> [[Double]]
forall a b. (a -> b) -> [a] -> [b]
map ((FieldStar -> Double) -> [FieldStar] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map FieldStar -> Double
unFieldStar) (Matrix FieldStar -> [[FieldStar]]
forall a. Matrix a -> [[a]]
toLists (Matrix FieldStar -> Matrix FieldStar
forall a. StarSemiring a => Matrix a -> Matrix a
starMatrix Matrix FieldStar
aMat)))
      pullback :: o -> (i, p)
pullback o
do_ =
        let -- Cross-coupling @B · do_@ for this particular output cotangent.
            cdc :: [Double]
cdc = ([Double], i) -> [Double]
forall a b. (a, b) -> a
fst ((([Double], i), p) -> ([Double], i)
forall a b. (a, b) -> a
fst (([Double], o) -> (([Double], i), p)
backward ([Double]
zeroV, o
do_)))
            dj :: [Double]
dj = Matrix Double -> [Double] -> [Double]
forall a. (Additive a, Multiplicative a) => Matrix a -> [a] -> [a]
matVec Matrix Double
aStar [Double]
cdc
            (([Double]
_, i
di), p
dp) = ([Double], o) -> (([Double], i), p)
backward ([Double]
dj, o
do_)
         in (i
di, p
dp)
   in if [Double] -> Double
forall {a}. (Ord a, Num a) => [a] -> a
normInf ([Double] -> [Double] -> [Double]
diffV [Double]
aNext [Double]
a) Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
tol
        then [Char] -> (o, o -> (i, p))
forall a. HasCallStack => [Char] -> a
error ([Char]
"traceDiffPMatrix: primal fixed point did not converge within " [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ Int -> [Char]
forall a. Show a => a -> [Char]
show Int
maxIter [Char] -> [Char] -> [Char]
forall a. [a] -> [a] -> [a]
++ [Char]
" iterations")
        else (o
o, o -> (i, p)
pullback)