hegg 0.3.0.0 → 0.4.0.0
raw patch · 24 files changed
+116/−153 lines, 24 filesdep −deriving-compatdep ~basePVP ok
version bump matches the API change (PVP)
Dependencies removed: deriving-compat
Dependency ranges changed: base
API changes (from Hackage documentation)
- Data.Equality.Graph.Classes: instance (GHC.Show.Show a, Data.Functor.Classes.Show1 l) => GHC.Show.Show (Data.Equality.Graph.Classes.EClass a l)
- Data.Equality.Graph.Nodes: instance (Data.Functor.Classes.Show1 l, GHC.Show.Show a) => GHC.Show.Show (Data.Equality.Graph.Nodes.NodeMap l a)
- Data.Equality.Graph.Nodes: instance Data.Functor.Classes.Eq1 l => GHC.Classes.Eq (Data.Equality.Graph.Nodes.ENode l)
- Data.Equality.Graph.Nodes: instance Data.Functor.Classes.Eq1 l => GHC.Classes.Eq (Data.Equality.Graph.Nodes.Operator l)
- Data.Equality.Graph.Nodes: instance Data.Functor.Classes.Ord1 l => GHC.Base.Monoid (Data.Equality.Graph.Nodes.NodeMap l a)
- Data.Equality.Graph.Nodes: instance Data.Functor.Classes.Ord1 l => GHC.Base.Semigroup (Data.Equality.Graph.Nodes.NodeMap l a)
- Data.Equality.Graph.Nodes: instance Data.Functor.Classes.Ord1 l => GHC.Classes.Ord (Data.Equality.Graph.Nodes.ENode l)
- Data.Equality.Graph.Nodes: instance Data.Functor.Classes.Ord1 l => GHC.Classes.Ord (Data.Equality.Graph.Nodes.Operator l)
- Data.Equality.Graph.Nodes: instance Data.Functor.Classes.Show1 l => GHC.Show.Show (Data.Equality.Graph.Nodes.ENode l)
- Data.Equality.Graph.Nodes: instance Data.Functor.Classes.Show1 l => GHC.Show.Show (Data.Equality.Graph.Nodes.Operator l)
- Data.Equality.Matching.Pattern: instance Data.Functor.Classes.Eq1 l => GHC.Classes.Eq (Data.Equality.Matching.Pattern.Pattern l)
- Data.Equality.Matching.Pattern: instance Data.Functor.Classes.Ord1 l => GHC.Classes.Ord (Data.Equality.Matching.Pattern.Pattern l)
- Data.Equality.Matching.Pattern: instance Data.Functor.Classes.Show1 lang => GHC.Show.Show (Data.Equality.Matching.Pattern.Pattern lang)
- Data.Equality.Saturation.Rewrites: instance Data.Functor.Classes.Show1 lang => GHC.Show.Show (Data.Equality.Saturation.Rewrites.Rewrite anl lang)
- Data.Equality.Utils: instance Data.Functor.Classes.Eq1 f => GHC.Classes.Eq (Data.Equality.Utils.Fix f)
- Data.Equality.Utils: instance Data.Functor.Classes.Show1 f => GHC.Show.Show (Data.Equality.Utils.Fix f)
+ Data.Equality.Graph.Classes: instance (GHC.Show.Show a, GHC.Show.Show (l Data.Equality.Graph.Classes.Id.ClassId)) => GHC.Show.Show (Data.Equality.Graph.Classes.EClass a l)
+ Data.Equality.Graph.Lens: _iclasses :: Traversal (EGraph a l) (EGraph b l) (ClassId, EClass a l) (EClass b l)
+ Data.Equality.Graph.Lens: allOf :: Traversal s t a b -> (a -> Bool) -> s -> Bool
+ Data.Equality.Graph.Nodes: instance (GHC.Show.Show a, GHC.Show.Show (l Data.Equality.Graph.Classes.Id.ClassId)) => GHC.Show.Show (Data.Equality.Graph.Nodes.NodeMap l a)
+ Data.Equality.Graph.Nodes: instance GHC.Classes.Eq (l ()) => GHC.Classes.Eq (Data.Equality.Graph.Nodes.Operator l)
+ Data.Equality.Graph.Nodes: instance GHC.Classes.Eq (l Data.Equality.Graph.Classes.Id.ClassId) => GHC.Classes.Eq (Data.Equality.Graph.Nodes.ENode l)
+ Data.Equality.Graph.Nodes: instance GHC.Classes.Ord (l ()) => GHC.Classes.Ord (Data.Equality.Graph.Nodes.Operator l)
+ Data.Equality.Graph.Nodes: instance GHC.Classes.Ord (l Data.Equality.Graph.Classes.Id.ClassId) => GHC.Base.Monoid (Data.Equality.Graph.Nodes.NodeMap l a)
+ Data.Equality.Graph.Nodes: instance GHC.Classes.Ord (l Data.Equality.Graph.Classes.Id.ClassId) => GHC.Base.Semigroup (Data.Equality.Graph.Nodes.NodeMap l a)
+ Data.Equality.Graph.Nodes: instance GHC.Classes.Ord (l Data.Equality.Graph.Classes.Id.ClassId) => GHC.Classes.Ord (Data.Equality.Graph.Nodes.ENode l)
+ Data.Equality.Graph.Nodes: instance GHC.Show.Show (l ()) => GHC.Show.Show (Data.Equality.Graph.Nodes.Operator l)
+ Data.Equality.Graph.Nodes: instance GHC.Show.Show (l Data.Equality.Graph.Classes.Id.ClassId) => GHC.Show.Show (Data.Equality.Graph.Nodes.ENode l)
+ Data.Equality.Language: instance (forall a. GHC.Classes.Ord a => GHC.Classes.Ord (l a), Data.Traversable.Traversable l) => Data.Equality.Language.Language l
+ Data.Equality.Matching.Pattern: instance (forall a. GHC.Classes.Eq a => GHC.Classes.Eq (l a)) => GHC.Classes.Eq (Data.Equality.Matching.Pattern.Pattern l)
+ Data.Equality.Matching.Pattern: instance (forall a. GHC.Classes.Eq a => GHC.Classes.Eq (l a), forall a. GHC.Classes.Ord a => GHC.Classes.Ord (l a)) => GHC.Classes.Ord (Data.Equality.Matching.Pattern.Pattern l)
+ Data.Equality.Matching.Pattern: instance (forall a. GHC.Show.Show a => GHC.Show.Show (lang a)) => GHC.Show.Show (Data.Equality.Matching.Pattern.Pattern lang)
+ Data.Equality.Saturation.Rewrites: instance (forall a. GHC.Show.Show a => GHC.Show.Show (lang a)) => GHC.Show.Show (Data.Equality.Saturation.Rewrites.Rewrite anl lang)
+ Data.Equality.Utils: instance (forall a. GHC.Classes.Eq a => GHC.Classes.Eq (f a)) => GHC.Classes.Eq (Data.Equality.Utils.Fix f)
+ Data.Equality.Utils: instance (forall a. GHC.Show.Show a => GHC.Show.Show (f a)) => GHC.Show.Show (Data.Equality.Utils.Fix f)
- Data.Equality.Graph.Lens: _classes :: Lens' (EGraph a l) (ClassIdMap (EClass a l))
+ Data.Equality.Graph.Lens: _classes :: Traversal (EGraph a l) (EGraph b l) (EClass a l) (EClass b l)
- Data.Equality.Graph.Lens: type Lens' s a = forall f. Functor f => (a -> f a) -> (s -> f s)
+ Data.Equality.Graph.Lens: type Traversal s t a b = forall f. Applicative f => (a -> f b) -> (s -> f t)
- Data.Equality.Graph.Nodes: deleteNM :: Ord1 l => ENode l -> NodeMap l a -> NodeMap l a
+ Data.Equality.Graph.Nodes: deleteNM :: Ord (l ClassId) => ENode l -> NodeMap l a -> NodeMap l a
- Data.Equality.Graph.Nodes: foldlWithKeyNM' :: Ord1 l => (b -> ENode l -> a -> b) -> b -> NodeMap l a -> b
+ Data.Equality.Graph.Nodes: foldlWithKeyNM' :: Ord (l ClassId) => (b -> ENode l -> a -> b) -> b -> NodeMap l a -> b
- Data.Equality.Graph.Nodes: foldrWithKeyNM' :: Ord1 l => (ENode l -> a -> b -> b) -> b -> NodeMap l a -> b
+ Data.Equality.Graph.Nodes: foldrWithKeyNM' :: Ord (l ClassId) => (ENode l -> a -> b -> b) -> b -> NodeMap l a -> b
- Data.Equality.Graph.Nodes: insertLookupNM :: Ord1 l => ENode l -> a -> NodeMap l a -> (Maybe a, NodeMap l a)
+ Data.Equality.Graph.Nodes: insertLookupNM :: Ord (l ClassId) => ENode l -> a -> NodeMap l a -> (Maybe a, NodeMap l a)
- Data.Equality.Graph.Nodes: insertNM :: Ord1 l => ENode l -> a -> NodeMap l a -> NodeMap l a
+ Data.Equality.Graph.Nodes: insertNM :: Ord (l ClassId) => ENode l -> a -> NodeMap l a -> NodeMap l a
- Data.Equality.Graph.Nodes: lookupNM :: Ord1 l => ENode l -> NodeMap l a -> Maybe a
+ Data.Equality.Graph.Nodes: lookupNM :: Ord (l ClassId) => ENode l -> NodeMap l a -> Maybe a
- Data.Equality.Language: class (Traversable l, Ord1 l) => Language l
+ Data.Equality.Language: class (forall a. Ord a => Ord (l a), Traversable l) => Language l
Files
- CHANGELOG.md +9/−0
- README.md +7/−13
- hegg.cabal +1/−3
- src/Data/Equality/Analysis.hs +1/−0
- src/Data/Equality/Extraction.hs +3/−2
- src/Data/Equality/Graph.hs +6/−1
- src/Data/Equality/Graph/Classes.hs +1/−3
- src/Data/Equality/Graph/Internal.hs +1/−3
- src/Data/Equality/Graph/Lens.hs +16/−4
- src/Data/Equality/Graph/Monad.hs +2/−0
- src/Data/Equality/Graph/Nodes.hs +24/−33
- src/Data/Equality/Language.hs +11/−2
- src/Data/Equality/Matching.hs +1/−0
- src/Data/Equality/Matching/Database.hs +1/−1
- src/Data/Equality/Matching/Pattern.hs +7/−7
- src/Data/Equality/Saturation.hs +4/−2
- src/Data/Equality/Saturation/Rewrites.hs +2/−3
- src/Data/Equality/Utils.hs +5/−6
- test/Invariants.hs +7/−7
- test/Lambda.hs +2/−16
- test/SimpleSym.hs +1/−13
- test/Sym.hs +2/−13
- test/T1.hs +1/−11
- test/T2.hs +1/−10
CHANGELOG.md view
@@ -2,6 +2,15 @@ ## Unreleased +## 0.4.0.0 -- 2023-06-24++* Make `Language` a constraint type synonym instead of a standalone empty class+* Use `QuantifiedConstraints` instead of `Eq1,Ord1,Show1` in the implementation,+ which results in the user only having to provide an `Eq a => Eq (language+ a)` instance rather than a `Eq1 language` one (which is much simpler and can+ usually be done automatically!)+* Make `_classes` a `Traversal` lens over all e-classes rather than a `Lens` into `IntMap EClass`+ ## 0.3.0.0 -- 2022-12-09 * A better `Analysis` tutorial in the README.
README.md view
@@ -47,12 +47,14 @@ represented in e-graph and on which equality saturation can be run: ```hs-class (Analysis l, Traversable l, Ord1 l) => Language l+type Language l = (Traversable l, ∀ a. Ord a => Ord (l a)) ``` -To declare a `Language` we must write the "base functor" of `SymExpr` -(i.e. use a type parameter where the recursion points used to be in the original `SymExpr`),-then instance `Traversable`, `Ord1`, and write an `Analysis` instance for it (see next section).+To declare a `Language` we must write the "base functor" of `SymExpr` (i.e. use+a type parameter where the recursion points used to be in the original+`SymExpr`), then instance `Traversable l`, `∀ a. Ord a => Ord (l a)` (we can do+it automatically through deriving), and write an `Analysis` instance for it (see+next section). ```hs data SymExpr a = Const Double@@ -60,7 +62,7 @@ | a :+: a | a :*: a | a :/: a- deriving (Functor, Foldable, Traversable)+ deriving (Eq, Ord, Show, Functor, Foldable, Traversable) infix 6 :+: infix 7 :*:, :/: ```@@ -76,14 +78,6 @@ e1 :: Fix SymExpr e1 = Fix (Fix (Fix (Symbol "x") :*: Fix (Const 2)) :/: (Fix (Const 2))) -- (x*2)/2 ```--We've already automagically derived `Functor`, `Foldable` and `Traversable`-instances, and can use the following template haskell functions from `derive-compat` to derive `Ord1`.-```hs-deriveEq1 ''SymExpr-deriveOrd1 ''SymExpr-```- Then, we define an `Analysis` for our `SymExpr`. ### Analysis
hegg.cabal view
@@ -1,6 +1,6 @@ cabal-version: 2.4 name: hegg-version: 0.3.0.0+version: 0.4.0.0 Tested-With: GHC ==9.4.2 || ==9.2.2 || ==9.0.2 || ==8.10.7 synopsis: Fast equality saturation in Haskell @@ -114,7 +114,6 @@ build-depends: base, hegg, containers,- deriving-compat >= 0.6 && < 0.7, tasty >= 1.4 && < 1.5, tasty-hunit >= 0.10 && < 0.11, tasty-quickcheck >= 0.10 && < 0.11@@ -127,7 +126,6 @@ type: exitcode-stdio-1.0 build-depends: base, hegg, containers,- deriving-compat, tasty, tasty-hunit, tasty-quickcheck,
src/Data/Equality/Analysis.hs view
@@ -6,6 +6,7 @@ {-# LANGUAGE FlexibleContexts #-} {-# LANGUAGE TypeFamilies #-} {-# LANGUAGE MultiParamTypeClasses #-}+{-# LANGUAGE ImpredicativeTypes #-} {-| E-class analysis, which allows the concise expression of a program analysis over
src/Data/Equality/Extraction.hs view
@@ -1,5 +1,6 @@ {-# LANGUAGE ScopedTypeVariables #-} {-# LANGUAGE ViewPatterns #-}+{-# LANGUAGE MonoLocalBinds #-} {-| Given an e-graph representing expressions of our language, we might want to extract, out of all expressions represented by some equivalence class, /the best/@@ -21,9 +22,9 @@ import qualified Data.Set as S import qualified Data.IntMap.Strict as IM +import Data.Equality.Graph.Internal (EGraph(classes)) import Data.Equality.Utils import Data.Equality.Graph-import Data.Equality.Graph.Lens -- vvvv and necessarily all the best sub-expressions from children equilalence classes @@ -51,7 +52,7 @@ -- picking up the best from the target e-class. In practice this shouldn't -- find the cost of unused nodes because the "topmost" e-class will be the -- target, and all sub-classes must be calculated?- let allCosts = findCosts (egr^._classes) mempty+ let allCosts = findCosts (classes egr) mempty in case findBest i allCosts of Just (CostWithExpr (_,n)) -> n
src/Data/Equality/Graph.hs view
@@ -5,6 +5,7 @@ {-# LANGUAGE FlexibleContexts #-} {-# LANGUAGE ScopedTypeVariables #-} {-# LANGUAGE TypeFamilies #-}+{-# LANGUAGE NamedFieldPuns #-} {-| An e-graph efficiently represents a congruence relation over many expressions. @@ -31,7 +32,10 @@ , module Data.Equality.Language ) where +-- ROMES:TODO: Is the E-Graph a Monad if the analysis data were the type arg? i.e. Monad (EGraph language)?+ -- import GHC.Conc+import Prelude hiding (lookup) import Data.Function import Data.Bifunctor@@ -272,7 +276,7 @@ repairAnal :: forall a l. (Analysis a l, Language l) => (ClassId, ENode l) -> EGraph a l -> EGraph a l repairAnal (repair_id, node) egr = let- c1 = (egr^._classes) IM.! repair_id+ c1 = egr^._class repair_id new_data = joinA @a @l (c1^._data) (makeA @a ((\i -> egr^._class i^._data @a) <$> unNode node)) (c2, added_nodes) = modifyA (c1 & _data .~ new_data) in@@ -304,6 +308,7 @@ {-# INLINE canonicalize #-} -- | Find the canonical representation of an e-class id in the e-graph+-- -- Invariant: The e-class id always exists. find :: ClassId -> EGraph a l -> ClassId find cid = findRepr cid . unionFind
src/Data/Equality/Graph/Classes.hs view
@@ -10,8 +10,6 @@ import qualified Data.Set as S -import Data.Functor.Classes- import Data.Equality.Graph.Classes.Id import Data.Equality.Graph.Nodes @@ -30,6 +28,6 @@ , eClassParents :: !(SList (ClassId, ENode language)) -- ^ E-nodes which are parents of this e-class and their corresponding e-class ids. } -instance (Show a, Show1 l) => Show (EClass a l) where+instance (Show a, Show (l ClassId)) => Show (EClass a l) where show (EClass a b d (SList c _)) = "Id: " <> show a <> "\nNodes: " <> show b <> "\nParents: " <> show c <> "\nData: " <> show d
src/Data/Equality/Graph/Internal.hs view
@@ -6,8 +6,6 @@ -} module Data.Equality.Graph.Internal where -import Data.Functor.Classes- import Data.Equality.Graph.ReprUnionFind import Data.Equality.Graph.Classes import Data.Equality.Graph.Nodes@@ -31,7 +29,7 @@ -- | Maintained worklist of e-class ids that need to be “upward merged” type Worklist l = [(ClassId, ENode l)] -instance (Show a, Show1 l) => Show (EGraph a l) where+instance (Show a, Show (l ClassId)) => Show (EGraph a l) where show (EGraph a b c d e) = "UnionFind: " <> show a <> "\n\nE-Classes: " <> show b <>
src/Data/Equality/Graph/Lens.hs view
@@ -13,6 +13,7 @@ import Data.Functor.Identity import Data.Functor.Const+import Data.Monoid import Data.Equality.Utils.SizedList import Data.Equality.Graph.Internal@@ -23,7 +24,7 @@ -- | A 'Lens'' as defined in other lenses libraries type Lens' s a = forall f. Functor f => (a -> f a) -> (s -> f s)-+type Traversal s t a b = forall f. Applicative f => (a -> f b) -> (s -> f t) -- outdated comment for "getClass": --@@ -50,11 +51,16 @@ _memo afa egr = (\m1 -> egr {memo = m1}) <$> afa (memo egr) {-# INLINE _memo #-} --- | Lens for the map of existing classes by id in an e-graph-_classes :: Lens' (EGraph a l) (ClassIdMap (EClass a l))-_classes afa egr = (\m1 -> egr {classes = m1}) <$> afa (classes egr)+-- | Traversal for the existing classes in an e-graph+_classes :: Traversal (EGraph a l) (EGraph b l) (EClass a l) (EClass b l)+_classes afb egr = (\m1 -> egr {classes = m1}) <$> traverse afb (classes egr) {-# INLINE _classes #-} +-- | Traversal for the existing classes in an e-graph+_iclasses :: Traversal (EGraph a l) (EGraph b l) (ClassId, EClass a l) (EClass b l)+_iclasses afb egr = (\m1 -> egr {classes = m1}) <$> IM.traverseWithKey (curry afb) (classes egr)+{-# INLINE _iclasses #-}+ -- | Lens for the 'Domain' of an e-class _data :: Lens' (EClass domain l) domain _data afa EClass{..} = (\d1 -> EClass eClassId eClassNodes d1 eClassParents) <$> afa eClassData@@ -102,4 +108,10 @@ over :: Lens' s a -> (a -> a) -> (s -> s) over ln f = runIdentity . ln (Identity . f) {-# INLINE over #-}++-- | Returns True if every target of a Traversable satisfies a predicate.+allOf :: Traversal s t a b -> (a -> Bool) -> s -> Bool+allOf trv f = getAll . getConst . trv (Const . All . f)+{-# INLINE allOf #-}+
src/Data/Equality/Graph/Monad.hs view
@@ -1,4 +1,5 @@ {-# LANGUAGE TupleSections #-}+{-# LANGUAGE MonoLocalBinds #-} {-| Monadic interface to e-graph stateful computations -}@@ -82,3 +83,4 @@ runEGraphM :: EGraph anl l -> EGraphM anl l a -> (a, EGraph anl l) runEGraphM = flip runState {-# INLINE runEGraphM #-}+
src/Data/Equality/Graph/Nodes.hs view
@@ -1,10 +1,12 @@ {-# LANGUAGE TypeFamilies #-}-{-# LANGUAGE TupleSections #-} {-# LANGUAGE FlexibleInstances #-}-{-# LANGUAGE DeriveGeneric #-}-{-# LANGUAGE ViewPatterns #-}+{-# LANGUAGE UnicodeSyntax #-} {-# LANGUAGE GeneralizedNewtypeDeriving #-} {-# LANGUAGE DeriveTraversable #-}+{-# LANGUAGE RankNTypes #-}+{-# LANGUAGE QuantifiedConstraints #-}+{-# LANGUAGE StandaloneDeriving #-}+{-# LANGUAGE UndecidableInstances #-} {-| Module defining e-nodes ('ENode'), the e-node function symbol ('Operator'), and@@ -13,7 +15,6 @@ -} module Data.Equality.Graph.Nodes where -import Data.Functor.Classes import Data.Foldable import Data.Bifunctor @@ -34,6 +35,10 @@ -- parametrized over 'ClassId', i.e. all recursive fields are rather e-class ids. newtype ENode l = Node { unNode :: l ClassId } +deriving instance Eq (l ClassId) => (Eq (ENode l))+deriving instance Ord (l ClassId) => (Ord (ENode l))+deriving instance Show (l ClassId) => (Show (ENode l))+ -- | Get the children e-class ids of an e-node children :: Traversable l => ENode l -> [ClassId] children = toList . unNode@@ -45,68 +50,54 @@ -- this means children e-classes are ignored. newtype Operator l = Operator { unOperator :: l () } +deriving instance Eq (l ()) => (Eq (Operator l))+deriving instance Ord (l ()) => (Ord (Operator l))+deriving instance Show (l ()) => (Show (Operator l))+ -- | Get the operator (function symbol) of an e-node operator :: Traversable l => ENode l -> Operator l operator = Operator . void . unNode {-# INLINE operator #-} -instance Eq1 l => (Eq (ENode l)) where- (==) (Node a) (Node b) = liftEq (==) a b- {-# INLINE (==) #-}--instance Ord1 l => (Ord (ENode l)) where- compare (Node a) (Node b) = liftCompare compare a b- {-# INLINE compare #-}--instance Show1 l => (Show (ENode l)) where- showsPrec p (Node l) = liftShowsPrec showsPrec showList p l--instance Eq1 l => (Eq (Operator l)) where- (==) (Operator a) (Operator b) = liftEq (\_ _ -> True) a b- {-# INLINE (==) #-}--instance Ord1 l => (Ord (Operator l)) where- compare (Operator a) (Operator b) = liftCompare (\_ _ -> EQ) a b- {-# INLINE compare #-}--instance Show1 l => (Show (Operator l)) where- showsPrec p (Operator l) = liftShowsPrec (const . const $ showString "") (const $ showString "") p l- -- * Node Map -- | A mapping from e-nodes of @l@ to @a@ newtype NodeMap (l :: Type -> Type) a = NodeMap { unNodeMap :: M.Map (ENode l) a } -- TODO: Investigate whether it would be worth it requiring a trie-map for the -- e-node definition. Probably it isn't better since e-nodes aren't recursive.- deriving (Show, Functor, Foldable, Traversable, Semigroup, Monoid)+ deriving (Functor, Foldable, Traversable) +deriving instance (Show a, Show (l ClassId)) => Show (NodeMap l a)+deriving instance Ord (l ClassId) => Semigroup (NodeMap l a)+deriving instance Ord (l ClassId) => Monoid (NodeMap l a)+ -- | Insert a value given an e-node in a 'NodeMap'-insertNM :: Ord1 l => ENode l -> a -> NodeMap l a -> NodeMap l a+insertNM :: Ord (l ClassId) => ENode l -> a -> NodeMap l a -> NodeMap l a insertNM e v (NodeMap m) = NodeMap (M.insert e v m) {-# INLINE insertNM #-} -- | Lookup an e-node in a 'NodeMap'-lookupNM :: Ord1 l => ENode l -> NodeMap l a -> Maybe a+lookupNM :: Ord (l ClassId) => ENode l -> NodeMap l a -> Maybe a lookupNM e = M.lookup e . unNodeMap {-# INLINE lookupNM #-} -- | Delete an e-node in a 'NodeMap'-deleteNM :: Ord1 l => ENode l -> NodeMap l a -> NodeMap l a+deleteNM :: Ord (l ClassId) => ENode l -> NodeMap l a -> NodeMap l a deleteNM e (NodeMap m) = NodeMap (M.delete e m) {-# INLINE deleteNM #-} -- | Insert a value and lookup by e-node in a 'NodeMap'-insertLookupNM :: Ord1 l => ENode l -> a -> NodeMap l a -> (Maybe a, NodeMap l a)+insertLookupNM :: Ord (l ClassId) => ENode l -> a -> NodeMap l a -> (Maybe a, NodeMap l a) insertLookupNM e v (NodeMap m) = second NodeMap $ M.insertLookupWithKey (\_ a _ -> a) e v m {-# INLINE insertLookupNM #-} -- | As 'Data.Map.foldlWithKeyNM'' but in a 'NodeMap'-foldlWithKeyNM' :: Ord1 l => (b -> ENode l -> a -> b) -> b -> NodeMap l a -> b +foldlWithKeyNM' :: Ord (l ClassId) => (b -> ENode l -> a -> b) -> b -> NodeMap l a -> b foldlWithKeyNM' f b = M.foldlWithKey' f b . unNodeMap {-# INLINE foldlWithKeyNM' #-} -- | As 'Data.Map.foldrWithKeyNM'' but in a 'NodeMap'-foldrWithKeyNM' :: Ord1 l => (ENode l -> a -> b -> b) -> b -> NodeMap l a -> b +foldrWithKeyNM' :: Ord (l ClassId) => (ENode l -> a -> b -> b) -> b -> NodeMap l a -> b foldrWithKeyNM' f b = M.foldrWithKey' f b . unNodeMap {-# INLINE foldrWithKeyNM' #-}
src/Data/Equality/Language.hs view
@@ -1,4 +1,11 @@ {-# LANGUAGE FlexibleContexts #-}+{-# LANGUAGE UnicodeSyntax #-}+{-# LANGUAGE RankNTypes #-}+{-# LANGUAGE QuantifiedConstraints #-}+{-# LANGUAGE ConstraintKinds #-}+{-# LANGUAGE StandaloneKindSignatures #-}+{-# LANGUAGE FlexibleInstances #-}+{-# LANGUAGE UndecidableInstances #-} {-| Defines 'Language', which is the required constraint on /expressions/ that are@@ -28,7 +35,7 @@ -} module Data.Equality.Language where -import Data.Functor.Classes+import Data.Kind -- | A 'Language' is the required constraint on /expressions/ that are to be -- represented in an e-graph.@@ -38,5 +45,7 @@ -- e-graphs), note that it must satisfy the other class constraints. In -- particular an 'Data.Equality.Analysis.Analysis' must be defined for the -- language.-class (Traversable l, Ord1 l) => Language l where+type Language :: (Type -> Type) -> Constraint+class (∀ a. Ord a => Ord (l a), Traversable l) => Language l+instance (∀ a. Ord a => Ord (l a), Traversable l) => Language l
src/Data/Equality/Matching.hs view
@@ -1,5 +1,6 @@ {-# LANGUAGE ViewPatterns #-} {-# LANGUAGE ScopedTypeVariables #-}+{-# LANGUAGE MonoLocalBinds #-} {-| Equality-matching, implemented using a relational database (defined in 'Data.Equality.Matching.Database') according to the paper
src/Data/Equality/Matching/Database.hs view
@@ -1,9 +1,9 @@ {-# LANGUAGE ViewPatterns #-}+{-# LANGUAGE MonoLocalBinds #-} {-# LANGUAGE BangPatterns #-} {-# LANGUAGE OverloadedLists #-} {-# LANGUAGE ScopedTypeVariables #-} {-# LANGUAGE LambdaCase #-}-{-# LANGUAGE TupleSections #-} {-| Custom database implemented with trie-maps specialized to run conjunctive queries using a (worst-case optimal) generic join algorithm.
src/Data/Equality/Matching/Pattern.hs view
@@ -1,10 +1,10 @@+{-# LANGUAGE QuantifiedConstraints, RankNTypes, UnicodeSyntax, UndecidableInstances #-} {-| Definition of 'Pattern' for use in equality matching ('Data.Equality.Matching'), where patterns are matched against the e-graph -} module Data.Equality.Matching.Pattern where -import Data.Functor.Classes import Data.String import Data.Equality.Utils@@ -77,20 +77,20 @@ pat :: lang (Pattern lang) -> Pattern lang pat = NonVariablePattern -instance Eq1 l => (Eq (Pattern l)) where- (==) (NonVariablePattern a) (NonVariablePattern b) = liftEq (==) a b+instance (∀ a. Eq a => Eq (l a)) => (Eq (Pattern l)) where+ (==) (NonVariablePattern a) (NonVariablePattern b) = (==) a b (==) (VariablePattern a) (VariablePattern b) = a == b (==) _ _ = False -instance Ord1 l => (Ord (Pattern l)) where+instance (∀ a. Eq a => Eq (l a), ∀ a. (Ord a) => Ord (l a)) => (Ord (Pattern l)) where compare (VariablePattern _) (NonVariablePattern _) = LT compare (NonVariablePattern _) (VariablePattern _) = GT compare (VariablePattern a) (VariablePattern b) = compare a b- compare (NonVariablePattern a) (NonVariablePattern b) = liftCompare compare a b+ compare (NonVariablePattern a) (NonVariablePattern b) = compare a b -instance Show1 lang => Show (Pattern lang) where+instance (∀ a. Show a => Show (lang a)) => Show (Pattern lang) where showsPrec _ (VariablePattern s) = showString (show s) -- ROMES:TODO don't ignore prec?- showsPrec d (NonVariablePattern x) = liftShowsPrec showsPrec showList d x+ showsPrec d (NonVariablePattern x) = showsPrec d x instance IsString (Pattern lang) where fromString = VariablePattern . hashString
src/Data/Equality/Saturation.hs view
@@ -5,6 +5,7 @@ {-# LANGUAGE LambdaCase #-} {-# LANGUAGE TupleSections #-} {-# LANGUAGE BlockArguments #-}+{-# LANGUAGE MonoLocalBinds #-} {-| Given an input program 𝑝, equality saturation constructs an e-graph 𝐸 that represents a large set of programs equivalent to 𝑝, and then extracts the@@ -51,6 +52,7 @@ import Data.Equality.Utils import Data.Equality.Graph.Nodes import Data.Equality.Graph.Lens+import Data.Equality.Graph.Internal (EGraph(classes)) import qualified Data.Equality.Graph as G import Data.Equality.Graph.Monad import Data.Equality.Language@@ -115,7 +117,7 @@ egr <- get - let (beforeMemo, beforeClasses) = (egr^._memo, egr^._classes)+ let (beforeMemo, beforeClasses) = (egr^._memo, classes egr) db = eGraphToDatabase egr -- Read-only phase, invariants are preserved@@ -129,7 +131,7 @@ -- Restore the invariants once per iteration rebuild - (afterMemo, afterClasses) <- gets (\g -> (g^._memo, g^._classes))+ (afterMemo, afterClasses) <- gets (\g -> (g^._memo, classes g)) -- ROMES:TODO: Node limit... -- ROMES:TODO: Actual Timeout... not just iteration timeout
src/Data/Equality/Saturation/Rewrites.hs view
@@ -1,3 +1,4 @@+{-# LANGUAGE QuantifiedConstraints, RankNTypes, UnicodeSyntax #-} {-| Definition of 'Rewrite' and 'RewriteCondition' used to define rewrite rules.@@ -8,8 +9,6 @@ -} module Data.Equality.Saturation.Rewrites where -import Data.Functor.Classes- import Data.Equality.Graph import Data.Equality.Matching import Data.Equality.Matching.Database@@ -53,6 +52,6 @@ type RewriteCondition anl lang = Subst -> EGraph anl lang -> Bool -instance Show1 lang => Show (Rewrite anl lang) where+instance (∀ a. Show a => Show (lang a)) => Show (Rewrite anl lang) where show (rw :| _) = show rw <> " :| <cond>" show (lhs := rhs) = show lhs <> " := " <> show rhs
src/Data/Equality/Utils.hs view
@@ -1,4 +1,4 @@-{-# LANGUAGE StandaloneDeriving #-}+{-# LANGUAGE UnicodeSyntax, RankNTypes, QuantifiedConstraints, UndecidableInstances #-} {-| Misc utilities used accross modules -}@@ -10,7 +10,6 @@ -- import qualified Data.Set as S -- import qualified Data.IntSet as IS-import Data.Functor.Classes -- | Fixed point newtype. --@@ -21,12 +20,12 @@ -- just e-graphs, but until I revert the decision we use this type. newtype Fix f = Fix { unFix :: f (Fix f) } -instance Eq1 f => Eq (Fix f) where- (==) (Fix a) (Fix b) = liftEq (==) a b+instance (∀ a. Eq a => Eq (f a)) => Eq (Fix f) where+ (==) (Fix a) (Fix b) = a == b {-# INLINE (==) #-} -instance Show1 f => Show (Fix f) where- showsPrec d (Fix f) = liftShowsPrec showsPrec showList d f+instance (∀ a. Show a => Show (f a)) => Show (Fix f) where+ showsPrec d (Fix f) = showsPrec d f {-# INLINE showsPrec #-} -- | Catamorphism
test/Invariants.hs view
@@ -14,17 +14,17 @@ import Test.Tasty import Test.Tasty.QuickCheck as QC hiding (classes) -import Data.Functor.Classes import Control.Monad import qualified Data.Containers.ListUtils as LU import qualified Data.Foldable as F import qualified Data.List as L-import qualified Data.Set as S+import qualified Data.IntSet as IS import qualified Data.IntMap.Strict as IM import Data.Equality.Graph.Monad as GM import Data.Equality.Graph.Lens+import Data.Equality.Graph.Internal (EGraph(classes)) import Data.Equality.Graph import Data.Equality.Extraction import Data.Equality.Saturation@@ -36,7 +36,7 @@ -- TODO: Use type level symbol to define the analysis type role SimpleExpr nominal newtype SimpleExpr l = SE (Expr l)- deriving (Functor, Foldable, Traversable, Show1, Eq1, Ord1, Language)+ deriving (Functor, Foldable, Traversable, Show, Eq, Ord) -- | When a rewrite of type "x":=c where x is a pattern variable and c is a -- constant is used in equality saturation of any expression, all e-classes@@ -45,7 +45,7 @@ patFoldAllClasses :: forall l. (Language l, Num (Pattern l)) => Fix l -> Integer -> Bool patFoldAllClasses expr i =- case IM.toList (eg^._classes) of+ case IM.toList (classes eg) of [_] -> True _ -> False where@@ -89,8 +89,8 @@ ematchSingletonVar v eg = let db = eGraphToDatabase eg- matches = S.fromList $ map matchClassId $ ematch db (VariablePattern v)- eclasses = S.fromList $ map fst $ IM.toList (eg^._classes)+ matches = IS.fromList $ map matchClassId $ ematch db (VariablePattern v)+ eclasses = IM.keysSet (classes eg) in matches == eclasses @@ -116,7 +116,7 @@ hashConsInvariant :: forall l. Language l => EGraph () l -> Bool hashConsInvariant eg =- all f (IM.toList (eg^._classes))+ allOf _iclasses f eg where -- e-node 𝑛 ∈ 𝑀 [𝑎] ⇐⇒ 𝐻 [canonicalize(𝑛)] = find(𝑎) f (i, EClass{eClassNodes=nodes}) = all g nodes
test/Lambda.hs view
@@ -1,12 +1,9 @@ {-# LANGUAGE LambdaCase #-} {-# LANGUAGE MultiParamTypeClasses #-} {-# LANGUAGE DataKinds #-}-{-# LANGUAGE DeriveGeneric #-} {-# LANGUAGE FlexibleInstances #-} {-# LANGUAGE OverloadedStrings #-}-{-# LANGUAGE ViewPatterns #-} {-# LANGUAGE TypeFamilies #-}-{-# LANGUAGE TemplateHaskell #-} {-# LANGUAGE DeriveTraversable #-} module Lambda where @@ -22,10 +19,6 @@ import Control.Applicative ((<|>)) -import Data.Eq.Deriving-import Data.Ord.Deriving-import Text.Show.Deriving- import Data.Equality.Graph import Data.Equality.Extraction import Data.Equality.Analysis@@ -49,14 +42,10 @@ | Lam a a | Let a a a | LFix a a- deriving ( Eq, Ord, Functor- , Foldable, Traversable+ deriving ( Eq, Ord, Show+ , Functor, Foldable, Traversable ) -deriveEq1 ''Lambda-deriveOrd1 ''Lambda-deriveShow1 ''Lambda- evalL :: Lambda (Maybe (Lambda ())) -> Maybe (Lambda ()) evalL = \case Bool n -> Just (Bool n)@@ -111,9 +100,6 @@ Subst a b c -> b <> a <> c joinA = (<>)---instance Language Lambda instance Num (Fix Lambda) where fromInteger = Fix . Num . fromInteger
test/SimpleSym.hs view
@@ -3,21 +3,15 @@ {-# LANGUAGE TypeFamilies #-} {-# LANGUAGE MultiParamTypeClasses #-} {-# LANGUAGE OverloadedStrings #-}-{-# LANGUAGE TemplateHaskell #-} {-# LANGUAGE DeriveTraversable #-} module SimpleSym where import Test.Tasty import Test.Tasty.HUnit -import Data.Eq.Deriving-import Data.Ord.Deriving-import Text.Show.Deriving- import Data.Equality.Utils import Data.Equality.Matching import Data.Equality.Saturation-import Data.Equality.Language import Data.Equality.Analysis import Data.Equality.Graph.Lens ((^.), _data) @@ -26,15 +20,9 @@ | a :+: a | a :*: a | a :/: a- deriving (Functor, Foldable, Traversable)+ deriving (Functor, Foldable, Traversable, Eq, Ord, Show) infix 6 :+: infix 7 :*:, :/:--deriveEq1 ''SymExpr-deriveOrd1 ''SymExpr-deriveShow1 ''SymExpr--instance Language SymExpr instance Analysis (Maybe Double) SymExpr where makeA = \case
test/Sym.hs view
@@ -1,6 +1,5 @@ {-# LANGUAGE MultiParamTypeClasses #-} {-# LANGUAGE TypeApplications #-}-{-# LANGUAGE TemplateHaskell #-} {-# LANGUAGE BangPatterns #-} {-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE DeriveTraversable #-}@@ -17,10 +16,6 @@ import Data.String import Data.Maybe (isJust) -import Data.Eq.Deriving-import Data.Ord.Deriving-import Text.Show.Deriving- import qualified Data.Foldable as F import Control.Applicative (liftA2)@@ -38,8 +33,8 @@ | Const !Double | UnOp !UOp !a | BinOp !BOp !a !a- deriving ( Eq, Ord, Functor- , Foldable, Traversable+ deriving ( Eq, Ord, Show+ , Functor, Foldable, Traversable ) data BOp = Add | Sub@@ -55,12 +50,6 @@ | Sqrt | Ln deriving (Eq, Ord, Show)--deriveEq1 ''Expr-deriveOrd1 ''Expr-deriveShow1 ''Expr--instance Language Expr instance IsString (Fix Expr) where fromString = Fix . Sym
test/T1.hs view
@@ -1,6 +1,5 @@ {-# language DeriveTraversable #-} {-# language LambdaCase #-}-{-# language TemplateHaskell #-} {-# LANGUAGE TypeFamilies #-} {-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE FlexibleInstances #-}@@ -8,9 +7,6 @@ module T1 (main) where import Test.Tasty.HUnit-import Data.Eq.Deriving-import Data.Ord.Deriving-import Text.Show.Deriving import Data.Equality.Graph import Data.Equality.Matching@@ -24,11 +20,7 @@ | MulF a a | DivF a a | LogF a- deriving (Functor, Foldable, Traversable)--deriveEq1 ''TreeF-deriveOrd1 ''TreeF-deriveShow1 ''TreeF+ deriving (Eq, Ord, Show, Functor, Foldable, Traversable) instance Num (Fix TreeF) where l + r = Fix $ AddF l r@@ -99,8 +91,6 @@ l ** r = undefined logBase l r = undefined--instance Language TreeF cost :: CostFunction TreeF Int cost = \case
test/T2.hs view
@@ -1,5 +1,4 @@ {-# LANGUAGE DeriveTraversable #-}-{-# LANGUAGE TemplateHaskell #-} {-# LANGUAGE TypeFamilies #-} {-# LANGUAGE OverloadedStrings #-} module T2 where@@ -9,9 +8,7 @@ import Prelude hiding (not) import Test.Tasty.HUnit-import Data.Deriving import Data.Equality.Matching-import Data.Equality.Language import Data.Equality.Extraction import Data.Equality.Saturation @@ -20,13 +17,7 @@ | Not a | ToElim a | Sym Int- deriving (Functor, Foldable, Traversable)--deriveEq1 ''Lang-deriveOrd1 ''Lang-deriveShow1 ''Lang--instance Language Lang+ deriving (Functor, Foldable, Traversable, Eq, Ord, Show) x, y :: Pattern Lang x = "x"