{-# LANGUAGE ScopedTypeVariables #-}

-- | Eshkol-style AD operators on top of the 'Circuit.Diff.Diff' carrier.
--
-- These are convenience wrappers around 'runDiff.  They expose a JAX-like
-- surface --- derivative, gradient, jacobian, hessian, divergence, curl,
-- laplacian --- while keeping the substrate's exact reverse-mode engine
-- underneath for first-order operators.
--
-- Higher-order scalar towers now use the 'Circuit.Diff.Taylor' carrier: a
-- 'Diff' function is sampled near the expansion point and a truncated Taylor
-- morphism is built from the finite-difference table.  This is an approximation
-- (the exact tower would require a dedicated forward-mode or Taylor-mode
-- carrier from the start), but it is enough to make 'derivativeN', 'taylor',
-- 'hessian' and 'laplacian' usable.
module Circuit.Diff.Operators
  ( -- * First-order operators
    derivative,
    gradient,
    jacobian,

    -- * Second-order operators
    hessian,
    divergence,
    curl,
    laplacian,

    -- * Higher-order towers (finite-difference bridge from Diff)
    derivativeN,
    taylor,

    -- * Exact higher-order towers (compositional Jet)
    derivativeNJ,
    taylorJ,
  )
where

import Circuit.Diff (Diff (..), runDiff)
import Circuit.Diff.Jet (Jet)
import Circuit.Diff.Jet qualified as Jet
import Circuit.Diff.Taylor (Taylor, approxTaylorFromDiff, evalTaylor)
import Data.Proxy (Proxy)
import GHC.TypeNats (SomeNat (..), someNatVal)
import Numeric.Natural (Natural)
import Prelude

-- $setup
-- >>> import Circuit.Diff (Diff (..))

-- | Scalar derivative: @f : R -> R@ at @x@.
--
-- >>> let sq = Diff (\x -> (x * x, \d -> 2 * x * d)) :: Diff () Double Double
-- >>> derivative sq 3.0
-- 6.0
derivative :: Diff p Double Double -> Double -> Double
derivative :: forall {k} (p :: k). Diff p Double Double -> Double -> Double
derivative Diff p Double Double
f Double
x =
  let (Double
_, Double -> Double
pb) = Diff p Double Double -> Double -> (Double, Double -> Double)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff p Double Double
f Double
x
   in Double -> Double
pb Double
1.0

-- | Gradient of a scalar function: @f : R^n -> R@ at @v@.
--
-- Returns the vector @∇f(v)@ by applying the pullback to the unit scalar.
gradient :: Diff p [Double] Double -> [Double] -> [Double]
gradient :: forall {k} (p :: k). Diff p [Double] Double -> [Double] -> [Double]
gradient Diff p [Double] Double
f [Double]
v =
  let (Double
_, Double -> [Double]
pb) = Diff p [Double] Double -> [Double] -> (Double, Double -> [Double])
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff p [Double] Double
f [Double]
v
   in Double -> [Double]
pb Double
1.0

-- | Jacobian of a vector function: @f : R^n -> R^m@ at @v@.
--
-- Returns an @m × n@ matrix: outer index is output component, inner index is
-- input component.  Each row is obtained by applying the pullback to one
-- output basis vector.
jacobian :: Diff p [Double] [Double] -> [Double] -> [[Double]]
jacobian :: forall {k} (p :: k).
Diff p [Double] [Double] -> [Double] -> [[Double]]
jacobian Diff p [Double] [Double]
f [Double]
v =
  let ([Double]
y, [Double] -> [Double]
pb) = Diff p [Double] [Double]
-> [Double] -> ([Double], [Double] -> [Double])
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff p [Double] [Double]
f [Double]
v
      m :: Int
m = [Double] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Double]
y
      basis :: Int -> [Double]
