module Circuit.RL.Estimator
(
Dual (..),
reinforceGrad,
pathwiseGrad,
closedFormGrad,
)
where
import Prelude
newtype Dual a = Dual {forall a. Dual a -> (a, a)
getDual :: (a, a)}
deriving stock (Dual a -> Dual a -> Bool
(Dual a -> Dual a -> Bool)
-> (Dual a -> Dual a -> Bool) -> Eq (Dual a)
forall a. Eq a => Dual a -> Dual a -> Bool
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: forall a. Eq a => Dual a -> Dual a -> Bool
== :: Dual a -> Dual a -> Bool
$c/= :: forall a. Eq a => Dual a -> Dual a -> Bool
/= :: Dual a -> Dual a -> Bool
Eq, Int -> Dual a -> ShowS
[Dual a] -> ShowS
Dual a -> String
(Int -> Dual a -> ShowS)
-> (Dual a -> String) -> ([Dual a] -> ShowS) -> Show (Dual a)
forall a. Show a => Int -> Dual a -> ShowS
forall a. Show a => [Dual a] -> ShowS
forall a. Show a => Dual a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> Dual a -> ShowS
showsPrec :: Int -> Dual a -> ShowS
$cshow :: forall a. Show a => Dual a -> String
show :: Dual a -> String
$cshowList :: forall a. Show a => [Dual a] -> ShowS
showList :: [Dual a] -> ShowS
Show)
instance (Num a) => Num (Dual a) where
Dual (a
v1, a
g1) + :: Dual a -> Dual a -> Dual a
+ Dual (a
v2, a
g2) = (a, a) -> Dual a
forall a. (a, a) -> Dual a
Dual (a
v1 a -> a -> a
forall a. Num a => a -> a -> a
+ a
v2, a
g1 a -> a -> a
forall a. Num a => a -> a -> a
+ a
g2)
Dual (a
v1, a
g1) * :: Dual a -> Dual a -> Dual a
* Dual (a
v2, a
g2) =
(a, a) -> Dual a
forall a. (a, a) -> Dual a
Dual (a
v1 a -> a -> a
forall a. Num a => a -> a -> a
* a
v2, a
v1 a -> a -> a
forall a. Num a => a -> a -> a
* a
g2 a -> a -> a
forall a. Num a => a -> a -> a
+ a
g1 a -> a -> a
forall a. Num a => a -> a -> a
* a
v2)
negate :: Dual a -> Dual a
negate (Dual (a
v, a
g)) = (a, a) -> Dual a
forall a. (a, a) -> Dual a
Dual (a -> a
forall a. Num a => a -> a
negate a
v, a -> a
forall a. Num a => a -> a
negate a
g)
abs :: Dual a -> Dual a
abs = String -> Dual a -> Dual a
forall a. HasCallStack => String -> a
error String
"Dual: abs not defined"
signum :: Dual a -> Dual a
signum = String -> Dual a -> Dual a
forall a. HasCallStack => String -> a
error String
"Dual: signum not defined"
fromInteger :: Integer -> Dual a
fromInteger Integer
n = (a, a) -> Dual a
forall a. (a, a) -> Dual a
Dual (Integer -> a
forall a. Num a => Integer -> a
fromInteger Integer
n, a
0)
instance (Fractional a) => Fractional (Dual a) where
fromRational :: Rational -> Dual a
fromRational Rational
r = (a, a) -> Dual a
forall a. (a, a) -> Dual a
Dual (Rational -> a
forall a. Fractional a => Rational -> a
fromRational Rational
r, a
0)
recip :: Dual a -> Dual a
recip (Dual (a
v, a
g)) =
(a, a) -> Dual a
forall a. (a, a) -> Dual a
Dual (a -> a
forall a. Fractional a => a -> a
recip a
v, a -> a
forall a. Num a => a -> a
negate a
g a -> a -> a
forall a. Fractional a => a -> a -> a
/ (a
v a -> a -> a
forall a. Num a => a -> a -> a
* a
v))
{-# INLINE recip #-}
gaussMoment :: Double -> Int -> Double
gaussMoment :: Double -> Int -> Double
gaussMoment Double
sigma Int
k
| Int -> Bool
forall a. Integral a => a -> Bool
odd Int
k = Double
0
| Bool
otherwise =
let sigmaSq :: Double
sigmaSq = Double
sigma Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
sigma
doubleFact :: a -> a
doubleFact a
n = [a] -> a
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
product [a
n, a
n a -> a -> a
forall a. Num a => a -> a -> a
- a
2 .. a
2]
in (Double
sigmaSq Double -> Double -> Double
forall a. Floating a => a -> a -> a
** Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int
k Int -> Int -> Int
forall a. Integral a => a -> a -> a
`div` Int
2)) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Int
forall {a}. (Num a, Enum a) => a -> a
doubleFact (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1))
closedFormGrad :: Double -> Double -> Double -> Double
closedFormGrad :: Double -> Double -> Double -> Double
closedFormGrad Double
theta Double
_sigma Double
aTarget = -(Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* (Double
theta Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
aTarget))
reinforceGrad :: Double -> Double -> Double -> Double
reinforceGrad :: Double -> Double -> Double -> Double
reinforceGrad Double
theta Double
sigma Double
aTarget =
let sigmaSq :: Double
sigmaSq = Double
sigma Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
sigma
delta :: Double
delta = Double
theta Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
aTarget
cubicTerm :: Double
cubicTerm = Double -> Int -> Double
gaussMoment Double
sigma Int
3
quadTerm :: Double
quadTerm = Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
delta Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Int -> Double
gaussMoment Double
sigma Int
2
linearTerm :: Double
linearTerm = Double
delta Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
delta Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Int -> Double
gaussMoment Double
sigma Int
1
in -((Double
cubicTerm Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
quadTerm Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
linearTerm) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
sigmaSq)
pathwiseGrad :: Double -> Double -> Double -> Double
pathwiseGrad :: Double -> Double -> Double -> Double
pathwiseGrad Double
theta Double
sigma Double
aTarget =
(-(Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
theta)) Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
aTarget) Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (-(Double
2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
sigma Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double -> Int -> Double
gaussMoment Double
1 Int
1))