-- | Policy-gradient estimator axes for the REINFORCE == pathwise oracle.
--
-- The instance table (loom/instance-table.md §1) names two estimator axes:
--
-- * REINFORCE (score-function): @Prob (->) (Dual r)@ — the continuation
--   carries a gradient component that accumulates via the dual-number product
--   rule: @(v,g)·(v',g') = (v·v', v·g' + g·v')@.  The log-policy score
--   function injects gradient-weighted-by-reward exactly when the continuation
--   multiplies by a score morphism.
--
-- * Pathwise (reparametrization): the derivative is taken through the sample
--   @a = θ + σε@, so @∂R/∂θ = ∂R/∂a · 1@.
--
-- The oracle is: on a 1-step Gaussian policy @N(θ, σ²=0.25)@ with quadratic
-- reward @R(a) = -(a-1)²@, both estimators analytically reduce to
-- @∇J = -2(θ-1)@.  At @θ₀=3@ this is @-4@ — exact in Double.
module Circuit.RL.Estimator
  ( -- * Dual scalar for score-function gradient accumulation
    Dual (..),

    -- * Estimators
    reinforceGrad,
    pathwiseGrad,
    closedFormGrad,
  )
where

import Prelude

-- ---------------------------------------------------------------------------
-- Dual scalar — the score-function axis
-- ---------------------------------------------------------------------------

-- | Dual-number scalar: value plus accumulated gradient.
--
-- Multiplication follows the dual-number product rule: the gradient
-- component records the cross-derivative.  This naturally gives the
-- REINFORCE coupling: when a continuation multiplies a downstream reward
-- (value component) by a score function (gradient component), the result
-- records the reward-weighted score.
--
-- >>> Dual (2, 3) * Dual (5, 7)
-- Dual (10, 29)
--
-- Check: 2·5 = 10, 2·7 + 3·5 = 14+15 = 29 ✓
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 #-}

-- ---------------------------------------------------------------------------
-- Gaussian central moments (exact, polynomial)
-- ---------------------------------------------------------------------------

-- | The k-th central moment E[(a-μ)^k] for a ~ N(μ, σ²).  For symmetric
-- Gaussians, odd moments vanish.  Even moments: E[z^(2m)] = σ^(2m)·(2m-1)!!.
--
-- >>> gaussMoment 0.5 0
-- 1.0
-- >>> gaussMoment 0.5 2
-- 0.25
-- >>> gaussMoment 0.5 4
-- 0.1875
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))

-- ---------------------------------------------------------------------------
-- Estimators
-- ---------------------------------------------------------------------------

-- | Closed-form gradient for Gaussian policy N(θ, σ²) with quadratic reward
-- @R(a) = -(a - a_target)²@.
--
-- @J(θ) = E[R(a)] = -(σ² + (θ - a_target)²)@, so @∇J = -2(θ - a_target)@.
--
-- >>> closedFormGrad 3.0 0.5 1.0
-- -4.0
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))

-- | REINFORCE (score-function) gradient estimator for Gaussian policy.
--
-- @∇J = E[R(a) · ∇_θ log π(a)]@ where @∇_θ log π(a) = (a-θ)/σ²@.
--
-- The analytical expectation uses Gaussian central moments.
-- Let @z = a-θ@, @δ = θ-a_target@:
--
-- @
--   R(a) = -(a - a_target)² = -(z+δ)²
--   R(a)(a-θ)/σ² = -(z+δ)²·z/σ² = -(z³ + 2δz² + δ²z)/σ²
--   E[R(a)(a-θ)/σ²] = -(1/σ²)(E[z³] + 2δ·E[z²] + δ²·E[z])
--                   = -(1/σ²)(0 + 2δσ² + 0) = -2δ
--                   = -2(θ-a_target)
-- @
--
-- >>> reinforceGrad 3.0 0.5 1.0
-- -4.0
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
      -- E[-(a-a_target)²·(a-θ)/σ²] = -E[z³+2δz²+δ²z]/σ²
      -- where z = a-θ ~ N(0,σ²)
      cubicTerm :: Double
cubicTerm = Double -> Int -> Double
gaussMoment Double
sigma Int
3 -- 0
      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 -- 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 -- 0
   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)

-- | Pathwise (reparametrization) gradient estimator for Gaussian policy.
--
-- @a = θ + σε@ with @ε ~ N(0,1)@.
-- @∂R/∂θ = ∂R/∂a · ∂a/∂θ = -2(a - a_target) · 1 = -2(θ+σε-a_target)@.
-- @E[∂R/∂θ] = -2(θ-a_target)@ since @E[ε] = 0@.
--
-- >>> pathwiseGrad 3.0 0.5 1.0
-- -4.0
pathwiseGrad :: Double -> Double -> Double -> Double
pathwiseGrad :: Double -> Double -> Double -> Double
pathwiseGrad Double
theta Double
sigma Double
aTarget =
  -- E[-2(θ+σε-a_target)] where ε ~ N(0,1), E[ε] = 0
  (-(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))