{-# LANGUAGE RebindableSyntax #-}

-- | Jets via truncated Taylor series.
--
-- A 'Jet' is a finite tower of Taylor coefficients
--
-- > c0 + c1*h + c2*h^2 + ... + cn*h^n
--
-- around a primal point.  Elementary functions act coefficient-wise via the
-- usual dual-number recurrences, so a NumHask-polymorphic function
-- @f :: (ExpField a, TrigField a) => a -> a@ applied to @'variable' n a@
-- returns the first @n+1@ Taylor coefficients of @f@ at @a@.
--
-- This is the "iterated" direction of @Diff@: where @Diff@ carries one
-- pullback, a jet carries the whole truncated tower.  The two interoperate
-- through 'jetFromDiff, which seeds the tower from a first-order pullback.
module Circuit.Diff.Jet
  ( -- * Jet type
    Jet (..),
    jetOrder,

    -- * Construction
    variable,
    constant,
    fromDiff,

    -- * Coefficient views
    taylorDers,
    taylor,

    -- * Series operations
    differentiate,
    integrate,
    scale,
    resize,
  )
where

import Circuit.Diff (Diff, runDiff)
import Circuit.Process qualified as CP (Process (..), scan)
import NumHask.Algebra.Additive (Additive (..), Subtractive (..), sum)
import NumHask.Algebra.Field (ExpField (..), TrigField (..))
import NumHask.Algebra.Multiplicative (Divisive (..), Multiplicative (..))
import NumHask.Data.Integral (FromInteger (..))
import NumHask.Prelude

-- | Truncated Taylor series stored as coefficients @[c0, c1, ..., cn]@
-- representing @c0 + c1*h + c2*h^2 + ... + cn*h^n@.
newtype Jet a = Jet {forall a. Jet a -> [a]
coefficients :: [a]}
  deriving (Jet a -> Jet a -> Bool
(Jet a -> Jet a -> Bool) -> (Jet a -> Jet a -> Bool) -> Eq (Jet a)
forall a. Eq a => Jet a -> Jet a -> Bool
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: forall a. Eq a => Jet a -> Jet a -> Bool
== :: Jet a -> Jet a -> Bool
$c/= :: forall a. Eq a => Jet a -> Jet a -> Bool
/= :: Jet a -> Jet a -> Bool
Eq, Int -> Jet a -> ShowS
[Jet a] -> ShowS
Jet a -> String
(Int -> Jet a -> ShowS)
-> (Jet a -> String) -> ([Jet a] -> ShowS) -> Show (Jet a)
forall a. Show a => Int -> Jet a -> ShowS
forall a. Show a => [Jet a] -> ShowS
forall a. Show a => Jet a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> Jet a -> ShowS
showsPrec :: Int -> Jet a -> ShowS
$cshow :: forall a. Show a => Jet a -> String
show :: Jet a -> String
$cshowList :: forall a. Show a => [Jet a] -> ShowS
showList :: [Jet a] -> ShowS
Show)

-- | Highest power of @h@ present.
jetOrder :: Jet a -> Int
jetOrder :: forall a. Jet a -> Int
jetOrder = Int -> Int
forall a. Enum a => a -> a
pred (Int -> Int) -> (Jet a -> Int) -> Jet a -> Int
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
. [a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length ([a] -> Int) -> (Jet a -> [a]) -> Jet a -> Int
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
. Jet a -> [a]
forall a. Jet a -> [a]
coefficients

-- | Truncate / pad to the given order.
resize :: (Additive a) => Int -> Jet a -> Jet a
resize :: forall a. Additive a => Int -> Jet a -> Jet a
resize Int
n (Jet [a]
cs) = [a] -> Jet a
forall a. [a] -> Jet a
Jet ([a] -> Jet a) -> [a] -> Jet a
forall a b. (a -> b) -> a -> b
$ Int -> [a] -> [a]
forall a. Int -> [a] -> [a]
take (Int
n Int -> Int -> Int
forall a. Additive a => a -> a -> a
+ Int
1) ([a]
cs [a] -> [a] -> [a]
forall a. [a] -> [a] -> [a]
++ a -> [a]
forall a. a -> [a]
repeat a
forall a. Additive a => a
zero)

-- | Align two jets to the same order by truncating the higher one.
align :: Jet a -> Jet a -> ([a], [a])
align :: forall a. Jet a -> Jet a -> ([a], [a])
align (Jet [a]
xs) (Jet [a]
ys) =
  let n :: Int
n = Int -> Int -> Int
forall a. Ord a => a -> a -> a
min ([a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
xs) ([a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
ys)
   in (Int -> [a] -> [a]
forall a. Int -> [a] -> [a]
take Int
n [a]
xs, Int -> [a] -> [a]
forall a. Int -> [a] -> [a]
take Int
n [a]
ys)

-- | Align two jets, lifting a length-1 (constant) jet to the order of the
-- other by padding with zeros.  This makes @one@, @zero@ and numeric literals
-- behave as scalars of arbitrary order.
alignLift :: (Additive a) => Jet a -> Jet a -> ([a], [a])
alignLift :: forall a. Additive a => Jet a -> Jet a -> ([a], [a])
alignLift (Jet [a
x]) (Jet [a]
ys) = (a
x a -> [a] -> [a]
forall a. a -> [a] -> [a]
: Int -> a -> [a]
forall a. Int -> a -> [a]
replicate ([a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
ys Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1) a
forall a. Additive a => a
zero, [a]
ys)
alignLift (Jet [a]
xs) (Jet [a
y]) = ([a]
xs, a
y a -> [a] -> [a]
forall a. a -> [a] -> [a]
: Int -> a -> [a]
forall a. Int -> a -> [a]
replicate ([a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
xs Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1) a
forall a. Additive a => a
zero)
alignLift (Jet [a]
xs) (Jet [a]
ys) = Jet a -> Jet a -> ([a], [a])
forall a. Jet a -> Jet a -> ([a], [a])
align ([a] -> Jet a
forall a. [a] -> Jet a
Jet [a]
xs) ([a] -> Jet a
forall a. [a] -> Jet a
Jet [a]
ys)

-- | Build a jet of order @n@ representing the input variable @a + h@.
variable :: (Additive a, Multiplicative a) => Int -> a -> Jet a
variable :: forall a. (Additive a, Multiplicative a) => Int -> a -> Jet a
variable Int
n a
a = [a] -> Jet a
forall a. [a] -> Jet a
Jet (a
a a -> [a] -> [a]
forall a. a -> [a] -> [a]
: a
forall a. Multiplicative a => a
one a -> [a] -> [a]
forall a. a -> [a] -> [a]
: Int -> a -> [a]
forall a. Int -> a -> [a]
replicate (Int
n Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1) a
forall a. Additive a => a
zero)

-- | Build a constant jet of order @n@.
constant :: (Additive a) => Int -> a -> Jet a
constant :: forall a. Additive a => Int -> a -> Jet a
constant Int
n a
c = [a] -> Jet a
forall a. [a] -> Jet a
Jet (a
c a -> [a] -> [a]
forall a. a -> [a] -> [a]
: Int -> a -> [a]
forall a. Int -> a -> [a]
replicate Int
n a
forall a. Additive a => a
zero)

-- | Seed a first-order jet from a 'Diff' first derivative.
--
-- Higher derivatives are /not/ recovered from a bare 'Diff'; use 'taylor'
-- with a NumHask-polymorphic function for automatic higher-order towers.
fromDiff :: (Multiplicative a) => Diff p a a -> a -> Jet a
fromDiff :: forall {k} a (p :: k). Multiplicative a => Diff p a a -> a -> Jet a
fromDiff Diff p a a
f a
a =
  let (a
y, a -> a
pb) = Diff p a a -> a -> (a, a -> a)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff p a a
f a
a
   in [a] -> Jet a
forall a. [a] -> Jet a
Jet [a
y, a -> a
pb a
forall a. Multiplicative a => a
one]

-- | Convert Taylor coefficients to raw derivatives.
--
-- > taylorDers (Jet [c0, c1, c2]) = [c0, 1!*c1, 2!*c2]
taylorDers :: (Additive a, Multiplicative a, FromInteger a) => Jet a -> [a]
taylorDers :: forall a.
(Additive a, Multiplicative a, FromInteger a) =>
Jet a -> [a]
taylorDers (Jet [a]
cs) = (a -> a -> a) -> [a] -> [a] -> [a]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith a -> a -> a
forall a. Multiplicative a => a -> a -> a
(*) [a]
cs [a]
factorials
  where
    factorials :: [a]
factorials = (a -> a -> a) -> a -> [a] -> [a]
forall b a. (b -> a -> b) -> b -> [a] -> [b]
scanl a -> a -> a
forall a. Multiplicative a => a -> a -> a
(*) a
forall a. Multiplicative a => a
one ((Integer -> a) -> [Integer] -> [a]
forall a b. (a -> b) -> [a] -> [b]
map ((a
forall a. Multiplicative a => a
one a -> a -> a
forall a. Additive a => a -> a -> a
+) (a -> a) -> (Integer -> a) -> Integer -> 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
. Integer -> a
forall a. FromInteger a => Integer -> a
fromInteger) [(Integer
0 :: Integer) ..])

-- | Apply a jet-level function at a point and return the raw
-- derivatives @[f(a), f'(a), f''(a), ..., f^(n)(a)]@.
taylor ::
  (ExpField a, FromInteger a) =>
  (Jet a -> Jet a) ->
  Int ->
  a ->
  [a]
taylor :: forall a.
(ExpField a, FromInteger a) =>
(Jet a -> Jet a) -> Int -> a -> [a]
taylor Jet a -> Jet a
f Int
n a
a = Jet a -> [a]
forall a.
(Additive a, Multiplicative a, FromInteger a) =>
Jet a -> [a]
taylorDers (Jet a -> Jet a
f (Int -> a -> Jet a
forall a. (Additive a, Multiplicative a) => Int -> a -> Jet a
variable Int
n a
a))

-- ---------------------------------------------------------------------------
-- NumHask instances
-- ---------------------------------------------------------------------------

instance (Additive a) => Additive (Jet a) where
  zero :: Jet a
zero = [a] -> Jet a
forall a. [a] -> Jet a
Jet [a
forall a. Additive a => a
zero]
  Jet [a]
xs + :: Jet a -> Jet a -> Jet a
+ Jet [a]
ys =
    let ([a]
xs', [a]
ys') = Jet a -> Jet a -> ([a], [a])
forall a. Additive a => Jet a -> Jet a -> ([a], [a])
alignLift ([a] -> Jet a
forall a. [a] -> Jet a
Jet [a]
xs) ([a] -> Jet a
forall a. [a] -> Jet a
Jet [a]
ys)
     in [a] -> Jet a
forall a. [a] -> Jet a
Jet ((a -> a -> a) -> [a] -> [a] -> [a]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith a -> a -> a
forall a. Additive a => a -> a -> a
(+) [a]
xs' [a]
ys')

instance (Subtractive a) => Subtractive (Jet a) where
  negate :: Jet a -> Jet a
negate (Jet [a]
xs) = [a] -> Jet a
forall a. [a] -> Jet a
Jet ((a -> a) -> [a] -> [a]
forall a b. (a -> b) -> [a] -> [b]
map a -> a
forall a. Subtractive a => a -> a
negate [a]
xs)
  Jet [a]
xs - :: Jet a -> Jet a -> Jet a
- Jet [a]
ys =
    let ([a]
xs', [a]
ys') = Jet a -> Jet a -> ([a], [a])
forall a. Additive a => Jet a -> Jet a -> ([a], [a])
alignLift ([a] -> Jet a
forall a. [a] -> Jet a
Jet [a]
xs) ([a] -> Jet a
forall a. [a] -> Jet a
Jet [a]
ys)
     in [a] -> Jet a
forall a. [a] -> Jet a
Jet ((a -> a -> a) -> [a] -> [a] -> [a]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith (-) [a]
xs' [a]
ys')

instance (Additive a, Multiplicative a) => Multiplicative (Jet a) where
  one :: Jet a
one = [a] -> Jet a
forall a. [a] -> Jet a
Jet [a
forall a. Multiplicative a => a
one]
  Jet [a
c] * :: Jet a -> Jet a -> Jet a
* Jet [a]
ys = [a] -> Jet a
forall a. [a] -> Jet a
Jet ((a -> a) -> [a] -> [a]
forall a b. (a -> b) -> [a] -> [b]
map (a
c a -> a -> a
forall a. Multiplicative a => a -> a -> a
*) [a]
ys)
  Jet [a]
xs * Jet [a
c] = [a] -> Jet a
forall a. [a] -> Jet a
Jet ((a -> a) -> [a] -> [a]
forall a b. (a -> b) -> [a] -> [b]
map (a -> a -> a
forall a. Multiplicative a => a -> a -> a
* a
c) [a]
xs)
  Jet [a]
xs * Jet [a]
ys =
    let n :: Int
n = Int -> Int -> Int
forall a. Ord a => a -> a -> a
min ([a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
xs) ([a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
ys)
        cauchy :: Int -> a
cauchy Int
k = [a] -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum [[a]
xs [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! Int
i a -> a -> a
forall a. Multiplicative a => a -> a -> a
* [a]
ys [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! (Int
k Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
i) | Int
i <- [Int
0 .. Int
k]]
     in [a] -> Jet a
forall a. [a] -> Jet a
Jet [Int -> a
cauchy Int
k | Int
k <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]

-- | Tail process for the reciprocal series.
--
-- If @u = u0 + u1*h + ...@ and @v = 1/u = v0 + v1*h + ...@ then
-- @v0 = 1/u0@ and @vk = -(sum_{i=1}^k ui * v{k-i}) / u0@.
-- The process consumes @u1, u2, ...@ and emits @v1, v2, ...@; the caller
-- prepends @v0@.
recipTailProcess ::
  (Subtractive a, Divisive a) =>
  a ->
  CP.Process a a
recipTailProcess :: forall a. (Subtractive a, Divisive a) => a -> Process a a
recipTailProcess a
u0 =
  (a -> ([a], [a]))
-> (([a], [a]) -> a -> ([a], [a]))
-> (([a], [a]) -> a)
-> Process a a
forall s a b. (a -> s) -> (s -> a -> s) -> (s -> b) -> Process a b
CP.Process a -> ([a], [a])
inject ([a], [a]) -> a -> ([a], [a])
step ([a], [a]) -> a
forall {a} {b}. ([a], b) -> a
extract
  where
    v0 :: a
v0 = a -> a
forall a. Divisive a => a -> a
recip a
u0
    inject :: a -> ([a], [a])
inject a
u = ([a], [a]) -> a -> ([a], [a])
step ([a
v0], []) a
u
    step :: ([a], [a]) -> a -> ([a], [a])
step ([a]
vs, [a]
us) a
u =
      let k :: Int
k = [a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
vs
          us' :: [a]
us' = [a]
us [a] -> [a] -> [a]
forall a. [a] -> [a] -> [a]
++ [a
u]
          vk :: a
vk = a -> a
forall a. Subtractive a => a -> a
negate ([a] -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum [[a]
us' [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! (Int
i Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1) a -> a -> a
forall a. Multiplicative a => a -> a -> a
* [a]
vs [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! (Int
k Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
i) | Int
i <- [Int
1 .. Int
k]]) a -> a -> a
forall a. Divisive a => a -> a -> a
/ a
u0
       in ([a]
vs [a] -> [a] -> [a]
forall a. [a] -> [a] -> [a]
++ [a
vk], [a]
us')
    extract :: ([a], b) -> a
extract ([a]
vs, b
_) = [a] -> a
forall a. HasCallStack => [a] -> a
last [a]
vs

-- | Reciprocal series as a coinductive stream.
--
-- @recipSeries u0 us@ produces @[v0, v1, ...]@ for @us = [u1, u2, ...]@.
-- The output length is one more than the input length, so truncation is
-- decided by the caller.
recipSeries ::
  (Subtractive a, Divisive a) =>
  a ->
  [a] ->
  [a]
recipSeries :: forall a. (Subtractive a, Divisive a) => a -> [a] -> [a]
recipSeries a
u0 [a]
us = a -> a
forall a. Divisive a => a -> a
recip a
u0 a -> [a] -> [a]
forall a. a -> [a] -> [a]
: Process a a -> [a] -> [a]
forall a b. Process a b -> [a] -> [b]
CP.scan (a -> Process a a
forall a. (Subtractive a, Divisive a) => a -> Process a a
recipTailProcess a
u0) [a]
us

instance
  (Additive a, Subtractive a, Multiplicative a, Divisive a) =>
  NumHask.Algebra.Multiplicative.Divisive (Jet a)
  where
  recip :: Jet a -> Jet a
recip (Jet []) = [a] -> Jet a
forall a. [a] -> Jet a
Jet []
  recip (Jet (a
u0 : [a]
us)) = [a] -> Jet a
forall a. [a] -> Jet a
Jet (a -> [a] -> [a]
forall a. (Subtractive a, Divisive a) => a -> [a] -> [a]
recipSeries a
u0 [a]
us)

-- ---------------------------------------------------------------------------
-- Field instances (exp / log / trig)
-- ---------------------------------------------------------------------------

-- | Term-by-term differentiation of a Taylor series.
--
-- > differentiate (Jet [c0, c1, c2, c3]) = Jet [c1, 2*c2, 3*c3]
differentiate ::
  (Multiplicative a, FromInteger a) =>
  Jet a ->
  Jet a
differentiate :: forall a. (Multiplicative a, FromInteger a) => Jet a -> Jet a
differentiate (Jet [a]
cs) =
  [a] -> Jet a
forall a. [a] -> Jet a
Jet [Integer -> a
forall a. FromInteger a => Integer -> a
fromInteger (Int -> Integer
forall a b. FromIntegral a b => b -> a
fromIntegral (Int
k :: Int)) a -> a -> a
forall a. Multiplicative a => a -> a -> a
* a
c | (Int
k, a
c) <- [Int] -> [a] -> [(Int, a)]
forall a b. [a] -> [b] -> [(a, b)]
zip [(Int
1 :: Int) ..] (Int -> [a] -> [a]
forall a. Int -> [a] -> [a]
drop Int
1 [a]
cs)]

-- | Term-by-term integration with supplied constant.
--
-- > integrate c0 (Jet [d0, d1, d2]) = Jet [c0, d0, d1/2, d2/3]
integrate ::
  (Divisive a, FromInteger a) =>
  a ->
  Jet a ->
  Jet a
integrate :: forall a. (Divisive a, FromInteger a) => a -> Jet a -> Jet a
integrate a
c0 (Jet [a]
ds) =
  [a] -> Jet a
forall a. [a] -> Jet a
Jet (a
c0 a -> [a] -> [a]
forall a. a -> [a] -> [a]
: [a
d a -> a -> a
forall a. Divisive a => a -> a -> a
/ Integer -> a
forall a. FromInteger a => Integer -> a
fromInteger (Int -> Integer
forall a b. FromIntegral a b => b -> a
fromIntegral (Int
k :: Int)) | (Int
k, a
d) <- [Int] -> [a] -> [(Int, a)]
forall a b. [a] -> [b] -> [(a, b)]
zip [(Int
1 :: Int) ..] [a]
ds])

-- | Scale every coefficient by a scalar.
scale :: (Multiplicative a) => a -> Jet a -> Jet a
scale :: forall a. Multiplicative a => a -> Jet a -> Jet a
scale a
s (Jet [a]
cs) = [a] -> Jet a
forall a. [a] -> Jet a
Jet ((a -> a) -> [a] -> [a]
forall a b. (a -> b) -> [a] -> [b]
map (a
s a -> a -> a
forall a. Multiplicative a => a -> a -> a
*) [a]
cs)

-- | Tail process for the mutual sin/cos series.
--
-- Around a primal point @u0@, the recurrences are
-- @m s_m = sum_{j=0}^{m-1} (m-j) c_j u_{m-j}@ and
-- @m c_m = -sum_{j=0}^{m-1} (m-j) s_j u_{m-j}@.
-- The process consumes @u1, u2, ...@ and emits @(s1, c1), (s2, c2), ...@;
-- the caller prepends @(sin u0, cos u0)@.
sinCosTailProcess ::
  (TrigField a, FromInteger a) =>
  a ->
  CP.Process a (a, a)
sinCosTailProcess :: forall a. (TrigField a, FromInteger a) => a -> Process a (a, a)
sinCosTailProcess a
u0 =
  (a -> ([a], [a], [a]))
-> (([a], [a], [a]) -> a -> ([a], [a], [a]))
-> (([a], [a], [a]) -> (a, a))
-> Process a (a, a)
forall s a b. (a -> s) -> (s -> a -> s) -> (s -> b) -> Process a b
CP.Process a -> ([a], [a], [a])
inject ([a], [a], [a]) -> a -> ([a], [a], [a])
forall {a}.
(FromInteger a, Divisive a, Subtractive a) =>
([a], [a], [a]) -> a -> ([a], [a], [a])
step ([a], [a], [a]) -> (a, a)
forall {a} {b} {c}. ([a], [b], c) -> (a, b)
extract
  where
    s0 :: a
s0 = a -> a
forall a. TrigField a => a -> a
sin a
u0
    c0 :: a
c0 = a -> a
forall a. TrigField a => a -> a
cos a
u0
    inject :: a -> ([a], [a], [a])
inject a
u = ([a], [a], [a]) -> a -> ([a], [a], [a])
forall {a}.
(FromInteger a, Divisive a, Subtractive a) =>
([a], [a], [a]) -> a -> ([a], [a], [a])
step ([a
s0], [a
c0], []) a
u
    step :: ([a], [a], [a]) -> a -> ([a], [a], [a])
step ([a]
ss, [a]
cs, [a]
us) a
u =
      let k :: Int
k = [a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
ss
          us' :: [a]
us' = [a]
us [a] -> [a] -> [a]
forall a. [a] -> [a] -> [a]
++ [a
u]
          m' :: a
m' = Integer -> a
forall a. FromInteger a => Integer -> a
fromInteger (Int -> Integer
forall a b. FromIntegral a b => b -> a
fromIntegral Int
k)
          sSum :: a
sSum = [a] -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum [Integer -> a
forall a. FromInteger a => Integer -> a
fromInteger (Int -> Integer
forall a b. FromIntegral a b => b -> a
fromIntegral (Int
k Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
j)) a -> a -> a
forall a. Multiplicative a => a -> a -> a
* ([a]
cs [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! Int
j) a -> a -> a
forall a. Multiplicative a => a -> a -> a
* ([a]
us' [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! (Int
k Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1 Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
j)) | Int
j <- [Int
0 .. Int
k Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]
          cSum :: a
cSum = [a] -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum [Integer -> a
forall a. FromInteger a => Integer -> a
fromInteger (Int -> Integer
forall a b. FromIntegral a b => b -> a
fromIntegral (Int
k Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
j)) a -> a -> a
forall a. Multiplicative a => a -> a -> a
* ([a]
ss [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! Int
j) a -> a -> a
forall a. Multiplicative a => a -> a -> a
* ([a]
us' [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! (Int
k Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1 Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
j)) | Int
j <- [Int
0 .. Int
k Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]
          sk :: a
sk = (a
forall a. Multiplicative a => a
one a -> a -> a
forall a. Divisive a => a -> a -> a
/ a
m') a -> a -> a
forall a. Multiplicative a => a -> a -> a
* a
sSum
          ck :: a
ck = a -> a
forall a. Subtractive a => a -> a
negate (a
forall a. Multiplicative a => a
one a -> a -> a
forall a. Divisive a => a -> a -> a
/ a
m') a -> a -> a
forall a. Multiplicative a => a -> a -> a
* a
cSum
       in ([a]
ss [a] -> [a] -> [a]
forall a. [a] -> [a] -> [a]
++ [a
sk], [a]
cs [a] -> [a] -> [a]
forall a. [a] -> [a] -> [a]
++ [a
ck], [a]
us')
    extract :: ([a], [b], c) -> (a, b)
extract ([a]
ss, [b]
cs, c
_) = ([a] -> a
forall a. HasCallStack => [a] -> a
last [a]
ss, [b] -> b
forall a. HasCallStack => [a] -> a
last [b]
cs)

-- | Simultaneously compute the Taylor coefficients of sin(u) and cos(u)
-- around a primal point @u0@.
sinCosSeries ::
  (TrigField a, FromInteger a) =>
  a ->
  [a] ->
  (Jet a, Jet a)
sinCosSeries :: forall a.
(TrigField a, FromInteger a) =>
a -> [a] -> (Jet a, Jet a)
sinCosSeries a
u0 [a]
us =
  let pairs :: [(a, a)]
pairs = (a -> a
forall a. TrigField a => a -> a
sin a
u0, a -> a
forall a. TrigField a => a -> a
cos a
u0) (a, a) -> [(a, a)] -> [(a, a)]
forall a. a -> [a] -> [a]
: Process a (a, a) -> [a] -> [(a, a)]
forall a b. Process a b -> [a] -> [b]
CP.scan (a -> Process a (a, a)
forall a. (TrigField a, FromInteger a) => a -> Process a (a, a)
sinCosTailProcess a
u0) [a]
us
      ([a]
ss, [a]
cs) = [(a, a)] -> ([a], [a])
forall a b. [(a, b)] -> ([a], [b])
unzip [(a, a)]
pairs
   in ([a] -> Jet a
forall a. [a] -> Jet a
Jet [a]
ss, [a] -> Jet a
forall a. [a] -> Jet a
Jet [a]
cs)

-- | Tail process for the square-root series.
--
-- If @u = v²@ with @u = u0 + u1*h + ...@ and @v = v0 + v1*h + ...@ then
-- @v0 = sqrt(u0)@ and @vk = (uk - sum_{i=1}^{k-1} vi v{k-i}) / (2 v0)@.
-- The process consumes @u1, u2, ...@ and emits @v1, v2, ...@; the caller
-- prepends @v0@.
sqrtTailProcess ::
  (ExpField a) =>
  a ->
  CP.Process a a
sqrtTailProcess :: forall a. ExpField a => a -> Process a a
sqrtTailProcess a
u0 =
  (a -> ([a], [a]))
-> (([a], [a]) -> a -> ([a], [a]))
-> (([a], [a]) -> a)
-> Process a a
forall s a b. (a -> s) -> (s -> a -> s) -> (s -> b) -> Process a b
CP.Process a -> ([a], [a])
inject ([a], [a]) -> a -> ([a], [a])
step ([a], [a]) -> a
forall {a} {b}. ([a], b) -> a
extract
  where
    v0 :: a
v0 = a -> a
forall a. ExpField a => a -> a
sqrt a
u0
    twoV0 :: a
twoV0 = a
v0 a -> a -> a
forall a. Additive a => a -> a -> a
+ a
v0
    inject :: a -> ([a], [a])
inject a
u = ([a], [a]) -> a -> ([a], [a])
step ([a
v0], []) a
u
    step :: ([a], [a]) -> a -> ([a], [a])
step ([a]
vs, [a]
us) a
u =
      let k :: Int
k = [a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
vs
          us' :: [a]
us' = [a]
us [a] -> [a] -> [a]
forall a. [a] -> [a] -> [a]
++ [a
u]
          inner :: a
inner = [a] -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum [[a]
vs [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! Int
i a -> a -> a
forall a. Multiplicative a => a -> a -> a
* [a]
vs [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! (Int
k Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
i) | Int
i <- [Int
1 .. Int
k Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]
          vk :: a
vk = ([a]
us' [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! (Int
k Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1) a -> a -> a
forall a. Subtractive a => a -> a -> a
- a
inner) a -> a -> a
forall a. Divisive a => a -> a -> a
/ a
twoV0
       in ([a]
vs [a] -> [a] -> [a]
forall a. [a] -> [a] -> [a]
++ [a
vk], [a]
us')
    extract :: ([a], b) -> a
extract ([a]
vs, b
_) = [a] -> a
forall a. HasCallStack => [a] -> a
last [a]
vs

-- | Square-root series as a coinductive stream.
sqrtSeries ::
  (ExpField a) =>
  a ->
  [a] ->
  [a]
sqrtSeries :: forall a. ExpField a => a -> [a] -> [a]
sqrtSeries a
u0 [a]
us = a -> a
forall a. ExpField a => a -> a
sqrt a
u0 a -> [a] -> [a]
forall a. a -> [a] -> [a]
: Process a a -> [a] -> [a]
forall a b. Process a b -> [a] -> [b]
CP.scan (a -> Process a a
forall a. ExpField a => a -> Process a a
sqrtTailProcess a
u0) [a]
us

instance (FromInteger a) => FromInteger (Jet a) where
  fromInteger :: Integer -> Jet a
fromInteger Integer
n = [a] -> Jet a
forall a. [a] -> Jet a
Jet [Integer -> a
forall a. FromInteger a => Integer -> a
fromInteger Integer
n]

-- | Tail process for the exponential series.
--
-- Solves @v' = v * u'@ with @v0 = exp(u0)@ coefficient-wise:
-- @m v_m = sum_{j=0}^{m-1} (m-j) v_j u_{m-j}@.
-- The process consumes @u1, u2, ...@ and emits @v1, v2, ...@; the caller
-- prepends @v0@.
expTailProcess ::
  (ExpField a, FromInteger a) =>
  a ->
  CP.Process a a
expTailProcess :: forall a. (ExpField a, FromInteger a) => a -> Process a a
expTailProcess a
u0 =
  (a -> ([a], [a]))
-> (([a], [a]) -> a -> ([a], [a]))
-> (([a], [a]) -> a)
-> Process a a
forall s a b. (a -> s) -> (s -> a -> s) -> (s -> b) -> Process a b
CP.Process a -> ([a], [a])
inject ([a], [a]) -> a -> ([a], [a])
forall {a}.
(FromInteger a, Divisive a, Additive a) =>
([a], [a]) -> a -> ([a], [a])
step ([a], [a]) -> a
forall {a} {b}. ([a], b) -> a
extract
  where
    v0 :: a
v0 = a -> a
forall a. ExpField a => a -> a
exp a
u0
    inject :: a -> ([a], [a])
inject a
u = ([a], [a]) -> a -> ([a], [a])
forall {a}.
(FromInteger a, Divisive a, Additive a) =>
([a], [a]) -> a -> ([a], [a])
step ([a
v0], []) a
u
    step :: ([a], [a]) -> a -> ([a], [a])
step ([a]
vs, [a]
us) a
u =
      let m :: Int
m = [a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
vs
          us' :: [a]
us' = [a]
us [a] -> [a] -> [a]
forall a. [a] -> [a] -> [a]
++ [a
u]
          m' :: a
m' = Integer -> a
forall a. FromInteger a => Integer -> a
fromInteger (Int -> Integer
forall a b. FromIntegral a b => b -> a
fromIntegral Int
m)
          vm :: a
vm =
            (a
forall a. Multiplicative a => a
one a -> a -> a
forall a. Divisive a => a -> a -> a
/ a
m')
              a -> a -> a
forall a. Multiplicative a => a -> a -> a
* [a] -> a
forall a (f :: * -> *). (Additive a, Foldable f) => f a -> a
sum
                [ Integer -> a
forall a. FromInteger a => Integer -> a
fromInteger (Int -> Integer
forall a b. FromIntegral a b => b -> a
fromIntegral (Int
m Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
j)) a -> a -> a
forall a. Multiplicative a => a -> a -> a
* ([a]
vs [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! Int
j) a -> a -> a
forall a. Multiplicative a => a -> a -> a
* ([a]
us' [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! (Int
m Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1 Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
j))
                | Int
j <- [Int
0 .. Int
m Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]
                ]
       in ([a]
vs [a] -> [a] -> [a]
forall a. [a] -> [a] -> [a]
++ [a
vm], [a]
us')
    extract :: ([a], b) -> a
extract ([a]
vs, b
_) = [a] -> a
forall a. HasCallStack => [a] -> a
last [a]
vs

-- | Exponential series as a coinductive stream.
expSeries ::
  (ExpField a, FromInteger a) =>
  a ->
  [a] ->
  [a]
expSeries :: forall a. (ExpField a, FromInteger a) => a -> [a] -> [a]
expSeries a
u0 [a]
us = a -> a
forall a. ExpField a => a -> a
exp a
u0 a -> [a] -> [a]
forall a. a -> [a] -> [a]
: Process a a -> [a] -> [a]
forall a b. Process a b -> [a] -> [b]
CP.scan (a -> Process a a
forall a. (ExpField a, FromInteger a) => a -> Process a a
expTailProcess a
u0) [a]
us

instance (Subtractive a, Divisive a, ExpField a, FromInteger a) => ExpField (Jet a) where
  exp :: Jet a -> Jet a
exp (Jet []) = [a] -> Jet a
forall a. [a] -> Jet a
Jet []
  exp (Jet (a
u0 : [a]
us)) = [a] -> Jet a
forall a. [a] -> Jet a
Jet (a -> [a] -> [a]
forall a. (ExpField a, FromInteger a) => a -> [a] -> [a]
expSeries a
u0 [a]
us)

  log :: Jet a -> Jet a
log (Jet []) = [a] -> Jet a
forall a. [a] -> Jet a
Jet []
  log (Jet (a
u0 : [a]
us)) =
    let u :: Jet a
u = [a] -> Jet a
forall a. [a] -> Jet a
Jet (a
u0 a -> [a] -> [a]
forall a. a -> [a] -> [a]
: [a]
us)
     in a -> Jet a -> Jet a
forall a. (Divisive a, FromInteger a) => a -> Jet a -> Jet a
integrate (a -> a
forall a. ExpField a => a -> a
log a
u0) (Jet a -> Jet a
forall a. (Multiplicative a, FromInteger a) => Jet a -> Jet a
differentiate Jet a
u Jet a -> Jet a -> Jet a
forall a. Divisive a => a -> a -> a
/ Jet a
u)

  sqrt :: Jet a -> Jet a
sqrt (Jet []) = [a] -> Jet a
forall a. [a] -> Jet a
Jet []
  sqrt (Jet (a
u0 : [a]
us)) = [a] -> Jet a
forall a. [a] -> Jet a
Jet (a -> [a] -> [a]
forall a. ExpField a => a -> [a] -> [a]
sqrtSeries a
u0 [a]
us)

instance (Subtractive a, Divisive a, ExpField a, TrigField a, FromInteger a) => TrigField (Jet a) where
  pi :: Jet a
pi = [a] -> Jet a
forall a. [a] -> Jet a
Jet [a
forall a. TrigField a => a
pi]

  sin :: Jet a -> Jet a
sin (Jet []) = [a] -> Jet a
forall a. [a] -> Jet a
Jet []
  sin (Jet (a
u0 : [a]
us)) =
    let (Jet a
ss, Jet a
_) = a -> [a] -> (Jet a, Jet a)
forall a.
(TrigField a, FromInteger a) =>
a -> [a] -> (Jet a, Jet a)
sinCosSeries a
u0 [a]
us
     in Jet a
ss

  cos :: Jet a -> Jet a
cos (Jet []) = [a] -> Jet a
forall a. [a] -> Jet a
Jet []
  cos (Jet (a
u0 : [a]
us)) =
    let (Jet a
_, Jet a
cs) = a -> [a] -> (Jet a, Jet a)
forall a.
(TrigField a, FromInteger a) =>
a -> [a] -> (Jet a, Jet a)
sinCosSeries a
u0 [a]
us
     in Jet a
cs

  asin :: Jet a -> Jet a
asin (Jet []) = [a] -> Jet a
forall a. [a] -> Jet a
Jet []
  asin (Jet (a
u0 : [a]
us)) =
    let u :: Jet a
u = [a] -> Jet a
forall a. [a] -> Jet a
Jet (a
u0 a -> [a] -> [a]
forall a. a -> [a] -> [a]
: [a]
us)
     in a -> Jet a -> Jet a
forall a. (Divisive a, FromInteger a) => a -> Jet a -> Jet a
integrate (a -> a
forall a. TrigField a => a -> a
asin a
u0) (Jet a -> Jet a
forall a. (Multiplicative a, FromInteger a) => Jet a -> Jet a
differentiate Jet a
u Jet a -> Jet a -> Jet a
forall a. Divisive a => a -> a -> a
/ Jet a -> Jet a
forall a. ExpField a => a -> a
sqrt (Jet a
forall a. Multiplicative a => a
one Jet a -> Jet a -> Jet a
forall a. Subtractive a => a -> a -> a
- Jet a
u Jet a -> Jet a -> Jet a
forall a. Multiplicative a => a -> a -> a
* Jet a
u))

  acos :: Jet a -> Jet a
acos (Jet []) = [a] -> Jet a
forall a. [a] -> Jet a
Jet []
  acos (Jet (a
u0 : [a]
us)) =
    let u :: Jet a
u = [a] -> Jet a
forall a. [a] -> Jet a
Jet (a
u0 a -> [a] -> [a]
forall a. a -> [a] -> [a]
: [a]
us)
     in a -> Jet a -> Jet a
forall a. (Divisive a, FromInteger a) => a -> Jet a -> Jet a
integrate (a -> a
forall a. TrigField a => a -> a
acos a
u0) (Jet a -> Jet a
forall a. Subtractive a => a -> a
negate (Jet a -> Jet a
forall a. (Multiplicative a, FromInteger a) => Jet a -> Jet a
differentiate Jet a
u Jet a -> Jet a -> Jet a
forall a. Divisive a => a -> a -> a
/ Jet a -> Jet a
forall a. ExpField a => a -> a
sqrt (Jet a
forall a. Multiplicative a => a
one Jet a -> Jet a -> Jet a
forall a. Subtractive a => a -> a -> a
- Jet a
u Jet a -> Jet a -> Jet a
forall a. Multiplicative a => a -> a -> a
* Jet a
u)))

  atan :: Jet a -> Jet a
atan (Jet []) = [a] -> Jet a
forall a. [a] -> Jet a
Jet []
  atan (Jet (a
u0 : [a]
us)) =
    let u :: Jet a
u = [a] -> Jet a
forall a. [a] -> Jet a
Jet (a
u0 a -> [a] -> [a]
forall a. a -> [a] -> [a]
: [a]
us)
     in a -> Jet a -> Jet a
forall a. (Divisive a, FromInteger a) => a -> Jet a -> Jet a
integrate (a -> a
forall a. TrigField a => a -> a
atan a
u0) (Jet a -> Jet a
forall a. (Multiplicative a, FromInteger a) => Jet a -> Jet a
differentiate Jet a
u Jet a -> Jet a -> Jet a
forall a. Divisive a => a -> a -> a
/ (Jet a
forall a. Multiplicative a => a
one Jet a -> Jet a -> Jet a
forall a. Additive a => a -> a -> a
+ Jet a
u Jet a -> Jet a -> Jet a
forall a. Multiplicative a => a -> a -> a
* Jet a
u))

  atan2 :: Jet a -> Jet a -> Jet a
atan2 Jet a
y Jet a
x =
    let y0 :: a
y0 = Jet a -> a
forall a. Jet a -> a
headCoeff Jet a
y
        x0 :: a
x0 = Jet a -> a
forall a. Jet a -> a
headCoeff Jet a
x
        deriv :: Jet a
deriv = (Jet a
x Jet a -> Jet a -> Jet a
forall a. Multiplicative a => a -> a -> a
* Jet a -> Jet a
forall a. (Multiplicative a, FromInteger a) => Jet a -> Jet a
differentiate Jet a
y Jet a -> Jet a -> Jet a
forall a. Subtractive a => a -> a -> a
- Jet a
y Jet a -> Jet a -> Jet a
forall a. Multiplicative a => a -> a -> a
* Jet a -> Jet a
forall a. (Multiplicative a, FromInteger a) => Jet a -> Jet a
differentiate Jet a
x) Jet a -> Jet a -> Jet a
forall a. Divisive a => a -> a -> a
/ (Jet a
x Jet a -> Jet a -> Jet a
forall a. Multiplicative a => a -> a -> a
* Jet a
x Jet a -> Jet a -> Jet a
forall a. Additive a => a -> a -> a
+ Jet a
y Jet a -> Jet a -> Jet a
forall a. Multiplicative a => a -> a -> a
* Jet a
y)
     in a -> Jet a -> Jet a
forall a. (Divisive a, FromInteger a) => a -> Jet a -> Jet a
integrate (a -> a -> a
forall a. TrigField a => a -> a -> a
atan2 a
y0 a
x0) Jet a
deriv

  sinh :: Jet a -> Jet a
sinh Jet a
u = (Jet a -> Jet a
forall a. ExpField a => a -> a
exp Jet a
u Jet a -> Jet a -> Jet a
forall a. Subtractive a => a -> a -> a
- Jet a -> Jet a
forall a. ExpField a => a -> a
exp (Jet a -> Jet a
forall a. Subtractive a => a -> a
negate Jet a
u)) Jet a -> Jet a -> Jet a
forall a. Divisive a => a -> a -> a
/ (Jet a
forall a. Multiplicative a => a
one Jet a -> Jet a -> Jet a
forall a. Additive a => a -> a -> a
+ Jet a
forall a. Multiplicative a => a
one)
  cosh :: Jet a -> Jet a
cosh Jet a
u = (Jet a -> Jet a
forall a. ExpField a => a -> a
exp Jet a
u Jet a -> Jet a -> Jet a
forall a. Additive a => a -> a -> a
+ Jet a -> Jet a
forall a. ExpField a => a -> a
exp (Jet a -> Jet a
forall a. Subtractive a => a -> a
negate Jet a
u)) Jet a -> Jet a -> Jet a
forall a. Divisive a => a -> a -> a
/ (Jet a
forall a. Multiplicative a => a
one Jet a -> Jet a -> Jet a
forall a. Additive a => a -> a -> a
+ Jet a
forall a. Multiplicative a => a
one)

  asinh :: Jet a -> Jet a
asinh Jet a
u = Jet a -> Jet a
forall a. ExpField a => a -> a
log (Jet a
u Jet a -> Jet a -> Jet a
forall a. Additive a => a -> a -> a
+ Jet a -> Jet a
forall a. ExpField a => a -> a
sqrt (Jet a
u Jet a -> Jet a -> Jet a
forall a. Multiplicative a => a -> a -> a
* Jet a
u Jet a -> Jet a -> Jet a
forall a. Additive a => a -> a -> a
+ Jet a
forall a. Multiplicative a => a
one))
  acosh :: Jet a -> Jet a
acosh Jet a
u = Jet a -> Jet a
forall a. ExpField a => a -> a
log (Jet a
u Jet a -> Jet a -> Jet a
forall a. Additive a => a -> a -> a
+ Jet a -> Jet a
forall a. ExpField a => a -> a
sqrt (Jet a
u Jet a -> Jet a -> Jet a
forall a. Multiplicative a => a -> a -> a
* Jet a
u Jet a -> Jet a -> Jet a
forall a. Subtractive a => a -> a -> a
- Jet a
forall a. Multiplicative a => a
one))
  atanh :: Jet a -> Jet a
atanh Jet a
u = Jet a -> Jet a
forall a. ExpField a => a -> a
log ((Jet a
forall a. Multiplicative a => a
one Jet a -> Jet a -> Jet a
forall a. Additive a => a -> a -> a
+ Jet a
u) Jet a -> Jet a -> Jet a
forall a. Divisive a => a -> a -> a
/ (Jet a
forall a. Multiplicative a => a
one Jet a -> Jet a -> Jet a
forall a. Subtractive a => a -> a -> a
- Jet a
u)) Jet a -> Jet a -> Jet a
forall a. Divisive a => a -> a -> a
/ (Jet a
forall a. Multiplicative a => a
one Jet a -> Jet a -> Jet a
forall a. Additive a => a -> a -> a
+ Jet a
forall a. Multiplicative a => a
one)

-- | Constant coefficient of a jet.
headCoeff :: Jet a -> a
headCoeff :: forall a. Jet a -> a
headCoeff (Jet []) = String -> a
forall a. HasCallStack => String -> a
error String
"Circuit.Diff.Jet.headCoeff: empty jet"
headCoeff (Jet (a
c : [a]
_)) = a
c