basis Int
i = [if Int
j Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
i then Double
1.0 else Double
0.0 | Int
j <- [Int
0 .. Int
m Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
   in [[Double] -> [Double]
pb (Int -> [Double]
basis Int
i) | Int
i <- [Int
0 .. Int
m Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]

-- | Trace of the Jacobian: @div f v = Σ_i ∂f_i/∂x_i@.
divergence :: Diff p [Double] [Double] -> [Double] -> Double
divergence :: forall {k} (p :: k). Diff p [Double] [Double] -> [Double] -> Double
divergence Diff p [Double] [Double]
f [Double]
v =
  let j :: [[Double]]
j = Diff p [Double] [Double] -> [Double] -> [[Double]]
forall {k} (p :: k).
Diff p [Double] [Double] -> [Double] -> [[Double]]
jacobian Diff p [Double] [Double]
f [Double]
v
      n :: Int
n = [Double] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Double]
v
   in [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [[[Double]]
j [[Double]] -> Int -> [Double]
forall a. HasCallStack => [a] -> Int -> a
!! Int
i [Double] -> Int -> Double
forall a. HasCallStack => [a] -> Int -> a
!! Int
i | Int
i <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]

-- | Curl of a 3-D vector field: @f : R^3 -> R^3@ at @v@.
curl :: Diff p [Double] [Double] -> [Double] -> [Double]
curl :: forall {k} (p :: k).
Diff p [Double] [Double] -> [Double] -> [Double]
curl Diff p [Double] [Double]
f [Double]
v =
  let j :: [[Double]]
j = Diff p [Double] [Double] -> [Double] -> [[Double]]
forall {k} (p :: k).
Diff p [Double] [Double] -> [Double] -> [[Double]]
jacobian Diff p [Double] [Double]
f [Double]
v
      at :: Int -> Int -> Double
at Int
r Int
c = ([[Double]]
j [[Double]] -> Int -> [Double]
forall a. HasCallStack => [a] -> Int -> a
!! Int
r) [Double] -> Int -> Double
forall a. HasCallStack => [a] -> Int -> a
!! Int
c
   in [ Int -> Int -> Double
at Int
2 Int
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Int -> Int -> Double
at Int
1 Int
2,
        Int -> Int -> Double
at Int
0 Int
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Int -> Int -> Double
at Int
2 Int
0,
        Int -> Int -> Double
at Int
1 Int
0 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Int -> Int -> Double
at Int
0 Int
1
      ]

-- | Hessian of a scalar function: @f : R^n -> R@ at @v@.
--
-- Implemented by central second-order finite differences; approximate.
hessian :: Diff p [Double] Double -> [Double] -> [[Double]]
hessian :: forall {k} (p :: k).
Diff p [Double] Double -> [Double] -> [[Double]]
hessian Diff p [Double] Double
f [Double]
v =
  let n :: Int
n = [Double] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Double]
v
      h :: Double
h = Double
1e-4 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
1.0 (Double -> Double
forall a. Floating a => a -> a
sqrt ([Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum ((Double -> Double) -> [Double] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map (Double -> Int -> Double
forall a b. (Num a, Integral b) => a -> b -> a
^ (Int
2 :: Int)) [Double]
v)))
      e :: Int -> [Double]
e Int
i = [if Int
j Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
i then Double
h else Double
0.0 | Int
j <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]
      val :: [Double] -> Double
val [Double]
u = (Double, Double -> [Double]) -> Double
forall a b. (a, b) -> a
fst (Diff p [Double] Double -> [Double] -> (Double, Double -> [Double])
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff p [Double] Double
f [Double]
u)
      fij :: Int -> Int -> Double
fij Int
i Int
j =
        ( [Double] -> Double
val ((Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) [Double]
v ((Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) (Int -> [Double]
e Int
i) (Int -> [Double]
e Int
j)))
            Double -> Double -> Double
forall a. Num a => a -> a -> a
- [Double] -> Double
val ((Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) [Double]
v (Int -> [Double]
e Int
i))
            Double -> Double -> Double
forall a. Num a => a -> a -> a
- [Double] -> Double
val ((Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) [Double]
v (Int -> [Double]
e Int
j))
            Double -> Double -> Double
forall a. Num a => a -> a -> a
+ [Double] -> Double
val [Double]
v
        )
          Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
h Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
h)
   in [[Int -> Int -> Double
fij Int
i Int
j | Int
j <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]] | Int
i <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]

-- | Laplacian of a scalar function: @f : R^n -> R@ at @v@.
--
-- Trace of the finite-difference Hessian.
laplacian :: Diff p [Double] Double -> [Double] -> Double
laplacian :: forall {k} (p :: k). Diff p [Double] Double -> [Double] -> Double
laplacian Diff p [Double] Double
f [Double]
v =
  let h :: [[Double]]
h = Diff p [Double] Double -> [Double] -> [[Double]]
forall {k} (p :: k).
Diff p [Double] Double -> [Double] -> [[Double]]
hessian Diff p [Double] Double
f [Double]
v
      n :: Int
n = [Double] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [Double]
v
   in [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [[[Double]]
h [[Double]] -> Int -> [Double]
forall a. HasCallStack => [a] -> Int -> a
!! Int
i [Double] -> Int -> Double
forall a. HasCallStack => [a] -> Int -> a
!! Int
i | Int
i <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1]]

-- | N-th derivative of a scalar function: @f : R -> R@ at @x@.
--
-- For @n <= 1@ this is exact reverse-mode AD.  Higher orders are read from
-- a finite-difference Taylor tower built with 'Circuit.Diff.Taylor'.
derivativeN :: Diff p Double Double -> Double -> Int -> Double
derivativeN :: forall {k} (p :: k).
Diff p Double Double -> Double -> Int -> Double
derivativeN Diff p Double Double
f Double
x Int
0 = (Double, Double -> Double) -> Double
forall a b. (a, b) -> a
fst (Diff p Double Double -> Double -> (Double, Double -> Double)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff p Double Double
f Double
x)
derivativeN Diff p Double Double
f Double
x Int
n
  | Int
n Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0 = [Char] -> Double
forall a. HasCallStack => [Char] -> a
error [Char]
"derivativeN: negative order"
  | Bool
otherwise = Diff p Double Double -> Double -> Int -> [Double]
forall {k} (p :: k).
Diff p Double Double -> Double -> Int -> [Double]
taylor Diff p Double Double
f Double
x Int
n [Double] -> Int -> Double
forall a. HasCallStack => [a] -> Int -> a
!! Int
n

-- | First @k@ raw derivatives of @f : R -> R@ at @x0@.
--
-- The result is @[f(x0), f'(x0), f''(x0), ..., f^(k)(x0)]@.  The constant
-- and linear terms are exact; higher terms come from a finite-difference
-- Taylor tower built with 'Circuit.Diff.Taylor'.
taylor :: Diff p Double Double -> Double -> Int -> [Double]
taylor :: forall {k} (p :: k).
Diff p Double Double -> Double -> Int -> [Double]
taylor Diff p Double Double
_ Double
_ Int
k | Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0 = [Char] -> [Double]
forall a. HasCallStack => [Char] -> a
error [Char]
"taylor: negative order"
taylor Diff p Double Double
f Double
x0 Int
k =
  case Natural -> SomeNat
someNatVal (Int -> Natural
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
k :: Natural) of
    SomeNat (Proxy n
_ :: Proxy n) ->
      let cs :: [Double]
cs = Taylor n Double Double -> Double -> [Double]
forall (n :: Natural).
KnownNat n =>
Taylor n Double Double -> Double -> [Double]
evalTaylor (Diff p Double Double -> Double -> Taylor n Double Double
forall {k} (n :: Natural) (p :: k).
KnownNat n =>
Diff p Double Double -> Double -> Taylor n Double Double
approxTaylorFromDiff Diff p Double Double
f Double
x0 :: Taylor n Double Double) Double
0
          facts :: [Double]
facts = (Double -> Double -> Double) -> Double -> [Double] -> [Double]
forall b a. (b -> a -> b) -> b -> [a] -> [b]
scanl Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Double
1.0 [Double
1.0 .. Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
k]
       in (Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) [Double]
cs [Double]
facts

-- | Exact N-th derivative of a scalar function expressed as a compositional
-- 'Circuit.Diff.Jet' tower.
--
-- The function must be built from NumHask-polymorphic operations; the tower is
-- propagated by the closed recurrences in 'Circuit.Diff.Jet'.  There is no
-- finite-difference approximation.
derivativeNJ :: (Jet Double -> Jet Double) -> Double -> Int -> Double
derivativeNJ :: (Jet Double -> Jet Double) -> Double -> Int -> Double
derivativeNJ Jet Double -> Jet Double
f Double
x Int
n
  | Int
n Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0 = [Char] -> Double
forall a. HasCallStack => [Char] -> a
error [Char]
"derivativeNJ: negative order"
  | Bool
otherwise = (Jet Double -> Jet Double) -> Double -> Int -> [Double]
taylorJ Jet Double -> Jet Double
f Double
x Int
n [Double] -> Int -> Double
forall a. HasCallStack => [a] -> Int -> a
!! Int
n

-- | Exact first @k@ raw derivatives via a compositional 'Circuit.Diff.Jet'
-- tower.
taylorJ :: (Jet Double -> Jet Double) -> Double -> Int -> [Double]
taylorJ :: (Jet Double -> Jet Double) -> Double -> Int -> [Double]
taylorJ Jet Double -> Jet Double
f Double
x Int
k
  | Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Int
0 = [Char] -> [Double]
forall a. HasCallStack => [Char] -> a
error [Char]
"taylorJ: negative order"
  | Bool
otherwise = (Jet Double -> Jet Double) -> Int -> Double -> [Double]
forall a.
(ExpField a, FromInteger a) =>
(Jet a -> Jet a) -> Int -> a -> [a]
Jet.taylor Jet Double -> Jet Double
f Int
k Double
x