{-# LANGUAGE CPP #-}
{-# LANGUAGE PolyKinds #-}
{-# LANGUAGE TupleSections #-}
{-# LANGUAGE UndecidableInstances #-}
{-# OPTIONS_GHC -Wno-pattern-namespace-specifier #-}

-- | Reverse-mode automatic differentiation as a NumHask carrier.
--
-- A 'Diff s b' is a smooth function @s -> b@ bundled with its pullback.
-- These instances turn it into a NumHask carrier in its own right, so any
-- function written with NumHask-polymorphic operators becomes differentiable
-- by instantiating at @Diff s b@.  The derivative rules live exactly where
-- they should: as the instance methods.
--
-- This is the "functorial lift" direction of the ecosystem membrane:
-- NumHask-polymorphic code wraps AD by using 'Diff as its carrier.
--
-- The lift in one doctest: the identity primitive is "the variable", every
-- operator applied to it is an instance method carrying its own derivative,
-- and the pullback of the composite is the chain rule assembled by
-- instance resolution.  @f s = sin s^2 + s^3@ at @s = 2@: value
-- @sin 4 + 8@, gradient @4*cos 4 + 12@.
--
-- >>> import Circuit.Diff
-- >>> import NumHask.Algebra.Additive qualified as NHA
-- >>> import NumHask.Algebra.Multiplicative qualified as NHM
-- >>> import NumHask.Algebra.Field qualified as NHF
-- >>> let x = Diff (\s -> (s, \db -> db)) :: Diff' Double Double
-- >>> let f = NHF.sin (x NHM.* x) NHA.+ x NHM.* x NHM.* x
-- >>> let (y, pb) = runDiff f 2.0
-- >>> abs (y - (sin 4 + 8)) < 1e-12
-- True
-- >>> abs (pb 1.0 - (4 * cos 4 + 12)) < 1e-12
-- True
module Circuit.Diff
  ( Diff (..),
    Diff',
  )
where

import Control.Category
import NumHask.Algebra.Additive qualified as NHA
import NumHask.Algebra.Field qualified as NHF
import NumHask.Algebra.Multiplicative qualified as NHM
import Prelude hiding (id, (.))
import Prelude qualified as P

-- | A reverse-mode differentiable function tagged by a phantom type @p@.
--
-- The phantom tag prevents perturbation confusion: values of type
-- @Diff p a b@ can only be composed with other @Diff p@ values.  Nested
-- AD introduces a fresh tag for each level.
--
-- @runDiff f a@ returns a pair @(b, pullback)@ where @b = f a@ and @pullback@
-- maps a cotangent @db@ on the output to a cotangent @da@ on the input.
newtype Diff (p :: k) a b = Diff
  { -- | Run the forward pass and return the backward pullback.
    forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff :: a -> (b, b -> a)
  }

-- | The untagged differentiable arrow.  Existing code can continue to use
-- this; it is simply @Diff ()@.
type Diff' = Diff ()

instance Category (Diff p) where
  id :: forall a. Diff p a a
id = (a -> (a, a -> a)) -> Diff p a a
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff (,a -> a
forall a. a -> a
forall {k} (cat :: k -> k -> *) (a :: k). Category cat => cat a a
id)
  Diff b -> (c, c -> b)
f . :: forall b c a. Diff p b c -> Diff p a b -> Diff p a c
. Diff a -> (b, b -> a)
g = (a -> (c, c -> a)) -> Diff p a c
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((a -> (c, c -> a)) -> Diff p a c)
-> (a -> (c, c -> a)) -> Diff p a c
forall a b. (a -> b) -> a -> b
$ \a
a ->
    let (b
b, b -> a
gb) = a -> (b, b -> a)
g a
a
        (c
c, c -> b
fc) = b -> (c, c -> b)
f b
b
     in (c
c, b -> a
gb (b -> a) -> (c -> b) -> c -> a
forall b c a. (b -> c) -> (a -> b) -> a -> c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. c -> b
fc)

-- | Additive structure: sum rule.
--
-- @zero@ is the constant zero function; @(+)@ adds outputs and fans-in
-- cotangents.
instance (NHA.Additive s, NHA.Additive b) => NHA.Additive (Diff p s b) where
  zero :: Diff p s b
zero = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((b, b -> s) -> s -> (b, b -> s)
forall a b. a -> b -> a
const (b
forall a. Additive a => a
NHA.zero, s -> b -> s
forall a b. a -> b -> a
const s
forall a. Additive a => a
NHA.zero))

  Diff s -> (b, b -> s)
f + :: Diff p s b -> Diff p s b -> Diff p s b
+ Diff s -> (b, b -> s)
g = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b1, b -> s
p1) = s -> (b, b -> s)
f s
s
        (b
b2, b -> s
p2) = s -> (b, b -> s)
g s
s
     in (b
b1 b -> b -> b
forall a. Additive a => a -> a -> a
NHA.+ b
b2, \b
db -> b -> s
p1 b
db s -> s -> s
forall a. Additive a => a -> a -> a
NHA.+ b -> s
p2 b
db)

-- | Subtractive structure: negation pushes through the pullback.
instance (NHA.Additive s, NHA.Subtractive s, NHA.Subtractive b) => NHA.Subtractive (Diff p s b) where
  negate :: Diff p s b -> Diff p s b
negate (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
     in (b -> b
forall a. Subtractive a => a -> a
NHA.negate b
b, s -> s
forall a. Subtractive a => a -> a
NHA.negate (s -> s) -> (b -> s) -> b -> s
forall b c a. (b -> c) -> (a -> b) -> a -> c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. b -> s
p)

  Diff s -> (b, b -> s)
f - :: Diff p s b -> Diff p s b -> Diff p s b
- Diff s -> (b, b -> s)
g = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b1, b -> s
p1) = s -> (b, b -> s)
f s
s
        (b
b2, b -> s
p2) = s -> (b, b -> s)
g s
s
     in (b
b1 b -> b -> b
forall a. Subtractive a => a -> a -> a
NHA.- b
b2, \b
db -> b -> s
p1 b
db s -> s -> s
forall a. Subtractive a => a -> a -> a
NHA.- b -> s
p2 b
db)

-- | Multiplicative structure: product rule.
--
-- @one@ is the constant one function.
instance (NHA.Additive s, NHM.Multiplicative b) => NHM.Multiplicative (Diff p s b) where
  one :: Diff p s b
one = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((b, b -> s) -> s -> (b, b -> s)
forall a b. a -> b -> a
const (b
forall a. Multiplicative a => a
NHM.one, s -> b -> s
forall a b. a -> b -> a
const s
forall a. Additive a => a
NHA.zero))

  Diff s -> (b, b -> s)
f * :: Diff p s b -> Diff p s b -> Diff p s b
* Diff s -> (b, b -> s)
g = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b1, b -> s
p1) = s -> (b, b -> s)
f s
s
        (b
b2, b -> s
p2) = s -> (b, b -> s)
g s
s
     in (b
b1 b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
b2, \b
db -> b -> s
p1 (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
b2) s -> s -> s
forall a. Additive a => a -> a -> a
NHA.+ b -> s
p2 (b
b1 b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
db))

-- | Divisive structure: reciprocal rule.
--
-- Division inherits the product rule via the default @'/' = '*' . 'recip'@.
instance
  (NHA.Additive s, NHA.Subtractive b, NHM.Multiplicative b, NHM.Divisive b) =>
  NHM.Divisive (Diff p s b)
  where
  recip :: Diff p s b -> Diff p s b
recip (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
        r :: b
r = b -> b
forall a. Divisive a => a -> a
NHM.recip b
b
        rr :: b
rr = b
r b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
r
     in (b
r, \b
db -> b -> s
p (b -> b
forall a. Subtractive a => a -> a
NHA.negate (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
rr)))

-- | Exponential field: @exp@, @log@, and the derived power/root family.
instance
  (NHF.ExpField b, NHA.Additive s, NHA.Subtractive s, NHM.Multiplicative b) =>
  NHF.ExpField (Diff p s b)
  where
  exp :: Diff p s b -> Diff p s b
exp (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
        e :: b
e = b -> b
forall a. ExpField a => a -> a
NHF.exp b
b
     in (b
e, \b
db -> b -> s
p (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
e))

  log :: Diff p s b -> Diff p s b
log (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
     in (b -> b
forall a. ExpField a => a -> a
NHF.log b
b, \b
db -> b -> s
p (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b -> b
forall a. Divisive a => a -> a
NHM.recip b
b))

-- | Trigonometric field: the elementary transcendental family.
instance
  ( NHF.TrigField b,
    NHF.ExpField b,
    NHA.Additive s,
    NHA.Subtractive s,
    NHM.Multiplicative b,
    NHM.Divisive b
  ) =>
  NHF.TrigField (Diff p s b)
  where
  pi :: Diff p s b
pi = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((b, b -> s) -> s -> (b, b -> s)
forall a b. a -> b -> a
const (b
forall a. TrigField a => a
NHF.pi, s -> b -> s
forall a b. a -> b -> a
const s
forall a. Additive a => a
NHA.zero))

  sin :: Diff p s b -> Diff p s b
sin (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
     in (b -> b
forall a. TrigField a => a -> a
NHF.sin b
b, \b
db -> b -> s
p (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b -> b
forall a. TrigField a => a -> a
NHF.cos b
b))

  cos :: Diff p s b -> Diff p s b
cos (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
     in (b -> b
forall a. TrigField a => a -> a
NHF.cos b
b, \b
db -> b -> s
p (b -> b
forall a. Subtractive a => a -> a
NHA.negate (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b -> b
forall a. TrigField a => a -> a
NHF.sin b
b)))

  asin :: Diff p s b -> Diff p s b
asin (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
        d :: b
d = b -> b
forall a. Divisive a => a -> a
NHM.recip (b -> b
forall a. ExpField a => a -> a
NHF.sqrt (b
forall a. Multiplicative a => a
NHM.one b -> b -> b
forall a. Subtractive a => a -> a -> a
NHA.- b
b b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
b))
     in (b -> b
forall a. TrigField a => a -> a
NHF.asin b
b, \b
db -> b -> s
p (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
d))

  acos :: Diff p s b -> Diff p s b
acos (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
        d :: b
d = b -> b
forall a. Divisive a => a -> a
NHM.recip (b -> b
forall a. ExpField a => a -> a
NHF.sqrt (b
forall a. Multiplicative a => a
NHM.one b -> b -> b
forall a. Subtractive a => a -> a -> a
NHA.- b
b b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
b))
     in (b -> b
forall a. TrigField a => a -> a
NHF.acos b
b, \b
db -> b -> s
p (b -> b
forall a. Subtractive a => a -> a
NHA.negate (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
d)))

  atan :: Diff p s b -> Diff p s b
atan (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
        d :: b
d = b -> b
forall a. Divisive a => a -> a
NHM.recip (b
forall a. Multiplicative a => a
NHM.one b -> b -> b
forall a. Additive a => a -> a -> a
NHA.+ b
b b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
b)
     in (b -> b
forall a. TrigField a => a -> a
NHF.atan b
b, \b
db -> b -> s
p (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
d))

  atan2 :: Diff p s b -> Diff p s b -> Diff p s b
atan2 (Diff s -> (b, b -> s)
f) (Diff s -> (b, b -> s)
g) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
y, b -> s
py) = s -> (b, b -> s)
f s
s
        (b
x, b -> s
px) = s -> (b, b -> s)
g s
s
        r :: b
r = b -> b -> b
forall a. TrigField a => a -> a -> a
NHF.atan2 b
y b
x
        denom :: b
denom = b
y b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
y b -> b -> b
forall a. Additive a => a -> a -> a
NHA.+ b
x b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
x
        dy :: b
dy = b
x b -> b -> b
forall a. Divisive a => a -> a -> a
NHM./ b
denom
        dx :: b
dx = b -> b
forall a. Subtractive a => a -> a
NHA.negate (b
y b -> b -> b
forall a. Divisive a => a -> a -> a
NHM./ b
denom)
     in (b
r, \b
db -> b -> s
py (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
dy) s -> s -> s
forall a. Additive a => a -> a -> a
NHA.+ b -> s
px (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
dx))

  sinh :: Diff p s b -> Diff p s b
sinh (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
     in (b -> b
forall a. TrigField a => a -> a
NHF.sinh b
b, \b
db -> b -> s
p (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b -> b
forall a. TrigField a => a -> a
NHF.cosh b
b))

  cosh :: Diff p s b -> Diff p s b
cosh (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
     in (b -> b
forall a. TrigField a => a -> a
NHF.cosh b
b, \b
db -> b -> s
p (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b -> b
forall a. TrigField a => a -> a
NHF.sinh b
b))

  asinh :: Diff p s b -> Diff p s b
asinh (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
        d :: b
d = b -> b
forall a. Divisive a => a -> a
NHM.recip (b -> b
forall a. ExpField a => a -> a
NHF.sqrt (b
b b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
b b -> b -> b
forall a. Additive a => a -> a -> a
NHA.+ b
forall a. Multiplicative a => a
NHM.one))
     in (b -> b
forall a. TrigField a => a -> a
NHF.asinh b
b, \b
db -> b -> s
p (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
d))

  acosh :: Diff p s b -> Diff p s b
acosh (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
        d :: b
d = b -> b
forall a. Divisive a => a -> a
NHM.recip (b -> b
forall a. ExpField a => a -> a
NHF.sqrt (b
b b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
b b -> b -> b
forall a. Subtractive a => a -> a -> a
NHA.- b
forall a. Multiplicative a => a
NHM.one))
     in (b -> b
forall a. TrigField a => a -> a
NHF.acosh b
b, \b
db -> b -> s
p (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
d))

  atanh :: Diff p s b -> Diff p s b
atanh (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
        d :: b
d = b -> b
forall a. Divisive a => a -> a
NHM.recip (b
forall a. Multiplicative a => a
NHM.one b -> b -> b
forall a. Subtractive a => a -> a -> a
NHA.- b
b b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
b)
     in (b -> b
forall a. TrigField a => a -> a
NHF.atanh b
b, \b
db -> b -> s
p (b
db b -> b -> b
forall a. Multiplicative a => a -> a -> a
NHM.* b
d))

-- | Mechanical 'Num' mirror so that base-polymorphic code can also use 'Diff
-- as its carrier.  Non-smooth methods ('abs', 'signum') raise an error; the
-- useful cases (literals, '+', '*', '-') work out of the box.
instance (P.Num s, P.Num b) => P.Num (Diff p s b) where
  Diff s -> (b, b -> s)
f + :: Diff p s b -> Diff p s b -> Diff p s b
+ Diff s -> (b, b -> s)
g = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b1, b -> s
p1) = s -> (b, b -> s)
f s
s
        (b
b2, b -> s
p2) = s -> (b, b -> s)
g s
s
     in (b
b1 b -> b -> b
forall a. Num a => a -> a -> a
P.+ b
b2, \b
db -> b -> s
p1 b
db s -> s -> s
forall a. Num a => a -> a -> a
P.+ b -> s
p2 b
db)

  Diff s -> (b, b -> s)
f * :: Diff p s b -> Diff p s b -> Diff p s b
* Diff s -> (b, b -> s)
g = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b1, b -> s
p1) = s -> (b, b -> s)
f s
s
        (b
b2, b -> s
p2) = s -> (b, b -> s)
g s
s
     in (b
b1 b -> b -> b
forall a. Num a => a -> a -> a
P.* b
b2, \b
db -> b -> s
p1 (b
db b -> b -> b
forall a. Num a => a -> a -> a
P.* b
b2) s -> s -> s
forall a. Num a => a -> a -> a
P.+ b -> s
p2 (b
b1 b -> b -> b
forall a. Num a => a -> a -> a
P.* b
db))

  negate :: Diff p s b -> Diff p s b
negate (Diff s -> (b, b -> s)
f) = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b, b -> s
p) = s -> (b, b -> s)
f s
s
     in (b -> b
forall a. Num a => a -> a
P.negate b
b, s -> s
forall a. Num a => a -> a
P.negate (s -> s) -> (b -> s) -> b -> s
forall b c a. (b -> c) -> (a -> b) -> a -> c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. b -> s
p)

  Diff s -> (b, b -> s)
f - :: Diff p s b -> Diff p s b -> Diff p s b
- Diff s -> (b, b -> s)
g = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((s -> (b, b -> s)) -> Diff p s b)
-> (s -> (b, b -> s)) -> Diff p s b
forall a b. (a -> b) -> a -> b
$ \s
s ->
    let (b
b1, b -> s
p1) = s -> (b, b -> s)
f s
s
        (b
b2, b -> s
p2) = s -> (b, b -> s)
g s
s
     in (b
b1 b -> b -> b
forall a. Num a => a -> a -> a
P.- b
b2, \b
db -> b -> s
p1 b
db s -> s -> s
forall a. Num a => a -> a -> a
P.- b -> s
p2 b
db)

  abs :: Diff p s b -> Diff p s b
abs Diff p s b
_ = [Char] -> Diff p s b
forall a. HasCallStack => [Char] -> a
P.error [Char]
"Circuit.Diff: abs is not differentiable at 0"
  signum :: Diff p s b -> Diff p s b
signum Diff p s b
_ = [Char] -> Diff p s b
forall a. HasCallStack => [Char] -> a
P.error [Char]
"Circuit.Diff: signum is not differentiable at 0"

  fromInteger :: Integer -> Diff p s b
fromInteger Integer
n = (s -> (b, b -> s)) -> Diff p s b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((b, b -> s) -> s -> (b, b -> s)
forall a b. a -> b -> a
const (Integer -> b
forall a. Num a => Integer -> a
P.fromInteger Integer
n, s -> b -> s
forall a b. a -> b -> a
const s
0))