packages feed

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