srtree 0.1.2.1 → 1.0.0.0
raw patch · 10 files changed
+738/−690 lines, 10 filesdep +HUnitdep +addep +criteriondep ~vectornew-component:exe:bench-srtreePVP ok
version bump matches the API change (PVP)
Dependencies added: HUnit, ad, criterion
Dependency ranges changed: vector
API changes (from Hackage documentation)
- Data.SRTree: (^.) :: OptIntPow a => a -> Int -> a
- Data.SRTree: Empty :: SRTree ix val
- Data.SRTree: Fun :: Function -> SRTree ix val -> SRTree ix val
- Data.SRTree: LogBase :: SRTree ix val -> SRTree ix val -> SRTree ix val
- Data.SRTree: Pow :: SRTree ix val -> Int -> SRTree ix val
- Data.SRTree: class OptIntPow a
- Data.SRTree: deriveParamBy :: (Eq ix, Eq val, Floating val, OptIntPow val) => ix -> SRTree ix val -> SRTree ix val
- Data.SRTree: evalTreeMap :: (Floating v1, OptIntPow v1, Floating v2, OptIntPow v2) => (v1 -> v2) -> SRTree ix v1 -> Reader (ix -> Maybe v2) (Maybe v2)
- Data.SRTree: evalTreeWithMap :: (Ord ix, Floating val, OptIntPow val) => SRTree ix val -> Map ix val -> Maybe val
- Data.SRTree: evalTreeWithVector :: (Floating val, OptIntPow val) => SRTree Int val -> Vector val -> Maybe val
- Data.SRTree: infix 8 ^.
- Data.SRTree: relabelOccurrences :: forall ix val. Ord ix => SRTree ix val -> SRTree (ix, Int) val
- Data.SRTree: simplify :: (Eq ix, Eq val, Floating val, OptIntPow val) => SRTree ix val -> SRTree ix val
- Data.SRTree: traverseIx :: Applicative f => (ixa -> f ixb) -> SRTree ixa val -> f (SRTree ixb val)
- Data.SRTree.Internal: (^.) :: OptIntPow a => a -> Int -> a
- Data.SRTree.Internal: Empty :: SRTree ix val
- Data.SRTree.Internal: Fun :: Function -> SRTree ix val -> SRTree ix val
- Data.SRTree.Internal: LogBase :: SRTree ix val -> SRTree ix val -> SRTree ix val
- Data.SRTree.Internal: Pow :: SRTree ix val -> Int -> SRTree ix val
- Data.SRTree.Internal: class OptIntPow a
- Data.SRTree.Internal: deriveParamBy :: (Eq ix, Eq val, Floating val, OptIntPow val) => ix -> SRTree ix val -> SRTree ix val
- Data.SRTree.Internal: evalTreeMap :: (Floating v1, OptIntPow v1, Floating v2, OptIntPow v2) => (v1 -> v2) -> SRTree ix v1 -> Reader (ix -> Maybe v2) (Maybe v2)
- Data.SRTree.Internal: evalTreeWithMap :: (Ord ix, Floating val, OptIntPow val) => SRTree ix val -> Map ix val -> Maybe val
- Data.SRTree.Internal: evalTreeWithVector :: (Floating val, OptIntPow val) => SRTree Int val -> Vector val -> Maybe val
- Data.SRTree.Internal: infix 8 ^.
- Data.SRTree.Internal: instance (GHC.Classes.Eq ix, GHC.Classes.Eq val) => GHC.Classes.Eq (Data.SRTree.Internal.SRTree ix val)
- Data.SRTree.Internal: instance (GHC.Classes.Eq ix, GHC.Classes.Eq val, GHC.Float.Floating val) => GHC.Float.Floating (Data.SRTree.Internal.SRTree ix val)
- Data.SRTree.Internal: instance (GHC.Classes.Eq ix, GHC.Classes.Eq val, GHC.Num.Num val) => GHC.Num.Num (Data.SRTree.Internal.SRTree ix val)
- Data.SRTree.Internal: instance (GHC.Classes.Eq ix, GHC.Classes.Eq val, GHC.Num.Num val, Data.SRTree.Internal.OptIntPow val) => Data.SRTree.Internal.OptIntPow (Data.SRTree.Internal.SRTree ix val)
- Data.SRTree.Internal: instance (GHC.Classes.Eq ix, GHC.Classes.Eq val, GHC.Real.Fractional val) => GHC.Real.Fractional (Data.SRTree.Internal.SRTree ix val)
- Data.SRTree.Internal: instance (GHC.Classes.Ord ix, GHC.Classes.Ord val) => GHC.Classes.Ord (Data.SRTree.Internal.SRTree ix val)
- Data.SRTree.Internal: instance (GHC.Show.Show ix, GHC.Show.Show val) => GHC.Show.Show (Data.SRTree.Internal.SRTree ix val)
- Data.SRTree.Internal: instance Data.Bifunctor.Bifunctor Data.SRTree.Internal.SRTree
- Data.SRTree.Internal: instance Data.Foldable.Foldable (Data.SRTree.Internal.SRTree ix)
- Data.SRTree.Internal: instance Data.SRTree.Internal.OptIntPow GHC.Types.Double
- Data.SRTree.Internal: instance Data.SRTree.Internal.OptIntPow GHC.Types.Float
- Data.SRTree.Internal: instance Data.Traversable.Traversable (Data.SRTree.Internal.SRTree ix)
- Data.SRTree.Internal: instance GHC.Base.Applicative (Data.SRTree.Internal.SRTree ix)
- Data.SRTree.Internal: instance GHC.Base.Functor (Data.SRTree.Internal.SRTree ix)
- Data.SRTree.Internal: relabelOccurrences :: forall ix val. Ord ix => SRTree ix val -> SRTree (ix, Int) val
- Data.SRTree.Internal: simplify :: (Eq ix, Eq val, Floating val, OptIntPow val) => SRTree ix val -> SRTree ix val
- Data.SRTree.Internal: traverseIx :: Applicative f => (ixa -> f ixb) -> SRTree ixa val -> f (SRTree ixb val)
- Data.SRTree.Print: D :: (ix -> String) -> (ix -> String) -> (val -> String) -> (Function -> String) -> String -> String -> DisplayNodes ix val
- Data.SRTree.Print: [_displayFloatPow] :: DisplayNodes ix val -> String
- Data.SRTree.Print: [_displayFun] :: DisplayNodes ix val -> Function -> String
- Data.SRTree.Print: [_displayPar] :: DisplayNodes ix val -> ix -> String
- Data.SRTree.Print: [_displayPow] :: DisplayNodes ix val -> String
- Data.SRTree.Print: [_displayVal] :: DisplayNodes ix val -> val -> String
- Data.SRTree.Print: [_displayVar] :: DisplayNodes ix val -> ix -> String
- Data.SRTree.Print: data DisplayNodes ix val
- Data.SRTree.Print: showDefault :: (Show ix, Show val) => SRTree ix val -> String
- Data.SRTree.Print: showTree :: (Show ix, Show val) => SRTree ix val -> DisplayNodes ix val -> String
+ Data.SRTree: Bin :: Op -> val -> val -> SRTree val
+ Data.SRTree: Uni :: Function -> val -> SRTree val
+ Data.SRTree: constsToParam :: Fix SRTree -> (Fix SRTree, [Double])
+ Data.SRTree: countConsts :: Fix SRTree -> Int
+ Data.SRTree: countParams :: Fix SRTree -> Int
+ Data.SRTree: data Op
+ Data.SRTree: deriveByParam :: Int -> Fix SRTree -> Fix SRTree
+ Data.SRTree: deriveByVar :: Int -> Fix SRTree -> Fix SRTree
+ Data.SRTree: evalOp :: Floating a => Op -> a -> a -> a
+ Data.SRTree: floatConstsToParam :: Fix SRTree -> (Fix SRTree, [Double])
+ Data.SRTree: forwardMode :: (Show a, Num a, Floating a) => Vector a -> Vector Double -> (Double -> a) -> Fix SRTree -> [a]
+ Data.SRTree: gradParams :: (Show a, Num a, Floating a) => Vector a -> Vector Double -> (Double -> a) -> Fix SRTree -> (a, [a])
+ Data.SRTree: param :: Int -> Fix SRTree
+ Data.SRTree: var :: Int -> Fix SRTree
+ Data.SRTree.Internal: Bin :: Op -> val -> val -> SRTree val
+ Data.SRTree.Internal: Uni :: Function -> val -> SRTree val
+ Data.SRTree.Internal: constsToParam :: Fix SRTree -> (Fix SRTree, [Double])
+ Data.SRTree.Internal: countConsts :: Fix SRTree -> Int
+ Data.SRTree.Internal: countParams :: Fix SRTree -> Int
+ Data.SRTree.Internal: data Op
+ Data.SRTree.Internal: deriveByParam :: Int -> Fix SRTree -> Fix SRTree
+ Data.SRTree.Internal: deriveByVar :: Int -> Fix SRTree -> Fix SRTree
+ Data.SRTree.Internal: evalOp :: Floating a => Op -> a -> a -> a
+ Data.SRTree.Internal: floatConstsToParam :: Fix SRTree -> (Fix SRTree, [Double])
+ Data.SRTree.Internal: forwardMode :: (Show a, Num a, Floating a) => Vector a -> Vector Double -> (Double -> a) -> Fix SRTree -> [a]
+ Data.SRTree.Internal: gradParams :: (Show a, Num a, Floating a) => Vector a -> Vector Double -> (Double -> a) -> Fix SRTree -> (a, [a])
+ Data.SRTree.Internal: instance GHC.Base.Functor Data.SRTree.Internal.SRTree
+ Data.SRTree.Internal: instance GHC.Base.Functor Data.SRTree.Internal.Tape
+ Data.SRTree.Internal: instance GHC.Classes.Eq Data.SRTree.Internal.Op
+ Data.SRTree.Internal: instance GHC.Classes.Eq val => GHC.Classes.Eq (Data.SRTree.Internal.SRTree val)
+ Data.SRTree.Internal: instance GHC.Classes.Ord Data.SRTree.Internal.Op
+ Data.SRTree.Internal: instance GHC.Classes.Ord val => GHC.Classes.Ord (Data.SRTree.Internal.SRTree val)
+ Data.SRTree.Internal: instance GHC.Enum.Enum Data.SRTree.Internal.Op
+ Data.SRTree.Internal: instance GHC.Float.Floating (Data.SRTree.Recursion.Fix Data.SRTree.Internal.SRTree)
+ Data.SRTree.Internal: instance GHC.Float.Floating a => GHC.Float.Floating (Data.SRTree.Internal.Tape a)
+ Data.SRTree.Internal: instance GHC.Num.Num (Data.SRTree.Recursion.Fix Data.SRTree.Internal.SRTree)
+ Data.SRTree.Internal: instance GHC.Num.Num a => GHC.Num.Num (Data.SRTree.Internal.Tape a)
+ Data.SRTree.Internal: instance GHC.Read.Read Data.SRTree.Internal.Op
+ Data.SRTree.Internal: instance GHC.Real.Fractional (Data.SRTree.Recursion.Fix Data.SRTree.Internal.SRTree)
+ Data.SRTree.Internal: instance GHC.Real.Fractional a => GHC.Real.Fractional (Data.SRTree.Internal.Tape a)
+ Data.SRTree.Internal: instance GHC.Show.Show Data.SRTree.Internal.Op
+ Data.SRTree.Internal: instance GHC.Show.Show a => GHC.Show.Show (Data.SRTree.Internal.Tape a)
+ Data.SRTree.Internal: instance GHC.Show.Show val => GHC.Show.Show (Data.SRTree.Internal.SRTree val)
+ Data.SRTree.Internal: param :: Int -> Fix SRTree
+ Data.SRTree.Internal: var :: Int -> Fix SRTree
+ Data.SRTree.Print: printLatex :: Fix SRTree -> IO ()
+ Data.SRTree.Print: printPython :: Fix SRTree -> IO ()
+ Data.SRTree.Print: printTikz :: Fix SRTree -> IO ()
+ Data.SRTree.Recursion: (:<) :: a -> f (Cofree f a) -> Cofree f a
+ Data.SRTree.Recursion: ConsF :: a -> b -> ListF a b
+ Data.SRTree.Recursion: Fix :: f (Fix f) -> Fix f
+ Data.SRTree.Recursion: LeafF :: TreeF a b
+ Data.SRTree.Recursion: NilF :: ListF a b
+ Data.SRTree.Recursion: NodeF :: b -> a -> b -> TreeF a b
+ Data.SRTree.Recursion: Op :: f (Free f a) -> Free f a
+ Data.SRTree.Recursion: Ret :: a -> Free f a
+ Data.SRTree.Recursion: StreamF :: a -> b -> StreamF a b
+ Data.SRTree.Recursion: SuccF :: a -> NatF a
+ Data.SRTree.Recursion: ZeroF :: NatF a
+ Data.SRTree.Recursion: [unfix] :: Fix f -> f (Fix f)
+ Data.SRTree.Recursion: accu :: Functor f => (forall x. f x -> p -> f (x, p)) -> (f a -> p -> a) -> Fix f -> p -> a
+ Data.SRTree.Recursion: ana :: Functor f => (a -> f a) -> a -> Fix f
+ Data.SRTree.Recursion: apo :: Functor f => (a -> f (Either (Fix f) a)) -> a -> Fix f
+ Data.SRTree.Recursion: cata :: Functor f => (f a -> a) -> Fix f -> a
+ Data.SRTree.Recursion: cataM :: (Functor f, Monad m) => (forall x. f (m x) -> m (f x)) -> (f a -> m a) -> Fix f -> m a
+ Data.SRTree.Recursion: chrono :: Functor f => (f (Cofree f b) -> b) -> (a -> f (Free f a)) -> a -> b
+ Data.SRTree.Recursion: data Cofree f a
+ Data.SRTree.Recursion: data Free f a
+ Data.SRTree.Recursion: data ListF a b
+ Data.SRTree.Recursion: data NatF a
+ Data.SRTree.Recursion: data StreamF a b
+ Data.SRTree.Recursion: data TreeF a b
+ Data.SRTree.Recursion: extract :: Cofree f a -> a
+ Data.SRTree.Recursion: fromList :: [a] -> Fix (ListF a)
+ Data.SRTree.Recursion: fromNat :: Fix NatF -> Int
+ Data.SRTree.Recursion: futu :: Functor f => (a -> f (Free f a)) -> a -> Fix f
+ Data.SRTree.Recursion: histo :: Functor f => (f (Cofree f a) -> a) -> Fix f -> a
+ Data.SRTree.Recursion: hylo :: Functor f => (f b -> b) -> (a -> f a) -> a -> b
+ Data.SRTree.Recursion: instance GHC.Base.Functor (Data.SRTree.Recursion.ListF a)
+ Data.SRTree.Recursion: instance GHC.Base.Functor (Data.SRTree.Recursion.StreamF a)
+ Data.SRTree.Recursion: instance GHC.Base.Functor (Data.SRTree.Recursion.TreeF a)
+ Data.SRTree.Recursion: instance GHC.Base.Functor Data.SRTree.Recursion.NatF
+ Data.SRTree.Recursion: mutu :: Functor f => (f (a, b) -> a) -> (f (a, b) -> b) -> (Fix f -> a, Fix f -> b)
+ Data.SRTree.Recursion: newtype Fix f
+ Data.SRTree.Recursion: para :: Functor f => (f (Fix f, a) -> a) -> Fix f -> a
+ Data.SRTree.Recursion: stream2list :: StreamF a [a] -> [a]
+ Data.SRTree.Recursion: toList :: Fix (ListF a) -> [a]
+ Data.SRTree.Recursion: toNat :: Int -> Fix NatF
+ Data.SRTree.Recursion: type Algebra f a = f a -> a
+ Data.SRTree.Recursion: type CoAlgebra f a = a -> f a
+ Data.SRTree.Recursion: unOp :: Free f a -> f (Free f a)
- Data.SRTree: Add :: SRTree ix val -> SRTree ix val -> SRTree ix val
+ Data.SRTree: Add :: Op
- Data.SRTree: Const :: val -> SRTree ix val
+ Data.SRTree: Const :: Double -> SRTree val
- Data.SRTree: Div :: SRTree ix val -> SRTree ix val -> SRTree ix val
+ Data.SRTree: Div :: Op
- Data.SRTree: Mul :: SRTree ix val -> SRTree ix val -> SRTree ix val
+ Data.SRTree: Mul :: Op
- Data.SRTree: Param :: ix -> SRTree ix val
+ Data.SRTree: Param :: Int -> SRTree val
- Data.SRTree: Power :: SRTree ix val -> SRTree ix val -> SRTree ix val
+ Data.SRTree: Power :: Op
- Data.SRTree: Sub :: SRTree ix val -> SRTree ix val -> SRTree ix val
+ Data.SRTree: Sub :: Op
- Data.SRTree: Var :: ix -> SRTree ix val
+ Data.SRTree: Var :: Int -> SRTree val
- Data.SRTree: arity :: SRTree ix val -> Int
+ Data.SRTree: arity :: Fix SRTree -> Int
- Data.SRTree: countNodes :: SRTree ix val -> Int
+ Data.SRTree: countNodes :: Fix SRTree -> Int
- Data.SRTree: countOccurrences :: Eq ix => SRTree ix val -> ix -> Int
+ Data.SRTree: countOccurrences :: Int -> Fix SRTree -> Int
- Data.SRTree: countVarNodes :: SRTree ix val -> Int
+ Data.SRTree: countVarNodes :: Fix SRTree -> Int
- Data.SRTree: data SRTree ix val
+ Data.SRTree: data SRTree val
- Data.SRTree: derivative :: (Eq ix, Eq val, Floating val) => Function -> SRTree ix val -> SRTree ix val
+ Data.SRTree: derivative :: Floating a => Function -> a -> a
- Data.SRTree: deriveBy :: (Eq ix, Eq val, Floating val, OptIntPow val) => ix -> SRTree ix val -> SRTree ix val
+ Data.SRTree: deriveBy :: Bool -> Int -> Fix SRTree -> Fix SRTree
- Data.SRTree: evalFun :: Floating val => Function -> val -> val
+ Data.SRTree: evalFun :: Floating a => Function -> a -> a
- Data.SRTree: evalTree :: (Floating val, OptIntPow val) => SRTree ix val -> Reader (ix -> Maybe val) (Maybe val)
+ Data.SRTree: evalTree :: (Num a, Floating a) => Vector a -> Vector Double -> (Double -> a) -> Fix SRTree -> a
- Data.SRTree: getChildren :: SRTree ix val -> [SRTree ix val]
+ Data.SRTree: getChildren :: Fix SRTree -> [Fix SRTree]
- Data.SRTree: relabelParams :: Num ix => SRTree ix val -> SRTree ix val
+ Data.SRTree: relabelParams :: Fix SRTree -> Fix SRTree
- Data.SRTree.Internal: Add :: SRTree ix val -> SRTree ix val -> SRTree ix val
+ Data.SRTree.Internal: Add :: Op
- Data.SRTree.Internal: Const :: val -> SRTree ix val
+ Data.SRTree.Internal: Const :: Double -> SRTree val
- Data.SRTree.Internal: Div :: SRTree ix val -> SRTree ix val -> SRTree ix val
+ Data.SRTree.Internal: Div :: Op
- Data.SRTree.Internal: Mul :: SRTree ix val -> SRTree ix val -> SRTree ix val
+ Data.SRTree.Internal: Mul :: Op
- Data.SRTree.Internal: Param :: ix -> SRTree ix val
+ Data.SRTree.Internal: Param :: Int -> SRTree val
- Data.SRTree.Internal: Power :: SRTree ix val -> SRTree ix val -> SRTree ix val
+ Data.SRTree.Internal: Power :: Op
- Data.SRTree.Internal: Sub :: SRTree ix val -> SRTree ix val -> SRTree ix val
+ Data.SRTree.Internal: Sub :: Op
- Data.SRTree.Internal: Var :: ix -> SRTree ix val
+ Data.SRTree.Internal: Var :: Int -> SRTree val
- Data.SRTree.Internal: arity :: SRTree ix val -> Int
+ Data.SRTree.Internal: arity :: Fix SRTree -> Int
- Data.SRTree.Internal: countNodes :: SRTree ix val -> Int
+ Data.SRTree.Internal: countNodes :: Fix SRTree -> Int
- Data.SRTree.Internal: countOccurrences :: Eq ix => SRTree ix val -> ix -> Int
+ Data.SRTree.Internal: countOccurrences :: Int -> Fix SRTree -> Int
- Data.SRTree.Internal: countVarNodes :: SRTree ix val -> Int
+ Data.SRTree.Internal: countVarNodes :: Fix SRTree -> Int
- Data.SRTree.Internal: data SRTree ix val
+ Data.SRTree.Internal: data SRTree val
- Data.SRTree.Internal: derivative :: (Eq ix, Eq val, Floating val) => Function -> SRTree ix val -> SRTree ix val
+ Data.SRTree.Internal: derivative :: Floating a => Function -> a -> a
- Data.SRTree.Internal: deriveBy :: (Eq ix, Eq val, Floating val, OptIntPow val) => ix -> SRTree ix val -> SRTree ix val
+ Data.SRTree.Internal: deriveBy :: Bool -> Int -> Fix SRTree -> Fix SRTree
- Data.SRTree.Internal: evalFun :: Floating val => Function -> val -> val
+ Data.SRTree.Internal: evalFun :: Floating a => Function -> a -> a
- Data.SRTree.Internal: evalTree :: (Floating val, OptIntPow val) => SRTree ix val -> Reader (ix -> Maybe val) (Maybe val)
+ Data.SRTree.Internal: evalTree :: (Num a, Floating a) => Vector a -> Vector Double -> (Double -> a) -> Fix SRTree -> a
- Data.SRTree.Internal: getChildren :: SRTree ix val -> [SRTree ix val]
+ Data.SRTree.Internal: getChildren :: Fix SRTree -> [Fix SRTree]
- Data.SRTree.Internal: relabelParams :: Num ix => SRTree ix val -> SRTree ix val
+ Data.SRTree.Internal: relabelParams :: Fix SRTree -> Fix SRTree
- Data.SRTree.Print: printExpr :: (Show ix, Show val) => SRTree ix val -> DisplayNodes ix val -> IO ()
+ Data.SRTree.Print: printExpr :: Fix SRTree -> IO ()
- Data.SRTree.Print: showExpr :: (Show ix, Show val) => SRTree ix val -> DisplayNodes ix val -> String
+ Data.SRTree.Print: showExpr :: Fix SRTree -> String
- Data.SRTree.Print: showLatex :: (Show ix, Show val) => SRTree ix val -> String
+ Data.SRTree.Print: showLatex :: Fix SRTree -> String
- Data.SRTree.Print: showPython :: (Show ix, Show val) => SRTree ix val -> String
+ Data.SRTree.Print: showPython :: Fix SRTree -> String
- Data.SRTree.Print: showTikz :: (Show ix, Show val, RealFrac val) => SRTree ix val -> String
+ Data.SRTree.Print: showTikz :: Fix SRTree -> String
- Data.SRTree.Random: P :: [ix] -> (val, val) -> (Int, Int) -> [Function] -> FullParams ix val
+ Data.SRTree.Random: P :: [Int] -> (Double, Double) -> (Int, Int) -> [Function] -> FullParams
- Data.SRTree.Random: data FullParams ix val
+ Data.SRTree.Random: data FullParams
- Data.SRTree.Random: randomConst :: (Ord val, Random val, HasVals p) => RndTree p ix val
+ Data.SRTree.Random: randomConst :: HasVals p => RndTree p
- Data.SRTree.Random: randomFunction :: HasFuns p => RndTree p ix val
+ Data.SRTree.Random: randomFunction :: HasFuns p => RndTree p
- Data.SRTree.Random: randomNode :: (Ord val, Random val, HasEverything p) => RndTree p ix val
+ Data.SRTree.Random: randomNode :: HasEverything p => RndTree p
- Data.SRTree.Random: randomNonTerminal :: (Ord val, Random val, HasEverything p) => RndTree p ix val
+ Data.SRTree.Random: randomNonTerminal :: HasEverything p => RndTree p
- Data.SRTree.Random: randomPow :: (Ord val, Random val, HasExps p) => RndTree p ix val
+ Data.SRTree.Random: randomPow :: HasExps p => RndTree p
- Data.SRTree.Random: randomTree :: (Ord val, Random val, HasEverything p) => Int -> RndTree p ix val
+ Data.SRTree.Random: randomTree :: HasEverything p => Int -> RndTree p
- Data.SRTree.Random: randomTreeBalanced :: (Ord val, Random val, HasEverything p) => Int -> RndTree p ix val
+ Data.SRTree.Random: randomTreeBalanced :: HasEverything p => Int -> RndTree p
- Data.SRTree.Random: randomVar :: HasVars p => RndTree p ix val
+ Data.SRTree.Random: randomVar :: HasVars p => RndTree p
- Data.SRTree.Random: type RndTree p ix val = ReaderT (p ix val) (StateT StdGen IO) (SRTree ix val)
+ Data.SRTree.Random: type RndTree p = ReaderT p (StateT StdGen IO) (Fix SRTree)
Files
- ChangeLog.md +12/−0
- README.md +6/−11
- app/Main.hs +41/−0
- src/Data/SRTree.hs +26/−17
- src/Data/SRTree/Internal.hs +369/−449
- src/Data/SRTree/Print.hs +65/−156
- src/Data/SRTree/Random.hs +36/−42
- src/Data/SRTree/Recursion.hs +89/−0
- srtree.cabal +24/−4
- test/Spec.hs +70/−11
ChangeLog.md view
@@ -1,5 +1,17 @@ # Changelog for srtree +## 1.0.0.0++- Complete refactoring of source+- We now work with the SRTree data type fixed point+- Symbolic derivative by Param or Var+- Forward mode autodiff+- Optimized gradient calculation of parameters if each parameter occurs only once in the expression++## 0.1.3.0++- `countParams` function+ ## 0.1.2.1 - Better bounds for base (compatible with stackage nightly)
README.md view
@@ -2,29 +2,24 @@ `srtree` is a Haskell library with a data structure and supporting functions to manipulate expression trees for symbolic regression. -The tree-like structure is parameterized by the type of the variables indexing and the return value when evaluating the tree. The most common is to index the variables with `Int` starting at $0$ and to return a `Double`. The Functor instance changes the type of the stored/returned values.+The tree-like structure is defined as a fixed-point of an n-ary tree. The variables and parameters of the regression model are indexed as `Int`type and the constant values are `Double`. The tree supports leaf nodes containing a variable, a free parameter, or a constant value; internal nodes that represents binary operators such as the four basic math operations, logarithm with custom base, and the power of two expressions; and unary functions specified by `Function` data type. -This library also defines the `OptInt` class with the operator `^.` that represents the integral power. This is needed to automatically simplify some constructs of the tree and also when using interval arithmetic, that requires a special case of integral power.- The `SRTree` structure has instances for `Num, Fractional, Floating` which allows to create an expression as a valid Haskell expression such as: ```haskell-x = Var 0-y = Var 1-expr = x * 2 + sin(y * pi + x) :: SRTree Int Double+x = var 0)+y = var 1+expr = x * 2 + sin(y * pi + x) :: Fix SRTree ``` -There is also a `Bifunctor` instance that allows to change the type of both parameters, and an `Applicative, Foldable, Traversable` instances. To traverse by the index type, there is a function called `traverseIx`.- ## Other features: -- simplification algorithm (`simplify`)-- derivative w.r.t. a variable (`deriveBy`) and w.r.t. a parameter (`deriveParamBy`)+- derivative w.r.t. a variable (`deriveByVar`) and w.r.t. a parameter (`deriveByParam`) - evaluation (`evalTree`) - relabel free parameters sequentially (`relabelParams`)-- relabel variables couting their occurrence (`relabelOccurrences`, used with interval arithmetic)+- gradient calculation with `forwardMode`, or optimized with `gradParams` if there is only a single occurrence of each parameter (most of the cases). ## TODO:
+ app/Main.hs view
@@ -0,0 +1,41 @@+module Main where++import Data.SRTree+import Data.SRTree.Print+import Data.SRTree.Random+import Data.SRTree.Recursion hiding (fromList)+import Data.Vector (fromList)+import System.Random+import Control.Monad.Reader+import Control.Monad.State+import Criterion.Main++t = 1 + var 0 * (3.1 + param 0 * var 1 + var 0 * param 1) - var 0++xs = fromList [1.0, 2.0]+ps = fromList [0.5, 0.3]++params = P [0,1] (-1.0, 1.0) (-2, 2) [Id ..] ++runRnd g ns p = flip evalStateT g $ traverse (\n -> runReaderT (randomTree n) p) ns+runRndBalance g ns p = flip evalStateT g $ traverse (\n -> runReaderT (randomTreeBalanced n) p) ns++lens :: [Int]+lens = replicate 100 10 <> replicate 100 100 <> replicate 100 1000++g = mkStdGen 42++benchTree f h = do ts <- f g lens params+ pure $ map (h xs ps id) ts++main :: IO ()+main = defaultMain [+ bgroup "unbalanced" + [ bench "forwardMode" $ nfIO (benchTree runRnd forwardMode)+ , bench "grad" $ nfIO (benchTree runRnd gradParams)+ ] ,+ bgroup "balanced" + [ bench "forwardMode" $ nfIO (benchTree runRndBalance forwardMode)+ , bench "grad" $ nfIO (benchTree runRndBalance gradParams)+ ] + ]
src/Data/SRTree.hs view
@@ -13,47 +13,56 @@ module Data.SRTree ( SRTree(..) , Function(..)- , OptIntPow(..)- , traverseIx+ , Op(..)+ , param+ , var , arity , getChildren , countNodes , countVarNodes+ , countConsts+ , countParams , countOccurrences , deriveBy- , deriveParamBy- , simplify+ , deriveByVar+ , deriveByParam , derivative+ , forwardMode+ , gradParams , evalFun+ , evalOp , inverseFunc , evalTree- , evalTreeMap- , evalTreeWithMap- , evalTreeWithVector- , relabelOccurrences , relabelParams+ , constsToParam+ , floatConstsToParam ) where -import Data.SRTree.Internal ( SRTree(..)+import Data.SRTree.Internal + ( SRTree(..) , Function(..)- , OptIntPow(..)- , traverseIx+ , Op(..)+ , param+ , var , arity , getChildren , countNodes , countVarNodes+ , countConsts+ , countParams , countOccurrences , deriveBy- , deriveParamBy- , simplify+ , deriveByVar+ , deriveByParam , derivative+ , forwardMode+ , gradParams , evalFun+ , evalOp , inverseFunc , evalTree- , evalTreeMap- , evalTreeWithMap- , evalTreeWithVector- , relabelOccurrences , relabelParams+ , constsToParam+ , floatConstsToParam )
src/Data/SRTree/Internal.hs view
@@ -1,5 +1,6 @@ {-# language FlexibleInstances, DeriveFunctor #-} {-# language ScopedTypeVariables #-}+{-# language RankNTypes #-} ----------------------------------------------------------------------------- -- | -- Module : Data.SRTree.Internal @@ -13,463 +14,291 @@ -- ----------------------------------------------------------------------------- -module Data.SRTree.Internal +module Data.SRTree.Internal ( SRTree(..) , Function(..)- , OptIntPow(..)- , traverseIx+ , Op(..)+ , param+ , var , arity , getChildren , countNodes , countVarNodes+ , countConsts+ , countParams , countOccurrences , deriveBy- , deriveParamBy- , simplify+ , deriveByVar+ , deriveByParam , derivative+ , forwardMode+ , gradParams , evalFun+ , evalOp , inverseFunc , evalTree- , evalTreeMap- , evalTreeWithMap- , evalTreeWithVector- , relabelOccurrences , relabelParams+ , constsToParam+ , floatConstsToParam ) where -import Data.Bifunctor+import Data.SRTree.Recursion ( Fix(Fix), cata, mutu, cataM ) -import Data.Map.Strict (Map(..), (!), (!?), insert, fromList)-import qualified Data.Map.Strict as M import qualified Data.Vector as V+import Data.Vector ((!)) import Control.Monad.State-import Control.Monad.Reader-import Control.Applicative hiding (Const) +import Debug.Trace (trace)+ -- | Tree structure to be used with Symbolic Regression algorithms.--- This structure is parametrized by the indexing type to retrieve the values--- of a variable and the type of the output value.-data SRTree ix val = - Empty - | Var ix- | Const val- | Param ix- | Fun Function (SRTree ix val)- | Pow (SRTree ix val) Int- | SRTree ix val `Add` SRTree ix val- | SRTree ix val `Sub` SRTree ix val- | SRTree ix val `Mul` SRTree ix val- | SRTree ix val `Div` SRTree ix val- | SRTree ix val `Power` SRTree ix val- | SRTree ix val `LogBase` SRTree ix val- deriving (Show, Eq, Ord, Functor)+-- This structure is a fixed point of a n-ary tree. +data SRTree val =+ Var Int -- ^ index of the variables+ | Param Int -- ^ index of the parameter+ | Const Double -- ^ constant value, can be converted to a parameter+ | Uni Function val -- ^ univariate function+ | Bin Op val val -- ^ binary operator+ deriving (Show, Eq, Ord, Functor) --- | Functions that can be applied to a subtree.-data Function = - Id +-- | Supported operators+data Op = Add | Sub | Mul | Div | Power+ deriving (Show, Read, Eq, Ord, Enum)++-- | Supported functions+data Function =+ Id | Abs- | Sin - | Cos - | Tan + | Sin+ | Cos+ | Tan | Sinh | Cosh- | Tanh - | ASin - | ACos + | Tanh+ | ASin+ | ACos | ATan | ASinh- | ACosh - | ATanh - | Sqrt - | Cbrt - | Square - | Log - | Exp + | ACosh+ | ATanh+ | Sqrt+ | Cbrt+ | Square+ | Log+ | Exp deriving (Show, Read, Eq, Ord, Enum) --- | A class for optimized `(^^)` operators for specific types.--- This was created because the integer power operator for--- interval arithmetic must be aware of the dependency problem,--- thus the default `(^)` doesn't work.-class OptIntPow a where- (^.) :: a -> Int -> a- infix 8 ^.- -instance OptIntPow Double where- (^.) = (^^)- {-# INLINE (^.) #-}-instance OptIntPow Float where- (^.) = (^^)- {-# INLINE (^.) #-}- - -instance (Eq ix, Eq val, Num val, OptIntPow val) => OptIntPow (SRTree ix val) where- t ^. 0 = 1- t ^. 1 = t- (Const c) ^. k = Const $ c ^. k - t ^. k = Pow t k- {-# INLINE (^.) #-}- -instance (Eq ix, Eq val, Num val) => Num (SRTree ix val) where- 0 + r = r- l + 0 = l- (Const c1) + (Const c2) = Const $ c1 + c2- l + r = Add l r+-- | create a tree with a single node representing a variable+var :: Int -> Fix SRTree+var ix = Fix (Var ix)++-- | create a tree with a single node representing a parameter+param :: Int -> Fix SRTree+param ix = Fix (Param ix)++instance Num (Fix SRTree) where+ Fix (Const 0) + r = r+ l + Fix (Const 0) = l+ Fix (Const c1) + Fix (Const c2) = Fix . Const $ c1 + c2+ l + r = Fix $ Bin Add l r {-# INLINE (+) #-} - 0 - r = (-1) * r- l - 0 = l- (Const c1) - (Const c2) = Const $ c1 - c2- l - r = Sub l r+ l - Fix (Const 0) = l+ Fix (Const 0) - r = negate r+ Fix (Const c1) - Fix (Const c2) = Fix . Const $ c1 - c2+ l - r = Fix $ Bin Sub l r {-# INLINE (-) #-} - 0 * r = 0- l * 0 = 0- 1 * r = r- l * 1 = l - (Const c1) * (Const c2) = Const $ c1 * c2- l * r = Mul l r+ Fix (Const 0) * _ = Fix (Const 0)+ _ * Fix (Const 0) = Fix (Const 0)+ Fix (Const 1) * r = r+ l * Fix (Const 1) = l+ Fix (Const c1) * Fix (Const c2) = Fix . Const $ c1 * c2+ l * r = Fix $ Bin Mul l r {-# INLINE (*) #-}- - abs = Fun Abs++ abs = Fix . Uni Abs {-# INLINE abs #-}- - negate (Const x) = Const (negate x)- negate t = Const (-1) * t++ negate (Fix (Const x)) = Fix $ Const (negate x)+ negate t = Fix (Const (-1)) * t {-# INLINE negate #-}- + signum t = case t of- Const x -> Const $ signum x- _ -> Const 0- fromInteger x = Const (fromInteger x)+ Fix (Const x) -> Fix . Const $ signum x+ _ -> Fix (Const 0)+ fromInteger x = Fix $ Const (fromInteger x) {-# INLINE fromInteger #-} -instance (Eq ix, Eq val, Fractional val) => Fractional (SRTree ix val) where- 0 / r = 0- l / 1 = l- (Const c1) / (Const c2) = Const $ c1/c2- l / r = Div l r+instance Fractional (Fix SRTree) where+ l / Fix (Const 1) = l+ Fix (Const c1) / Fix (Const c2) = Fix . Const $ c1/c2+ l / r = Fix $ Bin Div l r {-# INLINE (/) #-}- - fromRational = Const . fromRational++ fromRational = Fix . Const . fromRational {-# INLINE fromRational #-}- -instance (Eq ix, Eq val, Floating val) => Floating (SRTree ix val) where - pi = Const pi++instance Floating (Fix SRTree) where+ pi = Fix $ Const pi {-# INLINE pi #-}- exp = evalToConst . Fun Exp+ exp = Fix . Uni Exp {-# INLINE exp #-}- log = evalToConst . Fun Log+ log = Fix . Uni Log {-# INLINE log #-}- sqrt = evalToConst . Fun Sqrt+ sqrt = Fix . Uni Sqrt {-# INLINE sqrt #-}- sin = evalToConst . Fun Sin+ sin = Fix . Uni Sin {-# INLINE sin #-}- cos = evalToConst . Fun Cos+ cos = Fix . Uni Cos {-# INLINE cos #-}- tan = evalToConst . Fun Tan+ tan = Fix . Uni Tan {-# INLINE tan #-}- asin = evalToConst . Fun ASin+ asin = Fix . Uni ASin {-# INLINE asin #-}- acos = evalToConst . Fun ACos+ acos = Fix . Uni ACos {-# INLINE acos #-}- atan = evalToConst . Fun ATan+ atan = Fix . Uni ATan {-# INLINE atan #-}- sinh = evalToConst . Fun Sinh+ sinh = Fix . Uni Sinh {-# INLINE sinh #-}- cosh = evalToConst . Fun Cosh+ cosh = Fix . Uni Cosh {-# INLINE cosh #-}- tanh = evalToConst . Fun Tanh+ tanh = Fix . Uni Tanh {-# INLINE tanh #-}- asinh = evalToConst . Fun ASinh+ asinh = Fix . Uni ASinh {-# INLINE asinh #-}- acosh = evalToConst . Fun ACosh+ acosh = Fix . Uni ACosh {-# INLINE acosh #-}- atanh = evalToConst . Fun ATanh+ atanh = Fix . Uni ATanh {-# INLINE atanh #-} - 0 ** r = 0- 1 ** r = 1- l ** 0 = 1- l ** 1 = l- l ** r = evalToConst $ Power l r+ l ** Fix (Const 1) = l+ l ** Fix (Const 0) = Fix (Const 1)+ l ** r = Fix $ Bin Power l r {-# INLINE (**) #-} - logBase 1 r = 0- logBase l r = evalToConst $ LogBase l r+ logBase l (Fix (Const 1)) = Fix (Const 0)+ logBase l r = log l / log r {-# INLINE logBase #-} -instance Bifunctor SRTree where- first f Empty = Empty- first f (Var ix) = Var $ f ix- first f (Param ix) = Param $ f ix- first f (Fun g t) = Fun g $ first f t- first f (Pow t k) = Pow (first f t) k- first f (Add l r) = Add (first f l) (first f r)- first f (Sub l r) = Sub (first f l) (first f r)- first f (Mul l r) = Mul (first f l) (first f r)- first f (Div l r) = Div (first f l) (first f r)- first f (Power l r) = Power (first f l) (first f r)- first f (LogBase l r) = LogBase (first f l) (first f r)- {-# INLINE first #-}- - second = fmap- {-# INLINE second #-}--instance Applicative (SRTree ix) where- pure = Const-- Empty <*> t = Empty- Var ix <*> t = Var ix- Param ix <*> t = Param ix- Const f <*> t = fmap f t- Fun g tf <*> t = Fun g $ tf <*> t- Pow tf k <*> t = Pow (tf <*> t) k- Add lf rf <*> t = Add (lf <*> t) (rf <*> t)- Sub lf rf <*> t = Sub (lf <*> t) (rf <*> t)- Mul lf rf <*> t = Mul (lf <*> t) (rf <*> t)- Div lf rf <*> t = Div (lf <*> t) (rf <*> t)- Power lf rf <*> t = Power (lf <*> t) (rf <*> t)- LogBase lf rf <*> t = LogBase (lf <*> t) (rf <*> t)- -instance Foldable (SRTree ix) where- foldMap f Empty = mempty- foldMap f (Var ix) = mempty- foldMap f (Param ix) = mempty- foldMap f (Const c) = f c- foldMap f t = mconcat $ map (foldMap f) $ getChildren t--instance Traversable (SRTree ix) where- traverse mf Empty = pure Empty- traverse mf (Var ix) = pure $ Var ix- traverse mf (Param ix) = pure $ Param ix- traverse mf (Const c) = Const <$> mf c- traverse mf (Fun g t) = Fun g <$> traverse mf t- traverse mf (Pow t k) = (`Pow` k) <$> traverse mf t- traverse mf (Add l r) = Add <$> traverse mf l <*> traverse mf r- traverse mf (Sub l r) = Sub <$> traverse mf l <*> traverse mf r- traverse mf (Mul l r) = Mul <$> traverse mf l <*> traverse mf r- traverse mf (Div l r) = Div <$> traverse mf l <*> traverse mf r- traverse mf (Power l r) = Power <$> traverse mf l <*> traverse mf r- traverse mf (LogBase l r) = LogBase <$> traverse mf l <*> traverse mf r---- | Same as `traverse` but for the first type parameter.-traverseIx :: Applicative f => (ixa -> f ixb) -> SRTree ixa val -> f (SRTree ixb val)-traverseIx mf Empty = pure Empty-traverseIx mf (Var ix) = Var <$> mf ix-traverseIx mf (Param ix) = Param <$> mf ix-traverseIx mf (Const c) = pure $ Const c-traverseIx mf (Fun g t) = Fun g <$> traverseIx mf t-traverseIx mf (Pow t k) = (`Pow` k) <$> traverseIx mf t-traverseIx mf (Add l r) = Add <$> traverseIx mf l <*> traverseIx mf r-traverseIx mf (Sub l r) = Sub <$> traverseIx mf l <*> traverseIx mf r-traverseIx mf (Mul l r) = Mul <$> traverseIx mf l <*> traverseIx mf r-traverseIx mf (Div l r) = Div <$> traverseIx mf l <*> traverseIx mf r-traverseIx mf (Power l r) = Power <$> traverseIx mf l <*> traverseIx mf r-traverseIx mf (LogBase l r) = LogBase <$> traverseIx mf l <*> traverseIx mf r-{-# INLINE traverseIx #-}- -- | Arity of the current node-arity :: SRTree ix val -> Int-arity Empty = 0-arity (Var _) = 0-arity (Param _) = 0-arity (Const _) = 0-arity (Fun _ _) = 1-arity (Pow _ _) = 1-arity _ = 2+arity :: Fix SRTree -> Int+arity = cata alg+ where+ alg Var {} = 0+ alg Param {} = 0+ alg Const {} = 0+ alg Uni {} = 1+ alg Bin {} = 2 {-# INLINE arity #-} -- | Get the children of a node. Returns an empty list in case of a leaf node.-getChildren :: SRTree ix val -> [SRTree ix val]-getChildren Empty = []-getChildren (Var _) = []-getChildren (Param _) = []-getChildren (Const _) = []-getChildren (Fun _ t) = [t]-getChildren (Pow t _) = [t]-getChildren (Add l r) = [l, r]-getChildren (Sub l r) = [l, r]-getChildren (Mul l r) = [l, r]-getChildren (Div l r) = [l, r]-getChildren (Power l r) = [l, r]-getChildren (LogBase l r) = [l, r]+getChildren :: Fix SRTree -> [Fix SRTree]+getChildren (Fix (Var {})) = []+getChildren (Fix (Param {})) = []+getChildren (Fix (Const {})) = []+getChildren (Fix (Uni _ t)) = [t]+getChildren (Fix (Bin _ l r)) = [l, r] {-# INLINE getChildren #-} --- Support function to simplify operations applied to const subtrees.-evalToConst :: Floating val => SRTree ix val -> SRTree ix val -evalToConst (Fun g (Const c)) = Const $ evalFun g c-evalToConst (Power (Const c1) (Const c2)) = Const $ c1**c2-evalToConst (LogBase (Const c1) (Const c2)) = Const $ logBase c1 c2-evalToConst t = t-{-# INLINE evalToConst #-}---- Support function to sum the types of nodes specified by `f`.-sumCounts :: (SRTree ix val -> Int) -> Int -> SRTree ix val -> Int-sumCounts f val = foldr (\c v -> f c + v) val . getChildren-{-# INLINE sumCounts #-}- -- | Count the number of nodes in a tree.-countNodes :: SRTree ix val -> Int-countNodes Empty = 0-countNodes t = sumCounts countNodes 1 t+countNodes :: Fix SRTree -> Int+countNodes = cata alg+ where+ alg Var {} = 1+ alg Param {} = 1+ alg Const {} = 1+ alg (Uni _ t) = 1 + t+ alg (Bin _ l r) = 1 + l + r {-# INLINE countNodes #-} -- | Count the number of `Var` nodes-countVarNodes :: SRTree ix val -> Int-countVarNodes (Var _) = 1-countVarNodes t = sumCounts countVarNodes 0 t+countVarNodes :: Fix SRTree -> Int+countVarNodes = cata alg+ where+ alg Var {} = 1+ alg Param {} = 0+ alg Const {} = 0+ alg (Uni _ t) = 0 + t+ alg (Bin _ l r) = 0 + l + r {-# INLINE countVarNodes #-} --- | Count the occurrences of variable indexed as `ix`-countOccurrences :: Eq ix => SRTree ix val -> ix -> Int-countOccurrences (Var ix) iy = if ix==iy then 1 else 0-countOccurrences t iy = sumCounts (`countOccurrences` iy) 0 t-{-# INLINE countOccurrences #-}---- | Creates an `SRTree` representing the partial derivative of the input by the variable indexed by `ix`.-deriveBy :: (Eq ix, Eq val, Floating val, OptIntPow val) => ix -> SRTree ix val -> SRTree ix val-deriveBy _ Empty = Empty-deriveBy dx (Var ix)- | dx == ix = 1- | otherwise = 0-deriveBy dx (Param ix) = 0-deriveBy dx (Const val) = 0-deriveBy dx (Fun g t) =- case deriveBy dx t of- 0 -> 0- 1 -> derivative g t- t' -> derivative g t * t'-deriveBy dx (Pow t 0) = 0 -deriveBy dx (Pow t 1) = deriveBy dx t-deriveBy dx (Pow t k) = - case deriveBy dx t of- 0 -> 0- Const val -> Const (val * fromIntegral k) * (t ^. (k-1))- t' -> fromIntegral k * (t ^. (k-1)) * t'-deriveBy dx (Add l r) = deriveBy dx l + deriveBy dx r-deriveBy dx (Sub l r) = deriveBy dx l - deriveBy dx r-deriveBy dx (Mul l r) = deriveBy dx l * r + l * deriveBy dx r-deriveBy dx (Div l r) = (deriveBy dx l * r - l * deriveBy dx r) / r ^. 2-deriveBy dx (Power l r) = l ** (r-1) * (r * deriveBy dx l + l * log l * deriveBy dx r)-deriveBy dx (LogBase l r) = deriveBy dx (log l / log r)-{-# INLINE deriveBy #-}+-- | Count the number of `Param` nodes+countParams :: Fix SRTree -> Int+countParams = cata alg+ where+ alg Var {} = 0+ alg Param {} = 1+ alg Const {} = 0+ alg (Uni _ t) = 0 + t+ alg (Bin _ l r) = 0 + l + r+{-# INLINE countParams #-} --- | Creates an `SRTree` representing the partial derivative of the input by the parameter indexed by `ix`.-deriveParamBy :: (Eq ix, Eq val, Floating val, OptIntPow val) => ix -> SRTree ix val -> SRTree ix val-deriveParamBy _ Empty = Empty-deriveParamBy dx (Var ix) = 0-deriveParamBy dx (Param ix)- | dx == ix = 1- | otherwise = 0-deriveParamBy dx (Const val) = 0-deriveParamBy dx (Fun g t) =- case deriveParamBy dx t of- 0 -> 0- 1 -> derivative g t- t' -> derivative g t * t'-deriveParamBy dx (Pow t 0) = 0 -deriveParamBy dx (Pow t 1) = deriveParamBy dx t-deriveParamBy dx (Pow t k) = - case deriveParamBy dx t of- 0 -> 0- Const val -> Const (val * fromIntegral k) * (t ^. (k-1))- t' -> fromIntegral k * (t ^. (k-1)) * t'-deriveParamBy dx (Add l r) = deriveParamBy dx l + deriveParamBy dx r-deriveParamBy dx (Sub l r) = deriveParamBy dx l - deriveParamBy dx r-deriveParamBy dx (Mul l r) = deriveParamBy dx l * r + l * deriveParamBy dx r-deriveParamBy dx (Div l r) = (deriveParamBy dx l * r - l * deriveParamBy dx r) / r ^. 2-deriveParamBy dx (Power l r) = l ** (r-1) * (r * deriveParamBy dx l + l * log l * deriveParamBy dx r)-deriveParamBy dx (LogBase l r) = deriveParamBy dx (log l / log r)-{-# INLINE deriveParamBy #-}+-- | Count the number of const nodes+countConsts :: Fix SRTree -> Int+countConsts = cata alg+ where+ alg Var {} = 0+ alg Param {} = 0+ alg Const {} = 1+ alg (Uni _ t) = 0 + t+ alg (Bin _ l r) = 0 + l + r+{-# INLINE countConsts #-} --- | Simplifies the `SRTree`.-simplify :: (Eq ix, Eq val, Floating val, OptIntPow val) => SRTree ix val -> SRTree ix val-simplify (Fun g t) = evalToConst . Fun g $ simplify t-simplify (Pow t 0) = 1 -simplify (Pow t 1) = simplify t-simplify (Pow t k) =- case simplify t of- Const c -> Const $ c ^. k- t' -> Pow t' k- -simplify (Add l r)- | l' == r' = 2 * l' - | otherwise = l' + r' - where - l' = simplify l - r' = simplify r-simplify (Sub l r)- | l' == r' = 0- | otherwise = l' - r' - where - l' = simplify l - r' = simplify r-simplify (Mul l r)- | l' == r' = Pow l' 2- | otherwise = l' * r' - where - l' = simplify l - r' = simplify r-simplify (Div l r)- | l' == r' = 1- | otherwise = l' / r' - where - l' = simplify l - r' = simplify r+-- | Count the occurrences of variable indexed as `ix`+countOccurrences :: Int -> Fix SRTree -> Int+countOccurrences ix = sum . cata alg+ where+ alg (Var iy) = [1 | ix == iy]+ alg Param {} = []+ alg Const {} = []+ alg (Uni _ t) = t+ alg (Bin _ l r) = l <> r+{-# INLINE countOccurrences #-} -simplify (Power l r) = simplify l ** simplify r-simplify (LogBase l r) = logBase (simplify l) (simplify r)-simplify t = t-{-# INLINE simplify #-}+-- | Evaluates the tree given a vector of variable values, a vector of parameter values and a function that takes a Double and change to whatever type the variables have. This is useful when working with datasets of many values per variables.+evalTree :: (Num a, Floating a) => V.Vector a -> V.Vector Double -> (Double -> a) -> Fix SRTree -> a+evalTree xss params f = cata alg+ where+ alg (Var ix) = xss ! ix+ alg (Param ix) = f $ params ! ix+ alg (Const c) = f c+ alg (Uni g t) = evalFun g t+ alg (Bin op l r) = evalOp op l r+{-# INLINE evalTree #-} --- | Derivative of a Function-derivative :: (Eq ix, Eq val, Floating val) => Function -> SRTree ix val -> SRTree ix val-derivative Id = const 1-derivative Abs = \x -> x / abs x-derivative Sin = cos-derivative Cos = negate.sin-derivative Tan = recip . (**2.0) . cos-derivative Sinh = cosh-derivative Cosh = sinh-derivative Tanh = (1-) . (**2.0) . tanh-derivative ASin = recip . sqrt . (1-) . (^2)-derivative ACos = negate . recip . sqrt . (1-) . (^2)-derivative ATan = recip . (1+) . (^2)-derivative ASinh = recip . sqrt . (1+) . (^2)-derivative ACosh = \x -> 1 / (sqrt (x-1) * sqrt (x+1))-derivative ATanh = recip . (1-) . (^2)-derivative Sqrt = recip . (2*) . sqrt-derivative Cbrt = recip . (3*) . cbrt . (^2)-derivative Square = (2*)-derivative Exp = exp-derivative Log = recip-{-# INLINE derivative #-}+evalOp :: Floating a => Op -> a -> a -> a+evalOp Add = (+)+evalOp Sub = (-)+evalOp Mul = (*)+evalOp Div = (/)+evalOp Power = (**)+{-# INLINE evalOp #-} --- | Evaluates a function.-evalFun :: Floating val => Function -> val -> val-evalFun Id = id-evalFun Abs = abs-evalFun Sin = sin-evalFun Cos = cos-evalFun Tan = tan-evalFun Sinh = sinh-evalFun Cosh = cosh-evalFun Tanh = tanh-evalFun ASin = asin-evalFun ACos = acos-evalFun ATan = atan-evalFun ASinh = asinh-evalFun ACosh = acosh-evalFun ATanh = atanh-evalFun Sqrt = sqrt-evalFun Cbrt = cbrt-evalFun Square = (^2)-evalFun Exp = exp-evalFun Log = log+evalFun :: Floating a => Function -> a -> a+evalFun Id = id+evalFun Abs = abs+evalFun Sin = sin+evalFun Cos = cos+evalFun Tan = tan+evalFun Sinh = sinh+evalFun Cosh = cosh+evalFun Tanh = tanh+evalFun ASin = asin+evalFun ACos = acos+evalFun ATan = atan+evalFun ASinh = asinh+evalFun ACosh = acosh+evalFun ATanh = atanh+evalFun Sqrt = sqrt+evalFun Cbrt = cbrt+evalFun Square = (^2)+evalFun Log = log+evalFun Exp = exp {-# INLINE evalFun #-} +-- | Cubic root cbrt :: Floating val => val -> val cbrt x = signum x * abs x ** (1/3) {-# INLINE cbrt #-}@@ -492,85 +321,176 @@ inverseFunc x = error $ show x ++ " has no support for inverse function" {-# INLINE inverseFunc #-} --- | Evaluates a tree with the variables stored in a `Reader` monad.-evalTree :: (Floating val, OptIntPow val) => SRTree ix val -> Reader (ix -> Maybe val) (Maybe val)-evalTree Empty = pure Nothing-evalTree (Const c) = pure $ Just c-evalTree (Var ix) = askAbout ix-evalTree (Param ix) = pure $ Just 1.0 -- TODO: askAbout paramIx-evalTree (Fun f t) = evalFun f <$$> evalTree t-evalTree (Pow t k) = (^. k) <$$> evalTree t-evalTree (Add l r) = (+) <$*> evalTree l <*> evalTree r-evalTree (Sub l r) = (-) <$*> evalTree l <*> evalTree r-evalTree (Mul l r) = (*) <$*> evalTree l <*> evalTree r-evalTree (Div l r) = (/) <$*> evalTree l <*> evalTree r-evalTree (Power l r) = (**) <$*> evalTree l <*> evalTree r-evalTree (LogBase l r) = logBase <$*> evalTree l <*> evalTree r+-- | Creates the symbolic partial derivative of a tree by variable `dx` (if `p` is `False`)+-- or parameter `dx` (if `p` is `True`).+deriveBy :: Bool -> Int -> Fix SRTree -> Fix SRTree+deriveBy p dx = fst (mutu alg1 alg2)+ where+ alg1 (Var ix) = if not p && ix == dx then 1 else 0+ alg1 (Param ix) = if p && ix == dx then 1 else 0+ alg1 (Const _) = 0+ alg1 (Uni f t) = derivative f (snd t) * fst t+ alg1 (Bin Add l r) = fst l + fst r+ alg1 (Bin Sub l r) = fst l - fst r+ alg1 (Bin Mul l r) = fst l * snd r + snd l * fst r+ alg1 (Bin Div l r) = (fst l * snd r - snd l * fst r) / snd r ** 2+ alg1 (Bin Power l r) = snd l ** (snd r - 1) * (snd r * fst l + snd l * log (snd l) * fst r) --- | Evaluates a tree with the variables stored in a `Reader` monad while mapping the constant --- values to a different type.-evalTreeMap :: (Floating v1, OptIntPow v1, Floating v2, OptIntPow v2) => (v1 -> v2) -> SRTree ix v1 -> Reader (ix -> Maybe v2) (Maybe v2)-evalTreeMap f Empty = pure Nothing-evalTreeMap f (Const c) = pure $ Just $ f c-evalTreeMap f (Var ix) = askAbout ix-evalTreeMap f (Param ix) = pure $ Just $ f 1.0 -- TODO: askAbout paramIx-evalTreeMap f (Fun g t) = evalFun g <$$> evalTreeMap f t-evalTreeMap f (Pow t k) = (^. k) <$$> evalTreeMap f t-evalTreeMap f (Add l r) = (+) <$*> evalTreeMap f l <*> evalTreeMap f r-evalTreeMap f (Sub l r) = (-) <$*> evalTreeMap f l <*> evalTreeMap f r-evalTreeMap f (Mul l r) = (*) <$*> evalTreeMap f l <*> evalTreeMap f r-evalTreeMap f (Div l r) = (/) <$*> evalTreeMap f l <*> evalTreeMap f r-evalTreeMap f (Power l r) = (**) <$*> evalTreeMap f l <*> evalTreeMap f r-evalTreeMap f (LogBase l r) = logBase <$*> evalTreeMap f l <*> evalTreeMap f r+ alg2 (Var ix) = var ix+ alg2 (Param ix) = param ix+ alg2 (Const c) = Fix (Const c)+ alg2 (Uni f t) = Fix (Uni f $ snd t)+ alg2 (Bin f l r) = Fix (Bin f (snd l) (snd r)) --- lift functions inside nested applicatives.-(<$$>) :: (Applicative f, Applicative g) => (a -> b) -> f (g a) -> f (g b)-(<$$>) = fmap . fmap-{-# INLINE (<$$>) #-}-(<$*>) :: (Applicative f, Applicative g) => (a -> b -> c) -> f (g a) -> f (g b -> g c)-op <$*> m = liftA2 op <$> m-{-# INLINE (<$*>) #-}+newtype Tape a = Tape { untape :: [a] } deriving (Show, Functor) --- applies the argument `x` in the function carried by the `Reader` monad.-askAbout :: x -> Reader (x -> a) a-askAbout x = asks ($ x)-{-# INLINE askAbout #-}+instance Num a => Num (Tape a) where+ (Tape x) + (Tape y) = Tape $ zipWith (+) x y+ (Tape x) - (Tape y) = Tape $ zipWith (-) x y+ (Tape x) * (Tape y) = Tape $ zipWith (*) x y+ abs (Tape x) = Tape (map abs x)+ signum (Tape x) = Tape (map signum x)+ fromInteger x = Tape [fromInteger x]+ negate (Tape x) = Tape $ map (*(-1)) x+instance Floating a => Floating (Tape a) where+ pi = Tape [pi]+ exp (Tape x) = Tape (map exp x)+ log (Tape x) = Tape (map log x)+ sqrt (Tape x) = Tape (map sqrt x)+ sin (Tape x) = Tape (map sin x)+ cos (Tape x) = Tape (map cos x)+ tan (Tape x) = Tape (map tan x)+ asin (Tape x) = Tape (map asin x)+ acos (Tape x) = Tape (map acos x)+ atan (Tape x) = Tape (map atan x)+ sinh (Tape x) = Tape (map sinh x)+ cosh (Tape x) = Tape (map cosh x)+ tanh (Tape x) = Tape (map tanh x)+ asinh (Tape x) = Tape (map asinh x)+ acosh (Tape x) = Tape (map acosh x)+ atanh (Tape x) = Tape (map atanh x)+ (Tape x) ** (Tape y) = Tape $ zipWith (**) x y+instance Fractional a => Fractional (Tape a) where+ fromRational x = Tape [fromRational x]+ (Tape x) / (Tape y) = Tape $ zipWith (/) x y+ recip (Tape x) = Tape $ map recip x --- | Example of using `evalTree` with a Map.-evalTreeWithMap :: (Ord ix, Floating val, OptIntPow val) => SRTree ix val -> Map ix val -> Maybe val-evalTreeWithMap t m = runReader (evalTree t) (m !?)-{-# INLINE evalTreeWithMap #-}+-- | Calculates the numerical derivative of a tree using forward mode+-- provided a vector of variable values `xss`, a vector of parameter values `theta` and+-- a function that changes a Double value to the type of the variable values.+forwardMode :: (Show a, Num a, Floating a) => V.Vector a -> V.Vector Double -> (Double -> a) -> Fix SRTree -> [a]+forwardMode xss theta f = untape . fst (mutu alg1 alg2)+ where+ n = V.length theta+ repMat v = Tape $ replicate n v+ zeroes = repMat $ f 0+ twos = repMat $ f 2+ tapeXs = [repMat $ xss ! ix | ix <- [0 .. V.length xss - 1]]+ tapeTheta = [repMat $ f (theta ! ix) | ix <- [0 .. n - 1]]+ paramVec = [ Tape [if ix==iy then f 1 else f 0 | iy <- [0 .. n-1]] | ix <- [0 .. n-1] ] --- | Example of using `evalTree` with a Vector.-evalTreeWithVector :: (Floating val, OptIntPow val) => SRTree Int val -> V.Vector val -> Maybe val-evalTreeWithVector t v = runReader (evalTree t) (v V.!?)-{-# INLINE evalTreeWithVector #-}+ alg1 (Var ix) = zeroes+ alg1 (Param ix) = paramVec !! ix+ alg1 (Const _) = zeroes+ alg1 (Uni f t) = derivative f (snd t) * fst t+ alg1 (Bin Add l r) = fst l + fst r+ alg1 (Bin Sub l r) = fst l - fst r+ alg1 (Bin Mul l r) = (fst l * snd r) + (snd l * fst r)+ alg1 (Bin Div l r) = ((fst l * snd r) - (snd l * fst r)) / snd r ** twos+ alg1 (Bin Power l r) = snd l ** (snd r - 1) * ((snd r * fst l) + (snd l * log (snd l) * fst r)) --- | Relabel occurences of a var into a tuple (ix, Int).-relabelOccurrences :: forall ix val . Ord ix => SRTree ix val -> SRTree (ix, Int) val-relabelOccurrences t = traverseIx updVar t `evalState` M.empty + alg2 (Var ix) = tapeXs !! ix+ alg2 (Param ix) = tapeTheta !! ix+ alg2 (Const c) = repMat $ f c+ alg2 (Uni g t) = fmap (evalFun g) (snd t)+ alg2 (Bin op l r) = evalOp op (snd l) (snd r)++-- | The function `gradParams` calculates the numerical gradient of the tree and evaluates the tree at the same time. It assumes that each parameter has a unique occurrence in the expression. This should be significantly faster than `forwardMode`.+gradParams :: (Show a, Num a, Floating a) => V.Vector a -> V.Vector Double -> (Double -> a) -> Fix SRTree -> (a, [a])+gradParams xss theta f = cata alg where- updVar :: ix -> State (Map ix Int) (ix, Int)- updVar ix = do- s <- get- case s !? ix of- Nothing -> do put $ insert ix 0 s- pure (ix, 0)- Just c -> do put $ insert ix (c+1) s- pure (ix, c+1)+ n = V.length theta --- | Relabel the parameters sequentially starting from 0-relabelParams :: Num ix => SRTree ix val -> SRTree ix val-relabelParams t = (toState t) `evalState` 0+ alg (Var ix) = (xss ! ix, [])+ alg (Param ix) = (f $ theta ! ix, [1])+ alg (Const c) = (f c, [])+ alg (Uni f (v, gs)) = let v' = evalFun f v in (v', map (* derivative f v) gs)+ alg (Bin Add (v1, l) (v2, r)) = (v1+v2, l ++ r)+ alg (Bin Sub (v1, l) (v2, r)) = (v1-v2, l ++ map negate r)+ alg (Bin Mul (v1, l) (v2, r)) = (v1*v2, map (*v2) l ++ map (*v1) r)+ alg (Bin Div (v1, l) (v2, r)) = (v1/v2, map (/v2) l ++ map ((/v2^2) . (*v1) . negate) r)+ alg (Bin Power (v1, l) (v2, r)) = (v1 ** v2, map (* (v1 ** (v2 - 1))) (map (*v2) l ++ map ((*v1).(* log v1)) r))+++derivative :: Floating a => Function -> a -> a+derivative Id = const 1+derivative Abs = \x -> x / abs x+derivative Sin = cos+derivative Cos = negate.sin+derivative Tan = recip . (**2.0) . cos+derivative Sinh = cosh+derivative Cosh = sinh+derivative Tanh = (1-) . (**2.0) . tanh+derivative ASin = recip . sqrt . (1-) . (^2)+derivative ACos = negate . recip . sqrt . (1-) . (^2)+derivative ATan = recip . (1+) . (^2)+derivative ASinh = recip . sqrt . (1+) . (^2)+derivative ACosh = \x -> 1 / (sqrt (x-1) * sqrt (x+1))+derivative ATanh = recip . (1-) . (^2)+derivative Sqrt = recip . (2*) . sqrt+derivative Cbrt = recip . (3*) . cbrt . (^2)+derivative Square = (2*)+derivative Exp = exp+derivative Log = recip+{-# INLINE derivative #-}++-- | Symbolic derivative by a variable+deriveByVar :: Int -> Fix SRTree -> Fix SRTree+deriveByVar = deriveBy False++-- | Symbolic derivative by a parameter+deriveByParam :: Int -> Fix SRTree -> Fix SRTree+deriveByParam = deriveBy True++-- | Relabel the parameters incrementaly starting from 0+relabelParams :: Fix SRTree -> Fix SRTree+relabelParams t = cataM lTor alg t `evalState` 0 where- toState :: Num ix => SRTree ix val -> State ix (SRTree ix val)- toState (Param x) = do n <- get; put (n+1); pure (Param n)- toState (Add l r) = do l' <- toState l; r' <- toState r; pure (Add l' r')- toState (Sub l r) = do l' <- toState l; r' <- toState r; pure (Sub l' r')- toState (Mul l r) = do l' <- toState l; r' <- toState r; pure (Mul l' r')- toState (Div l r) = do l' <- toState l; r' <- toState r; pure (Div l' r')- toState (Power l r) = do l' <- toState l; r' <- toState r; pure (Power l' r')- toState (LogBase l r) = do l' <- toState l; r' <- toState r; pure (LogBase l' r')- toState (Fun f n) = do n' <- toState n; pure (Fun f n')- toState (Pow n i) = do n' <- toState n; pure (Pow n' i)- toState n = pure n+ lTor (Uni f mt) = Uni f <$> mt;+ lTor (Bin f ml mr) = Bin f <$> ml <*> mr+ lTor (Var ix) = pure (Var ix)+ lTor (Param ix) = pure (Param ix)+ lTor (Const c) = pure (Const c)++ alg :: SRTree (Fix SRTree) -> State Int (Fix SRTree)+ alg (Var ix) = pure $ var ix+ alg (Param ix) = do iy <- get; modify (+1); pure (param iy)+ alg (Const c) = pure $ Fix $ Const c+ alg (Uni f t) = pure $ Fix (Uni f t)+ alg (Bin f l r) = pure $ Fix (Bin f l r)++-- | Change constant values to a parameter, returning the changed tree and a list+-- of parameter values+constsToParam :: Fix SRTree -> (Fix SRTree, [Double])+constsToParam = first relabelParams . cata alg+ where+ first f (x, y) = (f x, y)++ alg (Var ix) = (Fix $ Var ix, [])+ alg (Param ix) = (Fix $ Param ix, [1.0])+ alg (Const c) = (Fix $ Param 0, [c])+ alg (Uni f t) = (Fix $ Uni f (fst t), snd t)+ alg (Bin f l r) = (Fix (Bin f (fst l) (fst r)), snd l <> snd r)++-- | Same as `constsToParam` but does not change constant values that+-- can be converted to integer without loss of precision+floatConstsToParam :: Fix SRTree -> (Fix SRTree, [Double])+floatConstsToParam = first relabelParams . cata alg+ where+ first f (x, y) = (f x, y)++ alg (Var ix) = (Fix $ Var ix, [])+ alg (Param ix) = (Fix $ Param ix, [1.0])+ alg (Const c) = if floor c == ceiling c then (Fix $ Const c, []) else (Fix $ Param 0, [c])+ alg (Uni f t) = (Fix $ Uni f (fst t), snd t)+ alg (Bin f l r) = (Fix (Bin f (fst l) (fst r)), snd l <> snd r)
src/Data/SRTree/Print.hs view
@@ -11,14 +11,14 @@ -- ----------------------------------------------------------------------------- module Data.SRTree.Print - ( DisplayNodes(..)- , showExpr- , showTree+ ( showExpr , printExpr- , showDefault , showTikz+ , printTikz , showPython+ , printPython , showLatex+ , printLatex ) where @@ -26,150 +26,37 @@ import Data.Char ( toLower ) import Data.SRTree.Internal---- | Data structure containing the needed definitions to print a SRTree.-data DisplayNodes ix val = D- { _displayVar :: ix -> String- , _displayPar :: ix -> String- , _displayVal :: val -> String- , _displayFun :: Function -> String- , _displayPow :: String- , _displayFloatPow :: String- }---- Auxiliary function to print a tree as an infix expression-asExpr :: (Show ix, Show val) => SRTree ix val -> Reader (DisplayNodes ix val) String-asExpr Empty = pure ""-asExpr (Var ix) = do- display <- asks _displayVar- pure $ display ix-asExpr (Param ix) = do- display <- asks _displayPar- pure $ display ix-asExpr (Const val) = do- display <- asks _displayVal- pure $ display val -asExpr (Fun f t) = do- display <- asks _displayFun- st <- asExpr t- pure $ mconcat [display f, "(", st, ")"]-asExpr (Pow t ix) = do- st <- asExpr t- pow <- asks _displayPow- pure $ mconcat ["(", st, ")", pow, "(", show ix, ")"]-asExpr (Add l r) = do- sl <- asExpr l- sr <- asExpr r- pure $ mconcat ["(", sl, ") + (", sr, ")"]-asExpr (Sub l r) = do- sl <- asExpr l- sr <- asExpr r- pure $ mconcat ["(", sl, ") - (", sr, ")"]-asExpr (Mul l r) = do- sl <- asExpr l- sr <- asExpr r- pure $ mconcat ["(", sl, ") * (", sr, ")"]-asExpr (Div l r) = do- sl <- asExpr l- sr <- asExpr r- pure $ mconcat ["(", sl, ") / (", sr, ")"]-asExpr (Power l r) = do- sl <- asExpr l- sr <- asExpr r- pow <- asks _displayFloatPow- pure $ mconcat ["(", sl, ")", pow, "(", sr, ")"]-asExpr (LogBase l r) = do- sl <- asExpr l- sr <- asExpr r- pure $ mconcat ["log(", sl, ",", sr, ")"]---- Auxiliary function to print a tree as a tree-like structure-asTree :: (Show ix, Show val) => SRTree ix val -> Reader (DisplayNodes ix val) String-asTree Empty = pure ""-asTree (Var ix) = do- display <- asks _displayVar- pure $ mconcat ["[", display ix, "]\n"]-asTree (Param ix) = do- display <- asks _displayPar- pure $ mconcat ["[", display ix, "]\n"]-asTree (Const val) = do- display <- asks _displayVal- pure $ mconcat ["[", display val, "]\n"]-asTree (Fun f t) = do- display <- asks _displayFun- st <- asTree t- pure $ mconcat ["[", display f, "\n", st, "]\n"]-asTree (Pow t ix) = do- st <- asTree t- pow <- asks _displayPow- pure $ mconcat ["[", pow, "\n", st, "[", show ix, "]\n]"] -asTree (Add l r) = do- sl <- asTree l- sr <- asTree r- pure $ mconcat ["[+\n", sl, sr, "]\n"]-asTree (Sub l r) = do- sl <- asTree l- sr <- asTree r- pure $ mconcat ["[-\n", sl, sr, "]\n"]-asTree (Mul l r) = do- sl <- asTree l- sr <- asTree r- pure $ mconcat ["[×\n", sl, sr, "]\n"]-asTree (Div l r) = do- sl <- asTree l- sr <- asTree r- pure $ mconcat ["[÷\n", sl, sr, "]\n"]-asTree (Power l r) = do- sl <- asTree l- sr <- asTree r- pow <- asks _displayFloatPow- pure $ mconcat ["[", pow, "\n", sl, sr, "]\n"]-asTree (LogBase l r) = do- sl <- asTree l- sr <- asTree r- pure $ mconcat ["[log\n", sl, sr, "]\n"]---- | Converts a tree to a `String` using the specifications given by `DisplayNodes`-showExpr, showTree :: (Show ix, Show val) => SRTree ix val -> DisplayNodes ix val -> String-showExpr t = runReader (asExpr t)-{-# INLINE showExpr #-}-showTree t = runReader (asTree t)-{-# INLINE showTree #-}---- | Prints a tree as an expression using the specifications given by `DisplayNodes`-printExpr :: (Show ix, Show val) => SRTree ix val -> DisplayNodes ix val -> IO ()-printExpr t = putStrLn . showExpr t+import Data.SRTree.Recursion --- | Displays a tree as an expression-showDefault t = showExpr t d+showExpr :: Fix SRTree -> String+showExpr = cata alg where- d = D (\ix -> mconcat ["x", show ix])- (\ix -> mconcat ["t", show ix])- show- show- "^"- "**"+ alg (Var ix) = 'x' : show ix+ alg (Param ix) = 't' : show ix+ alg (Const c) = show c+ alg (Bin op l r) = concat ["(", l, " ", showOp op, " ", r, ")"]+ alg (Uni f t) = concat [show f, "(", t, ")"] --- | Displays a tree in Tikz format-showTikz :: (Show ix, Show val, RealFrac val) => SRTree ix val -> String -showTikz t = showTree t d- where- d = D (\ix -> mconcat ["$x_{", show ix, "}$"])- (\ix -> mconcat ["$\\theta_{", show ix, "}$"])- (\val -> mconcat ["$", show $ (/100) $ fromIntegral $ round $ val*100, "$"])- show- "\\^{}"- "**"+printExpr :: Fix SRTree -> IO ()+printExpr = putStrLn . showExpr +showOp Add = "+"+showOp Sub = "-"+showOp Mul = "*"+showOp Div = "/"+showOp Power = "^"+{-# INLINE showOp #-}+ -- | Displays a tree as a numpy compatible expression.-showPython t = showExpr t d+showPython :: Fix SRTree -> String+showPython = cata alg where- d = D (\ix -> mconcat ["x[:,", show ix, "]"])- (\ix -> mconcat ["t[", show ix, "]"])- show- pyFun- "**"- "**"+ alg (Var ix) = concat ["x[:, ", show ix, "]"]+ alg (Param ix) = concat ["t[:, ", show ix, "]"]+ alg (Const c) = show c+ alg (Bin Power l r) = concat [l, " ** ", r]+ alg (Bin op l r) = concat ["(", l, " ", showOp op, " ", r, ")"]+ alg (Uni f t) = concat [pyFun f, "(", t, ")"] pyFun Id = "" pyFun Abs = "np.abs"@@ -190,21 +77,43 @@ pyFun Log = "np.log" pyFun Exp = "np.exp" +printPython :: Fix SRTree -> IO ()+printPython = putStrLn . showPython+ -- | Displays a tree as a sympy compatible expression.-showLatex :: (Show ix, Show val) => SRTree ix val -> String-showLatex Empty = ""-showLatex (Var ix) = mconcat ["x_{", show ix, "}"]-showLatex (Param ix) = mconcat ["\\theta_{", show ix, "}"]-showLatex (Const val) = show val-showLatex (Fun Abs t) = mconcat ["\\left |", showLatex t, "\\right |"]-showLatex (Fun f t) = mconcat [showLatexFun f, "\\left(", showLatex t, "\\right)"]-showLatex (Pow t ix) = mconcat ["\\left(", showLatex t, "\\right)^{", show ix, "}"]-showLatex (Add l r) = mconcat ["\\left(", showLatex l, "\\right) + \\left(", showLatex r, "\\right)"]-showLatex (Sub l r) = mconcat ["\\left(", showLatex l, "\\right) - \\left(", showLatex r, "\\right)"]-showLatex (Mul l r) = mconcat ["\\left(", showLatex l, "\\right) \\left(", showLatex r, "\\right)"]-showLatex (Div l r) = mconcat ["\\frac{", showLatex l, "}{", showLatex r, "}"]-showLatex (Power l r) = mconcat ["\\left(", showLatex l, "\\right)^{", showLatex r, "}"]-showLatex (LogBase l r) = mconcat ["\\log_{", showLatex r, "}{", showLatex l, "}"]+showLatex :: Fix SRTree -> String+showLatex = cata alg+ where+ alg (Var ix) = concat ["x_{, ", show ix, "}"]+ alg (Param ix) = concat ["\\theta_{, ", show ix, "}"]+ alg (Const c) = show c+ alg (Bin Power l r) = concat [l, "^{", r, "}"]+ alg (Bin op l r) = concat ["\\left(", l, " ", showOp op, " ", r, "\\right)"]+ alg (Uni Abs t) = concat ["\\left |", t, "\\right |"]+ alg (Uni f t) = concat [showLatexFun f, "(", t, ")"] showLatexFun :: Function -> String showLatexFun f = mconcat ["\\operatorname{", map toLower $ show f, "}"]+{-# INLINE showLatexFun #-}++printLatex :: Fix SRTree -> IO ()+printLatex = putStrLn . showLatex++-- | Displays a tree in Tikz format+showTikz :: Fix SRTree -> String+showTikz = cata alg+ where+ roundN n x = let ten = 10^n in (/ ten) . fromIntegral . round $ x*ten+ alg (Var ix) = concat ["[$x_{, ", show ix, "}$]\n"]+ alg (Param ix) = concat ["[$\\theta_{, ", show ix, "}$]\n"]+ alg (Const c) = concat ["[$", show (roundN 2 c), "$]\n"]+ alg (Bin op l r) = concat ["[", showOpTikz op, l, r, "]\n"]+ alg (Uni f t) = concat ["[", map toLower $ show f, t, "]\n"]++ showOpTikz Add = "+\n"+ showOpTikz Sub = "-\n"+ showOpTikz Mul = "×\n"+ showOpTikz Div = "÷\n"+ showOpTikz Power = "\\^{}\n"++printTikz = putStrLn . showTikz
src/Data/SRTree/Random.hs view
@@ -35,6 +35,7 @@ import Data.Maybe (fromJust) import Data.SRTree.Internal+import Data.SRTree.Recursion -- * Class definition of properties that a certain parameter type has. --@@ -43,19 +44,19 @@ -- HasExps: does `p` provides a range for the integral exponentes? -- HasFuns: does `p` provides a list of allowed functions? class HasVars p where- _vars :: p ix val -> [ix]+ _vars :: p -> [Int] class HasVals p where- _range :: p ix val -> (val, val)+ _range :: p -> (Double, Double) class HasExps p where- _exponents :: p ix val -> (Int, Int)+ _exponents :: p -> (Int, Int) class HasFuns p where- _funs :: p ix val -> [Function]+ _funs :: p -> [Function] -- | Constraint synonym for all properties. type HasEverything p = (HasVars p, HasVals p, HasExps p, HasFuns p) -- | A structure with every property-data FullParams ix val = P [ix] (val, val) (Int, Int) [Function]+data FullParams = P [Int] (Double, Double) (Int, Int) [Function] instance HasVars FullParams where _vars (P ixs _ _ _) = ixs@@ -83,79 +84,72 @@ {-# INLINE randomRange #-} -- Replace the child of a unary tree.-replaceChild :: SRTree ix val -> SRTree ix val -> Maybe (SRTree ix val)-replaceChild (Fun g _) t = Just $ Fun g t-replaceChild (Pow _ k) t = Just $ Pow t k+replaceChild :: Fix SRTree -> Fix SRTree -> Maybe (Fix SRTree)+replaceChild (Fix (Uni f _)) t = Just $ Fix (Uni f t) replaceChild _ _ = Nothing {-# INLINE replaceChild #-} -- Replace the children of a binary tree.-replaceChildren :: SRTree ix val -> SRTree ix val -> SRTree ix val -> Maybe (SRTree ix val)-replaceChildren (Add _ _) l r = Just $ Add l r-replaceChildren (Sub _ _) l r = Just $ Sub l r-replaceChildren (Mul _ _) l r = Just $ Mul l r-replaceChildren (Div _ _) l r = Just $ Div l r-replaceChildren (Power _ _) l r = Just $ Power l r-replaceChildren (LogBase _ _) l r = Just $ LogBase l r+replaceChildren :: Fix SRTree -> Fix SRTree -> Fix SRTree -> Maybe (Fix SRTree)+replaceChildren (Fix (Bin f _ _)) l r = Just $ Fix (Bin f l r) replaceChildren _ _ _ = Nothing {-# INLINE replaceChildren #-} -- | RndTree is a Monad Transformer to generate random trees of type `SRTree ix val` -- given the parameters `p ix val` using the random number generator `StdGen`.-type RndTree p ix val = ReaderT (p ix val) (StateT StdGen IO) (SRTree ix val)+type RndTree p = ReaderT p (StateT StdGen IO) (Fix SRTree) -- | Returns a random variable, the parameter `p` must have the `HasVars` property-randomVar :: HasVars p => RndTree p ix val+randomVar :: HasVars p => RndTree p randomVar = do vars <- asks _vars- lift $ Var <$> randomFrom vars+ lift $ Fix . Var <$> randomFrom vars -- | Returns a random constant, the parameter `p` must have the `HasConst` property-randomConst :: (Ord val, Random val, HasVals p) => RndTree p ix val+randomConst :: HasVals p => RndTree p randomConst = do rng <- asks _range- lift $ Const <$> randomRange rng+ lift $ Fix . Const <$> randomRange rng -- | Returns a random integer power node, the parameter `p` must have the `HasExps` property-randomPow :: (Ord val, Random val, HasExps p) => RndTree p ix val+randomPow :: HasExps p => RndTree p randomPow = do rng <- asks _exponents- lift $ Pow Empty <$> randomRange rng+ lift $ Fix . Bin Power 0 . Fix . Const . fromIntegral <$> randomRange rng -- | Returns a random function, the parameter `p` must have the `HasFuns` property-randomFunction :: HasFuns p => RndTree p ix val+randomFunction :: HasFuns p => RndTree p randomFunction = do funs <- asks _funs- lift $ (`Fun` Empty) <$> randomFrom funs+ f <- lift $ randomFrom funs+ lift $ pure $ Fix (Uni f 0) -- | Returns a random node, the parameter `p` must have every property.-randomNode :: (Ord val, Random val, HasEverything p) => RndTree p ix val+randomNode :: HasEverything p => RndTree p randomNode = do- choice <- lift $ randomRange (0, 9 :: Int)+ choice <- lift $ randomRange (0, 8 :: Int) case choice of 0 -> randomVar 1 -> randomConst 2 -> randomFunction 3 -> randomPow- 4 -> pure $ Add Empty Empty- 5 -> pure $ Sub Empty Empty- 6 -> pure $ Mul Empty Empty- 7 -> pure $ Div Empty Empty- 8 -> pure $ Power Empty Empty- 9 -> pure $ LogBase Empty Empty+ 4 -> pure . Fix $ Bin Add 0 0+ 5 -> pure . Fix $ Bin Sub 0 0+ 6 -> pure . Fix $ Bin Mul 0 0+ 7 -> pure . Fix $ Bin Div 0 0+ 8 -> pure . Fix $ Bin Power 0 0 -- | Returns a random non-terminal node, the parameter `p` must have every property.-randomNonTerminal :: (Ord val, Random val, HasEverything p) => RndTree p ix val+randomNonTerminal :: HasEverything p => RndTree p randomNonTerminal = do- choice <- lift $ randomRange (0, 7 :: Int)+ choice <- lift $ randomRange (0, 6 :: Int) case choice of 0 -> randomFunction 1 -> randomPow- 2 -> pure $ Add Empty Empty- 3 -> pure $ Sub Empty Empty- 4 -> pure $ Mul Empty Empty- 5 -> pure $ Div Empty Empty- 6 -> pure $ Power Empty Empty- 7 -> pure $ LogBase Empty Empty+ 2 -> pure . Fix $ Bin Add 0 0+ 3 -> pure . Fix $ Bin Sub 0 0+ 4 -> pure . Fix $ Bin Mul 0 0+ 5 -> pure . Fix $ Bin Div 0 0+ 6 -> pure . Fix $ Bin Power 0 0 -- | Returns a random tree with a limited budget, the parameter `p` must have every property.-randomTree :: (Ord val, Random val, HasEverything p) => Int -> RndTree p ix val+randomTree :: HasEverything p => Int -> RndTree p randomTree 0 = do coin <- lift toss if coin@@ -169,7 +163,7 @@ 2 -> replaceChildren node <$> randomTree (budget `div` 2) <*> randomTree (budget `div` 2) -- | Returns a random tree with a approximately a number `n` of nodes, the parameter `p` must have every property.-randomTreeBalanced :: (Ord val, Random val, HasEverything p) => Int -> RndTree p ix val+randomTreeBalanced :: HasEverything p => Int -> RndTree p randomTreeBalanced n | n <= 1 = do coin <- lift toss if coin
+ src/Data/SRTree/Recursion.hs view
@@ -0,0 +1,89 @@+{-# language RankNTypes #-}+{-# language DeriveFunctor #-}+module Data.SRTree.Recursion where++import Control.Monad ( (>=>) )++data ListF a b = NilF | ConsF a b deriving Functor+data NatF a = ZeroF | SuccF a deriving Functor+data StreamF a b = StreamF a b deriving Functor+data TreeF a b = LeafF | NodeF b a b deriving Functor++newtype Fix f = Fix {unfix :: f (Fix f)}++type Algebra f a = f a -> a+type CoAlgebra f a = a -> f a++data Cofree f a = a :< f (Cofree f a)+data Free f a = Ret a | Op (f (Free f a))++extract :: Cofree f a -> a+extract (x :< _) = x++unOp :: Free f a -> f (Free f a)+unOp (Op x) = x+unOp _ = error "partial function unOp called on Ret"++cata :: Functor f => (f a -> a) -> Fix f -> a+cata alg = alg . fmap (cata alg) . unfix++cataM :: (Functor f, Monad m) => (forall x . f (m x) -> m (f x)) -> (f a -> m a) -> Fix f -> m a+cataM seq alg = cata (seq >=> alg)++ana :: Functor f => (a -> f a) -> a -> Fix f+ana coalg = Fix . fmap (ana coalg) . coalg++hylo :: Functor f => (f b -> b) -> (a -> f a) -> a -> b+hylo alg coalg = alg . fmap (cata alg . ana coalg) . coalg++para :: Functor f => (f (Fix f, a) -> a) -> Fix f -> a+para alg = alg . fmap (id &&& para alg) . unfix+ where (f &&& g) x = (f x, g x)++mutu :: Functor f => (f (a, b) -> a) -> (f (a, b) -> b) -> (Fix f -> a, Fix f -> b)+mutu alg1 alg2 = (fst . cata alg, snd . cata alg)+ where alg x = (alg1 x, alg2 x)++apo :: Functor f => (a -> f (Either (Fix f) a)) -> a -> Fix f+apo coalg = Fix . fmap (id ||| apo coalg) . coalg+ where + (f ||| g) (Left x) = f x+ (f ||| g) (Right y) = g y++accu :: Functor f => (forall x. f x -> p -> f (x, p)) -> (f a -> p -> a) -> Fix f -> p -> a+accu st alg (Fix t) p = alg (fmap (uncurry (accu st alg)) (st t p)) p++histo :: Functor f => (f (Cofree f a) -> a) -> Fix f -> a+histo alg = extract . cata (\x -> alg x :< x)++futu :: Functor f => (a -> f (Free f a)) -> a -> Fix f+futu coalg = ana coalg' . Ret+ where+ coalg' (Ret a) = coalg a+ coalg' (Op k) = k++chrono :: Functor f => (f (Cofree f b) -> b) -> (a -> f (Free f a)) -> a -> b+chrono alg coalg = extract . hylo alg' coalg' . Ret+ where+ alg' x = alg x :< x+ coalg' (Ret a) = coalg a+ coalg' (Op k) = k++fromList :: [a] -> Fix (ListF a)+fromList [] = Fix NilF+fromList (x:xs) = Fix (ConsF x (fromList xs))++toList :: Fix (ListF a) -> [a]+toList (Fix NilF) = []+toList (Fix (ConsF x xs)) = x : toList xs++stream2list :: StreamF a [a] -> [a]+stream2list (StreamF x y) = x : y++toNat :: Int -> Fix NatF+toNat 0 = Fix ZeroF+toNat n = Fix (SuccF (toNat (n-1)))++fromNat :: Fix NatF -> Int+fromNat (Fix ZeroF) = 0+fromNat (Fix (SuccF x)) = 1 + fromNat x
srtree.cabal view
@@ -5,7 +5,7 @@ -- see: https://github.com/sol/hpack name: srtree-version: 0.1.2.1+version: 1.0.0.0 synopsis: A general framework to work with Symbolic Regression expression trees. description: A Symbolic Regression Tree data structure to work with mathematical expressions with support to first order derivative and simplification; category: Math, Data, Data Structures@@ -31,6 +31,7 @@ Data.SRTree.Internal Data.SRTree.Print Data.SRTree.Random+ Data.SRTree.Recursion other-modules: Paths_srtree hs-source-dirs:@@ -40,9 +41,26 @@ , containers ==0.6.* , mtl ==2.2.* , random ==1.2.*- , vector ==0.12.*+ , vector >=0.12 && <=0.13 default-language: Haskell2010 +executable bench-srtree+ main-is: Main.hs+ other-modules:+ Paths_srtree+ hs-source-dirs:+ app+ ghc-options: -threaded -rtsopts -with-rtsopts=-N -O2 -optc-O3+ build-depends:+ base >=4.16 && <4.18+ , containers ==0.6.*+ , criterion >=1.6.0 && <1.7+ , mtl ==2.2.*+ , random ==1.2.*+ , srtree+ , vector >=0.12 && <=0.13+ default-language: Haskell2010+ test-suite srtree-test type: exitcode-stdio-1.0 main-is: Spec.hs@@ -52,10 +70,12 @@ test ghc-options: -threaded -rtsopts -with-rtsopts=-N build-depends:- base >=4.16 && <4.18+ HUnit+ , ad+ , base >=4.16 && <4.18 , containers ==0.6.* , mtl ==2.2.* , random ==1.2.* , srtree- , vector ==0.12.*+ , vector >=0.12 && <=0.13 default-language: Haskell2010
test/Spec.hs view
@@ -1,17 +1,76 @@ import Data.SRTree-import Data.SRTree.Random-import Data.SRTree.Print -import System.Random-import Control.Monad.State-import Control.Monad.Reader+import qualified Data.Vector as V+import Numeric.AD.Double ( grad )+import Test.HUnit -runThing g n = flip evalStateT g . runReaderT (randomTree n)+-- test expressions+exprs = [+ param 0 * sin ( param 1)+ , sin (param 0) + cos (param 1)+ , 0.5 * sin (param 0) + 0.7 * cos (param 1)+ , log (param 0) + param 0 * param 1 - sin (param 1)+ , 1 / param 0 * param 1+ , param 0 + param 1 + param 0 * param 1 + sin (param 0) + sin (param 1) + cos (param 0) + cos (param 1) + sin (param 0 * param 1) + cos (param 0 * param 1)+ , sin (exp (param 0) + param 1)+ ] +-- autodiff with multiple occurrences of vars+autoDiffMult :: [[Double]]+autoDiffMult = [ grad (\[x,y] -> x * sin y) [2,3]+ , grad (\[x,y] -> sin x + cos y) [2,3]+ , grad (\[x,y] -> 0.5 * sin x + 0.7 * cos y) [2,3]+ , grad (\[x,y] -> log x + x*y - sin y) [2,3]+ , grad (\[x,y] -> 1 / x * y) [2,3]+ , grad (\[x,y] -> x + y + x * y + sin x + sin y + cos x + cos y + sin (x * y) + cos (x * y)) [2,3]+ , grad (\[x,y] -> sin (exp x + y)) [2,3]+ ]++-- autodiff with single occurrences of vars+autoDiffSingle :: [[Double]]+autoDiffSingle = [ grad (\[x,y] -> x * sin y) [2,3]+ , grad (\[x,y] -> sin x + cos y) [2,3]+ , grad (\[x,y] -> 0.5 * sin x + 0.7 * cos y) [2,3]+ , grad (\[x,y,v,w] -> log x + y*v - sin w) [2,3,2,3]+ , grad (\[x,y] -> 1 / x * y) [2,3]+ , grad (\[a,b,c,d,e,f,g,h,i,j,k,l] -> a + b + c * d + sin e + sin f + cos g + cos h + sin (i * j) + cos (k * l)) [2,3,2,3,2,3,2,3,2,3,2,3]+ , grad (\[x,y] -> sin (exp x + y)) [2,3]+ ]++-- xs is empty since we are interested in theta+xs :: V.Vector a+xs = V.empty+-- theta values+thetaMulti, thetaSingle :: V.Vector Double+thetaMulti = V.fromList [2.0, 3.0]+thetaSingle = V.fromList [2.0, 3.0, 2.0, 3.0, 2.0, 3.0, 2.0, 3.0, 2.0, 3.0, 2.0, 3.0]++-- values from forward mode+forwardVals :: [[Double]]+forwardVals = map (forwardMode xs thetaMulti id) exprs++-- values from grad+-- we must relabel the parameters of the expression to sequence values+gradVals :: [(Double, [Double])]+gradVals = map (gradParams xs thetaSingle id . relabelParams) exprs++-- values of the evaluated expressions+exprVals :: [Double]+exprVals = map (evalTree xs thetaSingle id . relabelParams) exprs++refGrad :: [(Double, [Double])]+refGrad = zip exprVals autoDiffSingle++testDiff :: (Eq a, Show a) => String -> String -> a -> a -> Test+testDiff lbl name a b = TestLabel lbl $ TestCase (assertEqual name a b)++tests :: Test+tests = TestList $+ zipWith (testDiff "forward mode" "autodiff x forward mode") autoDiffMult forwardVals+ <> zipWith (testDiff "opt. grad. parameters" "(evalTree, autodiff) x gradVals") refGrad gradVals+ <> zipWith (testDiff "deriveByParam" "deriveByParam x autodiff") (map head autoDiffSingle) (map (head.snd) gradVals)+ main :: IO () main = do- g <- getStdGen- t <- runThing g 10 $ P [0,1] (-1.0, 1.0) (-3, 3) [Id, Sin]- print (t :: SRTree Int Double)- putStrLn $ showDefault t- putStrLn $ showTikz t+ result <- runTestTT tests+ putStrLn $ showCounts result