srtree-3.0.0.0: src/Algorithm/SRTree/AD.hs
-----------------------------------------------------------------------------
-- |
-- Module : Data.SRTree.AD
-- Copyright : (c) Fabricio Olivetti 2021 - 2024
-- License : BSD3
-- Maintainer : fabricio.olivetti@gmail.com
-- Stability : experimental
-- Portability : FlexibleInstances, DeriveFunctor, ScopedTypeVariables
--
-- Automatic Differentiation for Expression trees
--
-----------------------------------------------------------------------------
module Algorithm.SRTree.AD
( compileFunAndGrad
, ADBackEnd(..)
) where
import qualified Data.Vector.Unboxed as VU
import qualified Data.Vector.Storable as V
import Data.SRTree
import Algorithm.SRTree.AD.Unboxed
data ADBackEnd = SingleThread | MultiThread deriving (Read, Show)
compileFunAndGrad :: ADBackEnd -> [VU.Vector Double] -> VU.Vector Double -> Maybe (VU.Vector Double) -> Fix SRTree -> V.Vector Double -> (Double, V.Vector Double)
compileFunAndGrad SingleThread xss ys mYerr tree =
let ct = compileTree xss ys mYerr tree
in \theta -> evalGradVec ct theta
compileFunAndGrad MultiThread xss ys mYerr tree =
let cts = compileTreeMulti xss ys mYerr tree
in \theta -> evalGradMulti cts theta