hegg 0.3.0.0 → 0.6.0.0
raw patch · 31 files changed
Files
- CHANGELOG.md +38/−0
- README.md +29/−34
- hegg.cabal +23/−12
- src/Data/Equality/Analysis.hs +32/−28
- src/Data/Equality/Analysis/Monadic.hs +68/−0
- src/Data/Equality/Extraction.hs +3/−2
- src/Data/Equality/Graph.hs +233/−47
- src/Data/Equality/Graph/Classes.hs +1/−3
- src/Data/Equality/Graph/Dot.hs +8/−10
- src/Data/Equality/Graph/Internal.hs +1/−3
- src/Data/Equality/Graph/Lens.hs +39/−10
- src/Data/Equality/Graph/Monad.hs +46/−1
- src/Data/Equality/Graph/Nodes.hs +24/−33
- src/Data/Equality/Language.hs +11/−2
- src/Data/Equality/Matching.hs +72/−23
- src/Data/Equality/Matching/Database.hs +65/−21
- src/Data/Equality/Matching/Pattern.hs +9/−12
- src/Data/Equality/Saturation.hs +29/−23
- src/Data/Equality/Saturation/Rewrites.hs +3/−4
- src/Data/Equality/Saturation/Scheduler.hs +2/−1
- src/Data/Equality/Utils.hs +8/−16
- test/Bench.hs +31/−6
- test/Invariants.hs +19/−25
- test/Lambda.hs +17/−30
- test/SimpleSym.hs +9/−17
- test/Sym.hs +34/−39
- test/T1.hs +1/−11
- test/T2.hs +1/−10
- test/T3.hs +47/−0
- test/T32.hs +57/−0
- test/Test.hs +4/−0
CHANGELOG.md view
@@ -2,6 +2,44 @@ ## Unreleased +## 0.6.0.0 -- 2024-07-13++* Fix a soundness bug that would cause equality saturation to be broken when+ `VariablePattern` was used explicitly with low numbers such as 1,2,3...+ This bug could also be triggered in the unlikely case when the hash of the+ given IsString variable instance collided with the Var Ids internally+ associated to NonVariablePatterns. Fixes #20 and #32.++## 0.5.0.0 -- 2023-10-31++* Change `'modifyA'` to instead operate over e-graphs, instead of being+ constrained to editing the e-class that prompted the modification.+ (Remember that the e-graph lenses in `'Data.Equality.Graph.Lens'` are the+ preferred way to edit the e-graph and the desired e-class (by id), and its+ data, etc...)++* Fix compilation of Data.Equality.Graph.Dot, the graphviz rendering backend+ (despite there being some usability bugs still) (by @BinderDavid)++* Dropped support for GHC 9.0 because of the QuantifiedConstraints bug (by @phadej)++* Add `AnalysisM`, a class for e-graph analysis that are only well-defined+ within a certain monadic context. Accordingly, we also add versions of the+ current e-graph transformation functions (such as `add` and `merge`) for+ analysis defined monadically (such as `addM` and `mergeM`).++* Add operation to create empty e-classes with explicit domain data+ (experimental, not sure whether this is something good to keep in the API)++## 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@@ -123,7 +117,7 @@ e-node has a constant value or otherwise return `Nothing`: ```hs-makeA :: SymExpr (Maybe Double) -> Maybe Int+makeA :: SymExpr (Maybe Double) -> Maybe Double makeA = \case Const x -> Just x Symbol _ -> Nothing@@ -150,19 +144,22 @@ Finally, `modifyA` describes how an e-class should (optionally) be modified according to the e-class data and what new language expressions are to be added to the e-class also w.r.t. the e-class data.-Its type is `EClass domain l -> (EClass domain l, [Fix l])`, where the argument-is the class to modify, the first element of the return tuple is the-(optionally) modified e-class and the second element is a list of the-expressions to represent and merge with this e-class.-For our example, if the e-class has a constant value associated to it, we want-to create a new e-class with that constant value and merge it to this e-class.+Its type is `ClassId -> EGraph domain l -> EGraph domain l`, where the first argument+is the id of the class to modify (the class which prompted the modification),+and then receives and returns an e-graph, in which the e-class has been+modified. For our example, if the e-class has a constant value associated to+it, we want to create a new e-class with that constant value and merge it to+this e-class. ```hs--- import Data.Equality.Graph.Lens ((^.), _data)-modifyA :: EClass (Maybe Double) SymExpr -> (EClass (Maybe Double) SymExpr, [Fix SymExpr])-modifyA c = case c^._data of- Nothing -> (c, [])- Just i -> (c, [Fix (Const i)])+-- import Data.Equality.Graph.Lens ((^.), _class, _data)+modifyA :: ClassId -> EGraph (Maybe Double) SymExpr -> EGraph (Maybe Double) SymExpr+modifyA c egr+ = case egr ^._class c._data of+ Nothing -> egr+ Just i ->+ let (c', egr') = represent (Fix (Const i)) egr+ in snd $ merge c c' egr' ``` Modify is a bit trickier than the other methods, but it allows our e-graph to@@ -187,15 +184,6 @@ joinA _ _ = () ``` -### Language, again--With this setup, we can now express that `SymExpr` forms a `Language` which we-can represent and manipulate in an e-graph by simply instancing it (there are no-additional functions to define).-```hs-instance Language SymExpr-```- ### Equality saturation Equality saturation is defined as the function@@ -315,3 +303,10 @@ open hegg-test.svg open hegg-test.eventlog.html ```++## Coverage++```+cabal test hegg-test --enable-coverage --enable-library-coverage+```+
hegg.cabal view
@@ -1,7 +1,7 @@ cabal-version: 2.4 name: hegg-version: 0.3.0.0-Tested-With: GHC ==9.4.2 || ==9.2.2 || ==9.0.2 || ==8.10.7+version: 0.6.0.0+Tested-With: GHC ==9.10.1 || ==9.8.4 || ==9.6.7 synopsis: Fast equality saturation in Haskell description: Fast equality saturation and equality graphs based on "egg:@@ -38,7 +38,7 @@ maintainer: Rodrigo Mesquita <rodrigo.m.mesquita@gmail.com> copyright: Copyright (C) 2022 Rodrigo Mesquita category: Data-extra-source-files: CHANGELOG.md+extra-doc-files: CHANGELOG.md README.md source-repository head@@ -46,6 +46,13 @@ location: https://github.com/alt-romes/hegg library+ other-extensions: BlockArguments StandaloneKindSignatures++ -- QuantifiedConstraints are quite broken with GHC-9.0+ -- Ord (Pattern l) instance triggers same issue like+ -- https://github.com/haskell/mtl/issues/138+ build-depends: base >=4.16 || <4.15+ ghc-options: -Wall -Wcompat -- -fno-prof-auto@@ -77,6 +84,7 @@ Data.Equality.Extraction, Data.Equality.Language, Data.Equality.Analysis,+ Data.Equality.Analysis.Monadic, Data.Equality.Saturation.Scheduler, Data.Equality.Saturation.Rewrites, Data.Equality.Utils,@@ -91,11 +99,13 @@ -- other-modules: -- LANGUAGE extensions used by modules in this package.- build-depends: base >= 4.4 && < 5,- transformers >= 0.4 && < 0.7,- containers >= 0.4 && < 0.7+ -- base-4.13, because foldMap' is used.+ build-depends: base >= 4.13 && < 5,+ transformers >= 0.4 && < 0.8,+ containers >= 0.4 && < 0.8 if flag(vizdot)- build-depends: graphviz >= 2999.6 && < 2999.7+ build-depends: graphviz >= 2999.20 && < 2999.21,+ text hs-source-dirs: src default-language: Haskell2010 @@ -108,13 +118,12 @@ hs-source-dirs: test main-is: Test.hs other-modules: Invariants, Sym, Lambda, SimpleSym,- T1, T2+ T1, T2, T3, T32 other-extensions: OverloadedStrings 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,12 +136,14 @@ type: exitcode-stdio-1.0 build-depends: base, hegg, containers,- deriving-compat, tasty, tasty-hunit, tasty-quickcheck,- tasty-bench >= 0.2 && < 0.4- ghc-options: -with-rtsopts=-A32m -threaded+ tasty-bench >= 0.2 && < 0.4,+ deepseq >= 1.4 && < 1.6+ ghc-options: -with-rtsopts=-A32m+ if impl(ghc >= 8.6)+ ghc-options: -fproc-alignment=64 Flag vizdot Description: Compile 'Data.Equality.Graph.Dot' module to visualize e-graphs
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@@ -29,8 +30,10 @@ import Data.Kind (Type) import Control.Arrow ((***)) -import Data.Equality.Utils+import Data.Function ((&))+import Data.Equality.Graph.Lens import Data.Equality.Language+import Data.Equality.Graph.Internal (EGraph) import Data.Equality.Graph.Classes -- | An e-class analysis with domain @domain@ defined for a language @l@.@@ -69,31 +72,29 @@ -- e-node to c. Modify should be idempotent if no other changes occur to -- the e-class, i.e., modify(modify(c)) = modify(c) --- -- The return value of the modify function is both the modified class and- -- the expressions (in their fixed-point form) to add to this class. We- -- can't manually add them because not only would it skip some of the- -- internal steps of representing + merging, but also because it's- -- impossible to add any expression with depth > 0 without access to the- -- e-graph (since we must represent every sub-expression in the e-graph- -- first).- --- -- That's why we must return the modified class and the expressions to add- -- to this class.- -- -- === Example -- -- Pruning an e-class with a constant value of all its nodes except for the -- leaf values, and adding a constant value node -- -- @- -- -- Prune all except leaf e-nodes- -- modifyA cl =- -- case cl^._data of- -- Nothing -> (cl, [])- -- Just d -> ((_nodes %~ S.filter (F.null .unNode)) cl, [Fix (Const d)])+ -- modifyA cl eg0+ -- = case eg0^._class cl._data of+ -- Nothing -> eg0+ -- Just d ->+ -- -- Add constant as e-node+ -- let (new_c,eg1) = represent (Fix (Const d)) eg0+ -- (rep, eg2) = merge cl new_c eg1+ -- -- Prune all except leaf e-nodes+ -- in eg2 & _class rep._nodes %~ S.filter (F.null .unNode) -- @- modifyA :: EClass domain l -> (EClass domain l, [Fix l])- modifyA c = (c, [])+ modifyA :: ClassId+ -- ^ Id of class @c@ whose new data @d_c@ triggered the modify call+ -> EGraph domain l+ -- ^ E-graph where class @c@ being modified exists+ -> EGraph domain l+ -- ^ E-graph resulting from the modification+ modifyA _ = id {-# INLINE modifyA #-} @@ -104,13 +105,14 @@ joinA = (<>) --- This instance is not necessarily well behaved for any two analysis, so care+-- | This instance is not necessarily well behaved for any two analysis, so care -- must be taken when using it. -- -- A possible criterion is: -- -- For any two analysis, where 'modifyA' is called @m1@ and @m2@ respectively,--- this instance is well behaved if @m1@ and @m2@ commute.+-- this instance is well behaved if @m1@ and @m2@ commute, and the analysis+-- only change the e-class being modified. -- -- That is, if @m1@ and @m2@ satisfy the following law: -- @@@ -132,10 +134,12 @@ joinA :: (a,b) -> (a,b) -> (a,b) joinA (x,y) = joinA @a @l x *** joinA @b @l y - modifyA :: EClass (a, b) l -> (EClass (a, b) l, [Fix l])- modifyA c =- let (ca, la) = modifyA @a (c { eClassData = fst (eClassData c) })- (cb, lb) = modifyA @b (c { eClassData = snd (eClassData c) })- in ( EClass (eClassId c) (eClassNodes ca <> eClassNodes cb) (eClassData ca, eClassData cb) (eClassParents ca <> eClassParents cb)- , la <> lb- )+ modifyA :: ClassId -> EGraph (a, b) l -> EGraph (a, b) l+ modifyA c egr =+ let egra = modifyA @a c (egr & _classes._data %~ fst)+ egrb = modifyA @b c (egr & _classes._data %~ snd)+ ca = egra ^._class c+ cb = egrb ^._class c+ in+ egr &+ _class c .~ (EClass c (eClassNodes ca <> eClassNodes cb) (eClassData ca, eClassData cb) (eClassParents ca <> eClassParents cb))
+ src/Data/Equality/Analysis/Monadic.hs view
@@ -0,0 +1,68 @@+{-# LANGUAGE AllowAmbiguousTypes #-} -- joinA+{-# LANGUAGE InstanceSigs #-}+{-# LANGUAGE TypeApplications #-}+{-# LANGUAGE ScopedTypeVariables #-}+{-# LANGUAGE FlexibleInstances #-}+{-# LANGUAGE FlexibleContexts #-}+{-# LANGUAGE TypeFamilies #-}+{-# LANGUAGE MultiParamTypeClasses #-}+{-# LANGUAGE ImpredicativeTypes #-}+{-|++Like 'Data.Equality.Analysis' but for 'Analysis' that are only well-defined+within an (effectful) context. Mostly used with the monadic operations+'representM', 'addM', 'mergeM', and 'rebuildM'.++This effectful 'Analysis' could almost be trivially defined in terms of the+other, through a "contextful" domain and by means of the '_classes' 'Traversal'.++However, that would require an instance of 'Eq' for the monadic domain, which+is usually unnattainable.++Therefore, we do need this class for monadic 'Analysis'.++-}+module Data.Equality.Analysis.Monadic where+++import Data.Kind (Type)++import Data.Equality.Graph.Internal (EGraph)+import Data.Equality.Graph.Classes++-- | An e-class analysis with domain @domain@ defined for a language @l@, whose operations are only well-defined within some effectful context.+--+-- The @domain@ is the type of the domain of the e-class analysis, that is, the+-- type of the data stored in an e-class according to this e-class analysis+class (Monad m, Eq domain) => AnalysisM m domain (l :: Type -> Type) where++ -- | When a new e-node is added into a new, singleton e-class, construct a+ -- new value of the domain to be associated with the new e-class, by+ -- accessing the associated data of the node's children+ --+ -- The argument is the e-node term populated with its children data+ makeA :: l domain -> m domain++ -- | When e-classes c1 c2 are being merged into c, join d_c1 and+ -- d_c2 into a new value d_c to be associated with the new+ -- e-class c+ joinA :: domain -> domain -> m domain++ -- | Optionally modify the e-class c (based on d_c), typically by adding an+ -- e-node to c. Modify should be idempotent if no other changes occur to+ -- the e-class, i.e., modify(modify(c)) = modify(c)+ modifyA :: ClassId+ -- ^ Id of class @c@ whose new data @d_c@ triggered the modify call+ -> EGraph domain l+ -- ^ E-graph where class @c@ being modified exists+ -> m (EGraph domain l)+ -- ^ E-graph resulting from the modification+ modifyA _ = pure+ {-# INLINE modifyA #-}+++-- | The simplest analysis that defines the domain to be () and does nothing otherwise+instance Monad m => AnalysisM m () l where+ makeA _ = pure ()+ joinA _ _ = pure ()+
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,8 @@ {-# LANGUAGE FlexibleContexts #-} {-# LANGUAGE ScopedTypeVariables #-} {-# LANGUAGE TypeFamilies #-}+{-# LANGUAGE TupleSections #-}+{-# OPTIONS_GHC -ddump-to-file -ddump-simpl #-} {-| An e-graph efficiently represents a congruence relation over many expressions. @@ -12,32 +14,45 @@ -} module Data.Equality.Graph (- -- * Definition of e-graph+ -- ** Definition of e-graph EGraph - -- * Functions on e-graphs- , emptyEGraph-- -- ** Transformations+ -- ** E-graph transformations , represent, add, merge, rebuild -- , repair, repairAnal -- ** Querying , find, canonicalize + -- ** Functions on e-graphs+ , emptyEGraph, newEClass++ -- * E-graph transformations for monadic analysis+ -- | These are the same operations over e-graphs as above but over a monad in which the analysis is defined.+ -- It is common to only have a valid 'Analysis' under a monadic context.+ -- In that case, these are the functions to use -- they are just like the+ -- non-monadic ones, but have require an 'Analysis' defined in a+ -- monadic context ('AnalysisM').+ , representM, addM, mergeM, rebuildM+ -- * Re-exports , module Data.Equality.Graph.Classes , module Data.Equality.Graph.Nodes , module Data.Equality.Language ) where +-- ROMES:TODO: Is the E-Graph a Monad if the analysis data were the type arg? i.e. instance Monad (EGraph language)?+ -- import GHC.Conc+import Prelude hiding (lookup) import Data.Function+import Data.Foldable (foldlM) import Data.Bifunctor import Data.Containers.ListUtils import Control.Monad+import Control.Monad.Trans.Class import Control.Monad.Trans.State import Control.Exception (assert) @@ -51,6 +66,7 @@ import Data.Equality.Graph.Classes import Data.Equality.Graph.Nodes import Data.Equality.Analysis+import qualified Data.Equality.Analysis.Monadic as AM import Data.Equality.Language import Data.Equality.Graph.Lens @@ -85,12 +101,7 @@ (new_eclass_id, new_uf) = makeNewSet (unionFind egr) -- New singleton e-class stores the e-node and its analysis data- -- which is modified according to analysis- --- -- The modification also produces a list of expressions to- -- represent and merge with this class, which we'll do before- -- returning from this function- (new_eclass, added_nodes) = modifyA $ EClass new_eclass_id (S.singleton new_en) (makeA @a ((\i -> egr^._class i._data @a) <$> unNode new_en)) mempty+ new_eclass = EClass new_eclass_id (S.singleton new_en) (makeA @a ((\i -> egr^._class i._data @a) <$> unNode new_en)) mempty -- TODO:Performance: All updates can be done to the map first? Parallelize? --@@ -134,20 +145,14 @@ -- Add the e-node's e-class id at the e-node's id new_memo = insertNM new_en new_eclass_id (memo egr) - -- So we have our almost final e-graph. We just need to represent- -- and merge in it all expressions which resulted from 'modifyA'- -- above- egr1 = egr { unionFind = new_uf- , classes = new_classes- , worklist = new_worklist- , memo = new_memo- }-- egr2 = foldr (representAndMerge new_eclass_id) egr1 added_nodes-- in ( new_eclass_id- , egr2+ , egr { unionFind = new_uf+ , classes = new_classes+ , worklist = new_worklist+ , memo = new_memo+ }+ -- Modify created node according to analysis+ & modifyA new_eclass_id ) {-# INLINABLE add #-} @@ -180,11 +185,10 @@ -- Update leader class with all e-nodes and parents from the -- subsumed class- (updatedLeader, added_nodes) = leader_class- & _parents %~ (sub_class^._parents <>)- & _nodes %~ (sub_class^._nodes <>)- & _data .~ new_data- & modifyA+ updatedLeader = leader_class+ & _parents %~ (sub_class^._parents <>)+ & _nodes %~ (sub_class^._nodes <>)+ & _data .~ new_data new_data = joinA @a @l (leader_class^._data) (sub_class^._data) @@ -222,10 +226,10 @@ , worklist = new_worklist , analysisWorklist = new_analysis_worklist }-- egr2 = foldr (representAndMerge leader) egr1 added_nodes+ -- Modify according to analysis+ & modifyA new_id - in (new_id, egr2)+ in (new_id, egr1) {-# INLINEABLE merge #-} @@ -272,22 +276,18 @@ 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- 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)+ c = egr^._class repair_id+ new_data = joinA @a @l (c^._data) (makeA @a ((\i -> egr^._class i^._data @a) <$> unNode node)) in -- Take action if the new_data is different from the existing data- if c1^._data /= new_data+ if c^._data /= new_data -- Merge result is different from original class data, update class -- with new_data then- let- new_classes = IM.insert repair_id c2 (classes egr)- egr1 = egr { analysisWorklist = toListSL (c1^._parents) <> analysisWorklist egr- , classes = new_classes- }- egr2 = foldr (representAndMerge repair_id) egr1 added_nodes- in egr2+ egr { analysisWorklist = toListSL (c^._parents) <> analysisWorklist egr+ , classes = IM.adjust (_data .~ new_data) repair_id (classes egr)+ }+ & modifyA repair_id else egr {-# INLINE repairAnal #-} @@ -304,6 +304,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@@ -314,7 +315,192 @@ emptyEGraph = EGraph emptyUF mempty mempty mempty mempty {-# INLINE emptyEGraph #-} --- | Represent an expression (in fix-point form) and merge it with the e-class with the given id-representAndMerge :: (Analysis a l, Language l) => ClassId -> Fix l -> EGraph a l -> EGraph a l-representAndMerge o f g = case represent f g of- (i, e) -> snd $ merge o i e+-- | Creates an empty e-class in an e-graph, with the explicitly given domain analysis data.+-- (That is, an e-class with no e-nodes)+newEClass :: (Language l) => a -> EGraph a l -> (ClassId, EGraph a l)+newEClass adata egr =+ let+ -- Make new equivalence class with a new id in the union-find+ (new_eclass_id, new_uf) = makeNewSet (unionFind egr)++ -- New empty e-class stores just the analysis data+ new_eclass = EClass new_eclass_id S.empty adata mempty+ in ( new_eclass_id+ , egr { unionFind = new_uf+ , classes = IM.insert new_eclass_id new_eclass (classes egr)+ }+ )+{-# INLINE newEClass #-}++----- Monadic operations on e-graphs+-- Unfortunately, these cannot be defined in terms of the primary ones.+-- This could almost be done by defining the domain of the standard Analysis to+-- be (m a), for some Monad m, but this would require an instance Eq (m a),+-- which often won't exist.+--+-- Anyway, this allows us to have a better story for monadic analysis because+-- the type-class functions arguments don't need to be of the same monadic type+-- as the result (unlike if we were using a normal analysis with a monadic domain).+--+-- Be sure to update these functions if the above "canonical" versions are changed.++-- TODO: Move these to new module?++-- * E-graph operations for analysis defined monadically ('AM.AnalysisM')++-- | Like 'represent', but for a monadic analysis+representM :: forall a l m. (AM.AnalysisM m a l, Language l) => Fix l -> EGraph a l -> m (ClassId, EGraph a l)+representM = cata $ \l e -> do+ -- Canonical implementation is represent, this is just the monadic variant of it+ (l', e') <- (`runStateT` e) $ traverse (\f -> get >>= lift . f >>= StateT . const . pure) l+ addM (Node l') e'++-- | Like 'add', but for a monadic analysis+addM :: forall a l m. (AM.AnalysisM m a l, Language l) => ENode l -> EGraph a l -> m (ClassId, EGraph a l)+addM uncanon_e egr =+ -- Canonical implementation is add, this is just the monadic variant of it+ let !new_en = canonicalize uncanon_e egr++ in case lookupNM new_en (memo egr) of+ Just canon_enode_id -> pure (find canon_enode_id egr, egr)+ Nothing -> do+ let+ (new_eclass_id, new_uf) = makeNewSet (unionFind egr)++ new_data <- AM.makeA @m @a ((\i -> egr^._class i._data @a) <$> unNode new_en)++ let+ new_eclass = EClass new_eclass_id (S.singleton new_en) new_data mempty+ new_parents = ((new_eclass_id, new_en) |:)+ new_classes = IM.insert new_eclass_id new_eclass $+ foldr (IM.adjust (_parents %~ new_parents))+ (classes egr)+ (unNode new_en)++ new_worklist = (new_eclass_id, new_en):worklist egr++ new_memo = insertNM new_en new_eclass_id (memo egr)++ egr1 <- egr { unionFind = new_uf+ , classes = new_classes+ , worklist = new_worklist+ , memo = new_memo+ }+ & AM.modifyA new_eclass_id++ return ( new_eclass_id, egr1 )+{-# INLINABLE addM #-}++-- | Like 'merge', but for a monadic analysis+mergeM :: forall a l m. (AM.AnalysisM m a l, Language l) => ClassId -> ClassId -> EGraph a l -> m (ClassId, EGraph a l)+mergeM a b egr0 = do+ -- Canonical implementation is merge, this is just the monadic variant of it+ let+ a' = find a egr0+ b' = find b egr0+ in+ if a' == b'+ then return (a', egr0)+ else do+ let+ class_a = egr0 ^._class a'+ class_b = egr0 ^._class b'++ (leader, leader_class, sub, sub_class) =+ if sizeSL (class_a^._parents) < sizeSL (class_b^._parents)+ then (b', class_b, a', class_a) -- b is leader+ else (a', class_a, b', class_b) -- a is leader++ (new_id, new_uf) = unionSets leader sub (unionFind egr0)+ & first (\n -> assert (leader == n) n)++ new_data <- AM.joinA @m @a @l (leader_class^._data) (sub_class^._data)++ let+ updatedLeader = leader_class+ & _parents %~ (sub_class^._parents <>)+ & _nodes %~ (sub_class^._nodes <>)+ & _data .~ new_data++ new_classes = (IM.insert leader updatedLeader . IM.delete sub) (classes egr0)++ -- Add all subsumed parents to worklist We can do this instead of+ -- adding the new e-class itself to the worklist because it would end+ -- up adding its parents anyway+ new_worklist = toListSL (sub_class^._parents) <> worklist egr0++ -- If the new_data is different from the classes, the parents of the+ -- class whose data is different from the merged must be put on the+ -- analysisWorklist+ new_analysis_worklist =+ (if new_data /= (sub_class^._data)+ then toListSL (sub_class^._parents)+ else mempty) <>+ (if new_data /= (leader_class^._data)+ then toListSL (leader_class^._parents)+ else mempty) <>+ analysisWorklist egr0++ -- ROMES:TODO: The code that makes the -1 * cos test pass when some other things are tweaked+ -- new_memo = foldr (`insertNM` leader) (memo egr0) (sub_class^._nodes)++ -- Build new e-graph+ egr1 <- egr0 { unionFind = new_uf+ , classes = new_classes+ -- , memo = new_memo+ , worklist = new_worklist+ , analysisWorklist = new_analysis_worklist+ }+ & AM.modifyA new_id++ return (new_id, egr1)+{-# INLINEABLE mergeM #-}++-- | Like 'rebuild', but for a monadic analysis+rebuildM :: forall a l m. (AM.AnalysisM m a l, Language l) => EGraph a l -> m (EGraph a l)+rebuildM (EGraph uf cls mm wl awl) = do+ -- Canonical implementation is rebuild, this is just the monadic variant of it+ let+ emptiedEgr = EGraph uf cls mm mempty mempty++ wl' = nubOrd $ bimap (`find` emptiedEgr) (`canonicalize` emptiedEgr) <$> wl++ egr' <- foldlM (flip repairM) emptiedEgr wl'++ let awl' = nubIntOn fst $ first (`find` egr') <$> awl++ egr'' <- foldlM (flip repairAnalM) egr' awl'++ if null (worklist egr'') && null (analysisWorklist egr'')+ then return egr''+ else rebuildM egr''+{-# INLINEABLE rebuildM #-}++-- | Like 'repair', but for a monadic analysis+repairM :: forall a l m. (AM.AnalysisM m a l, Language l) => (ClassId, ENode l) -> EGraph a l -> m (EGraph a l)+repairM (repair_id, node) egr =+ -- Canonical implementation is repair, this is just the monadic variant of it+ case insertLookupNM node repair_id (memo egr) of++ (Nothing, memo') -> return $ egr { memo = memo' }++ (Just existing_class, memo') -> snd <$> (mergeM existing_class repair_id egr{memo = memo'})+{-# INLINE repairM #-}++-- | Like 'repairAnal', but for a monadic analysis+repairAnalM :: forall a l m. (AM.AnalysisM m a l, Language l) => (ClassId, ENode l) -> EGraph a l -> m (EGraph a l)+repairAnalM (repair_id, node) egr = do+ -- Canonical implementation is repairAnal, this is just the monadic variant of it+ let c = egr^._class repair_id++ new_data <- AM.joinA @m @a @l (c^._data) =<< AM.makeA @m @a ((\i -> egr^._class i^._data @a) <$> unNode node)++ if c^._data /= new_data+ then+ egr { analysisWorklist = toListSL (c^._parents) <> analysisWorklist egr+ , classes = IM.adjust (_data .~ new_data) repair_id (classes egr)+ }+ & AM.modifyA repair_id+ else+ return egr+{-# INLINE repairAnalM #-}
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/Dot.hs view
@@ -23,22 +23,21 @@ import Data.GraphViz.Attributes (style, dotted, textLabel) import Data.GraphViz.Attributes.Complete -import Data.Equality.Saturation import Data.Equality.Graph-import Data.Equality.Matching-import Database+import Data.Equality.Graph.Internal +txt :: Show a => a -> Text txt = pack . show -writeDemo :: (Functor f, Foldable f, Show (ENode f)) => EGraph f -> IO ()+writeDemo :: (Language language, Show (ENode language)) => EGraph analysis language -> IO () writeDemo = writeDotFile "demo.gv" . toDotGraph -toDotGraph :: (Functor f, Foldable f, Show (ENode f)) => EGraph f -> DotGraph Text+toDotGraph :: (Language language, Show (ENode language)) => EGraph analysis language -> DotGraph Text toDotGraph eg = digraph (Str "egraph") $ do graphAttrs [Compound True, ClusterRank Local] - forM_ (IM.toList $ classes eg) $ \(class_id, EClass _ nodes parents) ->+ forM_ (IM.toList $ classes eg) $ \(class_id, EClass _ nodes _ _) -> subgraph (Str ("cluster_" <> txt class_id)) $ do graphAttrs [style dotted]@@ -46,17 +45,16 @@ let n' = canonicalize n eg node (txt class_id <> "." <> txt (find i eg)) [textLabel (txt n')] - forM_ (IM.toList $ classes eg) $ \(class_id, EClass _ nodes parents) -> do+ forM_ (IM.toList $ classes eg) $ \(class_id, EClass _ nodes _ _) -> do forM_ (zip (S.toList nodes) [0..]) $ \(n, i_in_class) -> do let n' = canonicalize n eg let i_in_class' = find i_in_class eg - forM_ (zip (children n') [0..]) $ \(child, arg_i) -> do+ forM_ (zip (children n') [(0 :: Integer)..]) $ \(child, arg_i) -> do -- TODO: On anchors and labels...? let child_leader = find child eg if child_leader == class_id- then edge (txt class_id <> "." <> txt i_in_class') (txt class_id <> "." <> txt i_in_class') [textLabel (txt arg_i)] -- LHead ("cluster_" <> txt class_id), + then edge (txt class_id <> "." <> txt i_in_class') (txt class_id <> "." <> txt i_in_class') [textLabel (txt arg_i)] -- LHead ("cluster_" <> txt class_id), else edge (txt class_id <> "." <> txt i_in_class') (txt child <> ".0") [LHead ("cluster_" <> txt child_leader), textLabel (txt arg_i)]-
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@@ -21,9 +22,12 @@ import Data.Equality.Graph.Classes import Data.Equality.Graph.ReprUnionFind --- | A 'Lens'' as defined in other lenses libraries+-- | A 'Lens'' as defined in lens type Lens' s a = forall f. Functor f => (a -> f a) -> (s -> f s)-+-- | A 'Lens' as defined in lens+type Lens s t a b = forall f. Functor f => (a -> f b) -> (s -> f t)+-- | A 'Traversal' as defined in lens+type Traversal s t a b = forall f. Applicative f => (a -> f b) -> (s -> f t) -- outdated comment for "getClass": --@@ -35,7 +39,7 @@ -- -- Invariant: The e-class exists. --- | Lens for the e-class with the given id in an e-graph+-- | Lens for the e-class at the representative of the given id in an e-graph -- -- Calls 'error' when the e-class doesn't exist _class :: ClassId -> Lens' (EGraph a l) (EClass a l)@@ -50,14 +54,19 @@ _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+_data :: Lens (EClass domain l) (EClass domain' l) domain domain'+_data afb EClass{..} = (\d1 -> EClass eClassId eClassNodes d1 eClassParents) <$> afb eClassData {-# INLINE _data #-} -- | Lens for the parent e-classes of an e-class@@ -83,7 +92,7 @@ {-# INLINE (.~) #-} -- | Synonym for @'over'@-(%~) :: Lens' s a -> (a -> a) -> (s -> s)+(%~) :: ASetter s t a b -> (a -> b) -> (s -> t) (%~) = over infixr 4 %~ {-# INLINE (%~) #-}@@ -99,7 +108,27 @@ {-# INLINE set #-} -- | Applies a function to the target-over :: Lens' s a -> (a -> a) -> (s -> s)+over :: ASetter s t a b -> (a -> b) -> (s -> t) over ln f = runIdentity . ln (Identity . f) {-# INLINE over #-} +-- | Basically 'traverse' over a 'Traversal'+traverseOf :: Traversal s t a b -> forall f. Applicative f => (a -> f b) -> s -> f t +traverseOf t = t+{-# INLINE traverseOf #-}++-- | 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 #-}++-- * Utilities++-- We need to use 'ASetter' instead of 'Lens' in %~ to ensure type inference can+-- figure out the Functor or Applicative is 'Identity'. Otherwise, we won't be+-- able to use the 'Traversal' to modify something through a 'Lens'.++-- | Used instead of 'Lens' in 'over' and '%~' to ensure one can call those+-- combinators on 'Lens's and 'Traversal's. Essentially, it helps type+-- inference in such function applications+type ASetter s t a b = (a -> Identity b) -> s -> Identity t
src/Data/Equality/Graph/Monad.hs view
@@ -1,9 +1,15 @@ {-# LANGUAGE TupleSections #-}+{-# LANGUAGE MonoLocalBinds #-} {-| Monadic interface to e-graph stateful computations -} module Data.Equality.Graph.Monad- ( egraph+ (+ -- * Threading e-graphs in a stateful computation+ --+ -- | These are the same operations over e-graphs as in 'Data.Equality.Graph',+ -- but defined in the context of a 'State' monad threading around the e-graph.+ egraph , represent , add , merge@@ -12,9 +18,17 @@ , EG.find , EG.emptyEGraph + -- * E-graph transformations for monadic analysis+ --+ -- | The same e-graph operations in a stateful computation threading around+ -- the e-graph, but for 'Analysis' defined monadically ('AnalysisM').+ , representM, addM, mergeM, rebuildM+ -- * E-graph stateful computations , EGraphM+ , EGraphMT , runEGraphM+ , runEGraphMT -- * E-graph definition re-export , EG.EGraph@@ -29,11 +43,14 @@ import Data.Equality.Utils (Fix, cata) import Data.Equality.Analysis+import qualified Data.Equality.Analysis.Monadic as AM import Data.Equality.Graph (EGraph, ClassId, Language, ENode(..)) import qualified Data.Equality.Graph as EG -- | E-graph stateful computation type EGraphM a l = State (EGraph a l)+-- | E-graph stateful computation over an arbitrary monad+type EGraphMT a l = StateT (EGraph a l) -- | Run EGraph computation on an empty e-graph --@@ -82,3 +99,31 @@ runEGraphM :: EGraph anl l -> EGraphM anl l a -> (a, EGraph anl l) runEGraphM = flip runState {-# INLINE runEGraphM #-}++--------------------------------------------------------------------------------+-- Monadic Analysis interface++-- | Run 'EGraphM' computation on a given e-graph over a monadic analysis+runEGraphMT :: EGraph anl l -> EGraphMT anl l m a -> m (a, EGraph anl l)+runEGraphMT = flip runStateT+{-# INLINE runEGraphMT #-}++-- | Like 'represent', but for a monadic analysis+representM :: (AM.AnalysisM m anl l, Language l) => Fix l -> EGraphMT anl l m ClassId+representM = StateT . EG.representM+{-# INLINE representM #-}++-- | Like 'add', but for a monadic analysis+addM :: (AM.AnalysisM m anl l, Language l) => ENode l -> EGraphMT anl l m ClassId+addM = StateT . EG.addM+{-# INLINE addM #-}++-- | Like 'merge', but for a monadic analysis+mergeM :: (AM.AnalysisM m anl l, Language l) => ClassId -> ClassId -> EGraphMT anl l m ClassId+mergeM a b = StateT (EG.mergeM a b)+{-# INLINE mergeM #-}++-- | Like 'rebuild', but for a monadic analysis+rebuildM :: (AM.AnalysisM m anl l, Language l) => EGraphMT anl l m ()+rebuildM = StateT (fmap ((),) . EG.rebuildM)+{-# INLINE rebuildM #-}
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@@ -9,7 +10,11 @@ ( ematch , eGraphToDatabase , Match(..)++ -- * Compiling a Pattern to a Query , compileToQuery+ , VarsState(varNames), findVarName+ , userPatVars , module Data.Equality.Matching.Pattern )@@ -30,6 +35,7 @@ import Data.Equality.Graph.Lens import Data.Equality.Matching.Database import Data.Equality.Matching.Pattern+import Data.Coerce (coerce) -- | Matching a pattern on an e-graph returns the e-class in which the pattern -- was matched and an e-class substitution for every 'VariablePattern' in the pattern.@@ -40,7 +46,8 @@ -- TODO: Perhaps e-graph could carry database and rebuild it on rebuild --- | Match a pattern against a 'Database', which can be gotten from an 'EGraph' with 'eGraphToDatabase'+-- | Match a pattern 'Query', gotten from 'compileToQuery' on a 'Pattern',+-- against a 'Database', which is built from an 'EGraph' with 'eGraphToDatabase'. -- -- Returns a list of matches, one 'Match' for each set of valid substitutions -- for all variables and the equivalence class in which the pattern was matched.@@ -49,21 +56,20 @@ -- could be constructed only once and shared accross matching. ematch :: Language l => Database l- -> Pattern l+ -> (Query l, Var {- query root var -}) -> [Match]-ematch db patr =+ematch db (q, root) = let- (q, root) = compileToQuery patr- -- | Convert each substitution into a match by getting the class-id -- where we matched from the subst -- -- If the substitution is empty there is no match f :: Subst -> Maybe Match- f s = if IM.null s then Nothing- else case IM.lookup root s of- Nothing -> error "how is root not in map?"- Just found -> pure (Match s found)+ f s = if nullSubst s+ then Nothing+ else case lookupSubst root s of+ Nothing -> error "how is root not in map?"+ Just found -> pure (Match s found) in mapMaybe f (genericJoin db q) @@ -96,24 +102,28 @@ -- | Auxiliary result in 'compileToQuery' algorithm data AuxResult lang = {-# UNPACK #-} !Var :~ [Atom lang] + -- | Compiles a 'Pattern' to a 'Query' and returns the query root variable with -- it. -- The root variable's substitutions are the e-classes where the pattern -- matched-compileToQuery :: (Traversable lang) => Pattern lang -> (Query lang, Var)-compileToQuery (VariablePattern x) = (SelectAllQuery x, x)-compileToQuery pa@(NonVariablePattern _) =+compileToQuery :: (Traversable lang) => Pattern lang -> ((Query lang, Var), VarsState)+compileToQuery (VariablePattern n) = flip runState emptyVarsState $ do+ v <- getVarName n+ return (SelectAllQuery v, v)+compileToQuery pa = - let root :~ atoms = evalState (aux pa) 0- in (Query (nubInt $ root:vars pa) atoms, root)+ let (root :~ atoms, varsState) = runState (aux pa) emptyVarsState+ in ((Query (nubVars $ root:userPatVars varsState) atoms, root), varsState) where - aux :: (Traversable lang) => Pattern lang -> State Int (AuxResult lang)- aux (VariablePattern x) = return (x :~ []) -- from definition in relational e-matching paper (needed for as base case for recursion)+ aux :: (Traversable lang) => Pattern lang -> State VarsState (AuxResult lang)+ aux (VariablePattern x) = do+ v <- getVarName x+ return (v :~ []) -- from definition in relational e-matching paper (needed for as base case for recursion) aux (NonVariablePattern p) = do- v <- get- modify' (+1)+ v <- nextVar (toList -> auxs) <- traverse aux p let boundVars = map (\(b :~ _) -> b) auxs atoms = join $ map (\(_ :~ a) -> a) auxs@@ -127,9 +137,48 @@ -- taking from the bound vars array subPatsToVars :: Traversable lang => lang (Pattern lang) -> [Var] -> State Int (lang Var) subPatsToVars p' boundVars = traverse (const $ (boundVars !!) <$> (get >>= \i -> modify' (+1) >> return i)) p'-- -- | Return distinct variables in a pattern- vars :: Foldable lang => Pattern lang -> [Var]- vars (VariablePattern x) = [x]- vars (NonVariablePattern p) = nubInt $ join $ map vars $ toList p {-# INLINABLE compileToQuery #-}++--------------------------------------------------------------------------------+-- ** Vars utils for compileToQuery+--------------------------------------------------------------------------------++-- | Map user-given Variable names to internal 'Var's plus a counter+data VarsState = VarsState+ { varNames :: !(M.Map String Var)+ , nextVarId :: !Int+ }++-- | An empty 'VarsState'+emptyVarsState :: VarsState+emptyVarsState = VarsState mempty 0++-- | Compute the next internal 'Var' from the current 'VarNameMap'+nextVar :: State VarsState Var+nextVar = do+ n <- gets nextVarId+ modify' (\vs -> vs{nextVarId = nextVarId vs +1})+ return (MatchVar n)++-- | Add a name to the 'VarNameMap' and get the resulting 'Var'.+getVarName :: String -> State VarsState Var+getVarName s = do+ vm <- gets varNames+ case M.lookup s vm of+ Nothing -> do+ n <- nextVar+ modify' (\vs -> vs{varNames = M.insert s n (varNames vs)})+ return n+ Just v ->+ return v++findVarName :: VarsState -> String -> Var+findVarName vs s = (varNames vs) M.! s++-- | Return the variables given in a pattern by a user, from the 'VarsState'+userPatVars :: VarsState -> [Var]+userPatVars = M.elems . varNames++-- | Deduplicate list of pattern Vars+nubVars :: [Var] -> [Var]+nubVars = coerce . nubInt . coerce
src/Data/Equality/Matching/Database.hs view
@@ -1,9 +1,10 @@ {-# LANGUAGE ViewPatterns #-}+{-# LANGUAGE MonoLocalBinds #-} {-# LANGUAGE BangPatterns #-} {-# LANGUAGE OverloadedLists #-} {-# LANGUAGE ScopedTypeVariables #-} {-# LANGUAGE LambdaCase #-}-{-# LANGUAGE TupleSections #-}+{-# LANGUAGE CPP #-} {-| Custom database implemented with trie-maps specialized to run conjunctive queries using a (worst-case optimal) generic join algorithm.@@ -20,10 +21,15 @@ , Database(..) , Query(..) , IntTrie(..)- , Subst- , Var+ , Var(..) , Atom(..) , ClassIdOrVar(..)++ -- * Subst+ , Subst+ , lookupSubst+ , nullSubst+ , sizeSubst ) where import Data.List (sortBy)@@ -31,7 +37,10 @@ import Data.Maybe (mapMaybe) import Control.Monad -import Data.Foldable as F (toList, foldl', length)+#if !MIN_VERSION_base(4,20,0)+import Data.Foldable (foldl')+#endif+import qualified Data.Foldable as F (toList, length) import qualified Data.Map.Strict as M import qualified Data.IntMap.Strict as IM import qualified Data.IntSet as IS@@ -39,6 +48,7 @@ import Data.Equality.Graph.Classes.Id import Data.Equality.Graph.Nodes import Data.Equality.Language+import Data.Coerce -- | A variable in a query is identified by an 'Int'. -- This is much more efficient than using e.g. a 'String'.@@ -46,11 +56,8 @@ -- As a consequence, patterns also use 'Int' to represent a variable, but we -- can still have an 'Data.String.IsString' instance for variable patterns by hashing the -- string into a unique number.-type Var = Int---- | Mapping from 'Var' to 'ClassId'. In a 'Subst' there is only one--- substitution for each variable-type Subst = IM.IntMap ClassId+newtype Var = MatchVar { unsafeMatchVar :: Int }+ deriving (Show, Eq, Ord) -- | A value which is either a 'ClassId' or a 'Var' data ClassIdOrVar = CClassId {-# UNPACK #-} !ClassId@@ -66,7 +73,9 @@ -- | A conjunctive query to be run on the database data Query lang- = Query ![Var] ![Atom lang]+ = Query -- Q(x1,...,xk) <- R1(x11,...x1n), ..., Rn(xn1,...,xnn)+ ![Var] {- query head -}+ ![Atom lang] {- query body -} | SelectAllQuery {-# UNPACK #-} !Var -- | The relational representation of an e-graph, as described in section 3.1@@ -113,19 +122,18 @@ -- We want to match against ANYTHING, so we return a valid substitution for -- all existing e-class: get all relations and make a substition for each class in that relation, then join all substitutions across all classes-genericJoin (DB m) (SelectAllQuery x) = concatMap (map (IM.singleton x) . IS.toList . tkeys) (M.elems m)+genericJoin (DB m) (SelectAllQuery x) = concatMap (map (singleSubstVar x) . IS.toList . tkeys) (M.elems m) -- This is the last variable, so we return a valid substitution for every -- possible value for the variable (hence, we prepend @x@ to each and make it -- its own substitution)--- ROMES:TODO: Start here. Map vars to indexs in an array and substitute in the resulting subst genericJoin d q@(Query _ atoms) = genericJoin' atoms (orderedVarsInQuery q) where genericJoin' :: [Atom l] -> [Var] -> [Subst] genericJoin' atoms' = \case - [] -> mempty <$> atoms'+ [] -> const emptySubst <$> atoms' (!x):xs -> do @@ -134,7 +142,7 @@ -- Each valid sub-query assumes x -> x_in_D substitution y <- genericJoin' (substitute x x_in_D atoms') xs - return $! IM.insert x x_in_D y -- TODO: A bit contrieved, perhaps better to avoid map ?+ return $! insertSubstVar x x_in_D y domainX :: Var -> [Atom l] -> [Int] domainX x = IS.toList . intersectAtoms x d . filter (x `elemOfAtom`)@@ -166,11 +174,14 @@ -- described in the paper (such as batching) orderedVarsInQuery :: (Functor lang, Foldable lang) => Query lang -> [Var] orderedVarsInQuery (SelectAllQuery x) = [x]-orderedVarsInQuery (Query _ atoms) = IS.toList . IS.fromAscList $ sortBy (compare `on` varCost) $ mapMaybe toVar $ foldl' f mempty atoms+orderedVarsInQuery (Query _ atoms) = coerce . IS.toList . IS.fromAscList . coerce $+ sortBy (compare `on` varCost) $+ mapMaybe toVar $+ foldl' f mempty atoms where f :: Foldable lang => [ClassIdOrVar] -> Atom lang -> [ClassIdOrVar]- f s (Atom v (toList -> l)) = v:(l <> s)+ f s (Atom v (F.toList -> l)) = v:(l <> s) {-# INLINE f #-} -- First, prioritize variables that occur in many relations; second,@@ -210,7 +221,7 @@ -- If needed relation does exist, find intersection in it -- Add list of found intersections to existing- Just r -> case intersectInTrie var mempty r (v:toList l) of+ Just r -> case intersectInTrie var emptySubst r (v:F.toList l) of Nothing -> error "intersectInTrie should return valid substitution for variable query" Just xs -> xs @@ -235,7 +246,7 @@ -- -- TODO: Really, a valid substitution is one which isn't empty... intersectInTrie :: Var -- ^ The variable whose domain we are looking for- -> IM.IntMap ClassId -- ^ A mapping from variables that have been substituted+ -> Subst -- ^ A mapping from variables that have been substituted -> IntTrie -- ^ The trie -> [ClassIdOrVar] -- ^ The "query" -> Maybe IS.IntSet -- ^ The resulting domain for a valid substitution@@ -269,7 +280,7 @@ -- continue the recursion again to validate the assumption and -- possibly find the domain of the variable we're looking for ahead --- CVar x:xs -> case IM.lookup x substs of+ CVar x:xs -> case lookupSubst x substs of -- (2) or (4), we simply continue Just varVal -> IM.lookup varVal m >>= \next -> intersectInTrie var substs next xs -- (1) or (3)@@ -283,14 +294,14 @@ if all (isVarDifferentFrom x) xs then trieKeys else IM.foldrWithKey (\k ls (!acc) ->- case intersectInTrie var (IM.insert x k substs) ls xs of+ case intersectInTrie var (insertSubstVar x k substs) ls xs of Nothing -> acc Just _ -> k `IS.insert` acc ) mempty m -- (3) -- else IS.unions $ IM.elems $ IM.mapMaybeWithKey (\k ls -> intersectInTrie var (IM.insert x k substs) ls xs) m else IM.foldrWithKey (\k ls (!acc) ->- case intersectInTrie var (IM.insert x k substs) ls xs of+ case intersectInTrie var (insertSubstVar x k substs) ls xs of Nothing -> acc Just rs -> rs <> acc) mempty m where@@ -300,3 +311,36 @@ isVarDifferentFrom _ (CClassId _) = False isVarDifferentFrom x (CVar y) = x /= y {-# INLINE isVarDifferentFrom #-}++--------------------------------------------------------------------------------+-- * Subst+--------------------------------------------------------------------------------++-- | Mapping from 'Var' to 'ClassId'. In a 'Subst' there is only one+-- substitution for each variable+newtype Subst = Subst (IM.IntMap ClassId)++-- | Insert a 'Var' in a 'Subst'+insertSubstVar :: Var -> ClassId -> Subst -> Subst+insertSubstVar (MatchVar k) v (Subst s) = Subst (IM.insert k v s)++-- | Make a singleton 'Var' in a 'Subst'+singleSubstVar :: Var -> ClassId -> Subst+singleSubstVar (MatchVar k) v = Subst (IM.singleton k v)++-- | Lookup a 'Var' in a 'Subst'+lookupSubst :: Var -> Subst -> Maybe ClassId+lookupSubst (MatchVar k) (Subst s) = IM.lookup k s++-- | An empty substitution+emptySubst :: Subst+emptySubst = Subst IM.empty++-- | Is the 'Subst' empty?+nullSubst :: Subst -> Bool+nullSubst (Subst s) = IM.null s++sizeSubst :: Subst -> Int+sizeSubst (Subst s) = IM.size s++
src/Data/Equality/Matching/Pattern.hs view
@@ -1,15 +1,12 @@+{-# 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-import Data.Equality.Matching.Database- -- | A pattern can be either a variable or an non-variable expression of -- patterns. --@@ -64,7 +61,7 @@ -- @ data Pattern lang = NonVariablePattern (lang (Pattern lang))- | VariablePattern Var -- ^ Should be a >0 positive number+ | VariablePattern String -- | Synonym for 'NonVariablePattern'. --@@ -77,21 +74,21 @@ 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+ fromString = VariablePattern
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@@ -45,12 +46,12 @@ import qualified Data.IntMap.Strict as IM -import Data.Bifunctor import Control.Monad 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,13 +116,15 @@ 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 -- With backoff scheduler -- ROMES:TODO parMap with chunks- let (!matches, newStats) = mconcat (fmap (matchWithScheduler db i stats) (zip [1..] rewrites))+ let (!matches, newStats) = mconcat (fmap (\(rw_id,rw) ->+ let (ms, ss, vss) = matchWithScheduler db i stats rw_id rw+ in (map (\m -> (rw,m,vss)) ms, ss)) (zip [1..] rewrites)) -- Write-only phase, temporarily break invariants forM_ matches applyMatchesRhs@@ -129,7 +132,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@@ -139,61 +142,64 @@ && IM.size afterClasses == IM.size beforeClasses) (runEqualitySaturation' (i+1) newStats) - matchWithScheduler :: Database l -> Int -> IM.IntMap (Stat schd) -> (Int, Rewrite a l) -> ([(Rewrite a l, Match)], IM.IntMap (Stat schd))- matchWithScheduler db i stats = \case- (rw_id, rw :| cnd) -> first (map (first (:| cnd))) $ matchWithScheduler db i stats (rw_id, rw)- (rw_id, lhs := rhs) -> do+ matchWithScheduler :: Database l -> Int -> IM.IntMap (Stat schd) -> Int -> Rewrite a l+ -> ([Match], IM.IntMap (Stat schd), VarsState {- the vars mapping resulting from compiling the query -})+ matchWithScheduler db i stats rw_id = \case+ rw :| _ -> matchWithScheduler db i stats rw_id rw+ lhs := _ -> do+ let (lhs_query, varsState) = compileToQuery lhs+ case IM.lookup rw_id stats of -- If it's banned until some iteration, don't match this rule -- against anything.- Just s | isBanned @schd i s -> ([], stats)+ Just s | isBanned @schd i s -> ([], stats, varsState) -- Otherwise, match and update stats x -> do -- Match pattern- let matches' = ematch db lhs -- Add rewrite to the e-match substitutions+ let matches' = ematch db lhs_query -- Add rewrite to the e-match substitutions -- Backoff scheduler: update stats let newStats = updateStats schd i rw_id x stats matches' - (map (lhs := rhs,) matches', newStats)+ (matches', newStats, varsState) - applyMatchesRhs :: (Rewrite a l, Match) -> EGraphM a l ()+ applyMatchesRhs :: (Rewrite a l, Match, VarsState) -> EGraphM a l () applyMatchesRhs = \case- (rw :| cond, m@(Match subst _)) -> do+ (rw :| cond, m@(Match subst _), vss) -> do -- If the rewrite condition is satisfied, applyMatchesRhs on the rewrite rule. egr <- get- when (cond subst egr) $- applyMatchesRhs (rw, m)+ when (cond vss subst egr) $+ applyMatchesRhs (rw, m, vss) - (_ := VariablePattern v, Match subst eclass) -> do+ (_ := VariablePattern v, Match subst eclass, vss) -> do -- rhs is equal to a variable, simply merge class where lhs -- pattern was found (@eclass@) and the eclass the pattern -- variable matched (@lookup v subst@)- case IM.lookup v subst of+ case lookupSubst (findVarName vss v) subst of Nothing -> error "impossible: couldn't find v in subst" Just n -> do _ <- merge n eclass return () - (_ := NonVariablePattern rhs, Match subst eclass) -> do+ (_ := NonVariablePattern rhs, Match subst eclass, vss) -> do -- rhs is (at the top level) a non-variable pattern, so substitute -- all pattern variables in the pattern and create a new e-node (and -- e-class that represents it), then merge the e-class of the -- substituted rhs with the class that matched the left hand side- eclass' <- reprPat subst rhs+ eclass' <- reprPat vss subst rhs _ <- merge eclass eclass' return () -- | Represent a pattern in the e-graph a pattern given substitions- reprPat :: Subst -> l (Pattern l) -> EGraphM a l ClassId- reprPat subst = add . Node <=< traverse \case+ reprPat :: VarsState -> Subst -> l (Pattern l) -> EGraphM a l ClassId+ reprPat vss subst = add . Node <=< traverse \case VariablePattern v ->- case IM.lookup v subst of+ case lookupSubst (findVarName vss v) subst of Nothing -> error "impossible: couldn't find v in subst?" Just i -> return i- NonVariablePattern p -> reprPat subst p+ NonVariablePattern p -> reprPat vss subst p {-# INLINEABLE runEqualitySaturation #-}
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@@ -50,9 +49,9 @@ -- Just class_id -> -- egr^._class class_id._data /= Just 0 -- @-type RewriteCondition anl lang = Subst -> EGraph anl lang -> Bool+type RewriteCondition anl lang = VarsState -> 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/Saturation/Scheduler.hs view
@@ -16,6 +16,7 @@ import qualified Data.IntMap.Strict as IM import Data.Equality.Matching+import Data.Equality.Matching.Database (sizeSubst) -- | A 'Scheduler' determines whether a certain rewrite rule is banned from -- being used based on statistics it defines and collects on applied rewrite@@ -76,7 +77,7 @@ where -- TODO: Overall difficult, and buggy at the moment.- total_len = sum (map (length . matchSubst) matches)+ total_len = sum (map (sizeSubst . matchSubst) matches) bannedN = case currentStat of Nothing -> 0;
src/Data/Equality/Utils.hs view
@@ -1,16 +1,17 @@-{-# LANGUAGE StandaloneDeriving #-}+{-# LANGUAGE UnicodeSyntax, RankNTypes, QuantifiedConstraints, UndecidableInstances, CPP #-} {-| Misc utilities used accross modules -} module Data.Equality.Utils where -- import GHC.Conc+#if MIN_VERSION_base(4,20,0)+#else import Data.Foldable-import Data.Bits+#endif -- import qualified Data.Set as S -- import qualified Data.IntSet as IS-import Data.Functor.Classes -- | Fixed point newtype. --@@ -21,27 +22,18 @@ -- 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 cata :: Functor f => (f a -> a) -> (Fix f -> a) cata f = f . fmap (cata f) . unFix {-# INLINE cata #-}---- | Get the hash of a string.------ This util is currently used to generate an 'Int' used for the internal--- pattern variable representation from the external pattern variable--- representation ('String')-hashString :: String -> Int-hashString = foldl' (\h c -> 33*h `xor` fromEnum c) 5381-{-# INLINE hashString #-} -- -- | We don't have the parallel package, so roll our own simple parMap -- parMap :: (a -> b) -> [a] -> [b]
test/Bench.hs view
@@ -1,18 +1,43 @@ {-# LANGUAGE OverloadedStrings #-}+{-# LANGUAGE StandaloneDeriving, FlexibleInstances, DeriveAnyClass, RankNTypes, QuantifiedConstraints, UndecidableInstances, DeriveGeneric #-} import Test.Tasty.Bench +import GHC.Generics+import Control.DeepSeq+ import Data.Equality.Utils import Invariants import Sym-import Lambda-import SimpleSym +-- Instances for benchmarking. It's amazing this works!+deriving instance (forall a. Generic a => Generic (f a)) => Generic (Fix f)+deriving instance NFData UOp+deriving instance NFData BOp+deriving instance NFData a => NFData (Expr a)+deriving instance (forall a. NFData a => NFData (f a), forall a. Generic a => Generic (f a)) => NFData (Fix f)+ tests :: [Benchmark] tests = [ bgroup "Tests"- [ symTests- , lambdaTests- , simpleSymTests- , invariants+ [ bgroup "Symbolic bench"+ [ bench "i1" $+ nf rewrite (Fix $ BinOp Integral 1 "x")++ , bench "i2" $+ nf rewrite (Fix $ BinOp Integral (Fix $ UnOp Cos "x") "x")++ , bench "i3" $+ nf rewrite (Fix $ BinOp Integral (Fix $ BinOp Pow "x" 1) "x")++ , bench "i4" $+ nf rewrite (_i ((*) "x" (_cos "x")) "x")++ , bench "i5" $+ nf rewrite (_i ((*) (_cos "x") "x") "x")++ , bench "i6" $+ nf rewrite (_i (_ln "x") "x")+ ]+ -- , invariants ] ] main :: IO ()
test/Invariants.hs view
@@ -14,17 +14,16 @@ 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 +35,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,12 +44,12 @@ 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 eg :: EGraph () l- eg = snd $ equalitySaturation expr [VariablePattern 1:=fromInteger i] (error "Cost function shouldn't be used" :: CostFunction l Int)+ eg = snd $ equalitySaturation expr [VariablePattern "1":=fromInteger i] (error "Cost function shouldn't be used" :: CostFunction l Int) -- | Test 'compileToQuery'. --@@ -60,15 +59,14 @@ -- The number of atoms should also match the number of non variable patterns -- since we should create an additional atom (with a new bound variable) for each. testCompileToQuery :: Traversable lang => Pattern lang -> Bool-testCompileToQuery p = case fst $ compileToQuery p of- -- Handle special case for selectAll queries...- SelectAllQuery x -> [x] == vars p && numNonVarPatterns p == 0- q@(Query _ atoms)- | [] <- queryHeadVars q -> False- | _:xs <- queryHeadVars q ->- L.sort xs == L.sort (vars p)- && length atoms == numNonVarPatterns p- _ -> error "impossible! testCompileToQuery"+testCompileToQuery p = case compileToQuery p of+ -- Handle special case for selectAll queries...+ ((SelectAllQuery x, _), vss) -> [x] == userPatVars vss && numNonVarPatterns p == 0+ ((q@(Query _ atoms), _), vss)+ | _:xs <- queryHeadVars q ->+ L.sort xs == L.sort (userPatVars vss)+ && length atoms == numNonVarPatterns p+ | otherwise -> False where numNonVarPatterns :: Foldable lang => Pattern lang -> Int numNonVarPatterns (VariablePattern _) = 0@@ -78,19 +76,15 @@ queryHeadVars (SelectAllQuery x) = [x] queryHeadVars (Query qv _) = qv - -- | Return distinct variables in a pattern- vars :: Foldable lang => Pattern lang -> [Var]- vars (VariablePattern x) = [x]- vars (NonVariablePattern p') = LU.nubInt $ join $ map vars $ F.toList p'- -- | If we match a singleton variable pattern against an e-graph, we should get -- a match on all e-classes in the e-graph-ematchSingletonVar :: Language lang => Var -> EGraph () lang -> Bool+ematchSingletonVar :: Language lang => String -> EGraph () lang -> Bool 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)+ (q, _) = compileToQuery (VariablePattern v)+ matches = IS.fromList $ map matchClassId $ ematch db q+ eclasses = IM.keysSet (classes eg) in matches == eclasses @@ -116,7 +110,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@@ -182,7 +176,7 @@ instance Arbitrary (Pattern SimpleExpr) where arbitrary = sized p' where- p' 0 = VariablePattern <$> oneof (return <$> [1..16])+ p' 0 = VariablePattern <$> arbitrary p' n = NonVariablePattern <$> resize (n `div` 2) arbitrary newtype Name = Name { un :: String }
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 @@ -17,15 +14,10 @@ import Data.Maybe -import qualified Data.IntMap as IM import qualified Data.Set as S 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 +41,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)@@ -84,14 +72,16 @@ instance Analysis (Maybe (Lambda ())) Lambda where makeA = evalL joinA = (<|>)- modifyA c = case c^._data of- Nothing -> (c, [])- Just v -> (c, [f v])- where- f = \case- Bool b -> Fix $ Bool b- Num i -> Fix $ Num i- _ -> error "impossible, lambda () can't construct this"+ modifyA c eg+ = case eg^._class c._data of+ Nothing -> eg+ Just v -> let (c', eg') = represent (f v) eg+ in snd $ merge c c' eg'+ where+ f = \case+ Bool b -> Fix $ Bool b+ Num i -> Fix $ Num i+ _ -> error "impossible, lambda () can't construct this" -- Free variable analysis for lambda@@ -112,9 +102,6 @@ joinA = (<>) --instance Language Lambda- instance Num (Fix Lambda) where fromInteger = Fix . Num . fromInteger (+) a b = Fix $ Add a b@@ -123,17 +110,17 @@ abs = error "todo..." signum = error "todo..." -unsafeGetSubst :: Pattern Lambda -> D.Subst -> ClassId-unsafeGetSubst (NonVariablePattern _) _ = error "unsafeGetSubst: NonVariablePattern; expecting VariablePattern"-unsafeGetSubst (VariablePattern v) subst = case IM.lookup v subst of+unsafeGetSubst :: Pattern Lambda -> VarsState -> D.Subst -> ClassId+unsafeGetSubst (NonVariablePattern _) _ _ = error "unsafeGetSubst: NonVariablePattern; expecting VariablePattern"+unsafeGetSubst (VariablePattern v) vss subst = case lookupSubst (findVarName vss v) subst of Nothing -> error "Searching for non existent bound var in conditional" Just class_id -> class_id isConst :: Pattern Lambda -> RewriteCondition LA Lambda-isConst v subst egr = isJust $ snd $ egr^._class (unsafeGetSubst v subst)._data+isConst v vss subst egr = isJust $ snd $ egr^._class (unsafeGetSubst v vss subst)._data isNotSameVar :: Pattern Lambda -> Pattern Lambda -> RewriteCondition LA Lambda-isNotSameVar v1 v2 subst egr = find (unsafeGetSubst v1 subst) egr /= find (unsafeGetSubst v2 subst) egr+isNotSameVar v1 v2 vss subst egr = find (unsafeGetSubst v1 vss subst) egr /= find (unsafeGetSubst v2 vss subst) egr rules :: [Rewrite LA Lambda] rules =
test/SimpleSym.hs view
@@ -3,39 +3,28 @@ {-# 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)+import Data.Equality.Graph+import Data.Equality.Graph.Lens data SymExpr a = Const Double | Symbol String | 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 Const x -> Just x@@ -49,9 +38,12 @@ joinA Nothing Nothing = Nothing joinA (Just x) (Just y) = if x == y then Just x else error "ouch, that shouldn't have happened" - modifyA c = case c^._data of- Nothing -> (c, [])- Just i -> (c, [Fix (Const i)])+ modifyA c eg+ = case eg^._class c._data of+ Nothing -> eg+ Just i ->+ let (c', eg') = represent (Fix (Const i)) eg+ in snd $ merge c c' eg' cost :: CostFunction SymExpr Int cost = \case
test/Sym.hs view
@@ -1,31 +1,32 @@ {-# LANGUAGE MultiParamTypeClasses #-} {-# LANGUAGE TypeApplications #-}-{-# LANGUAGE TemplateHaskell #-} {-# LANGUAGE BangPatterns #-} {-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE DeriveTraversable #-} {-# LANGUAGE TypeFamilies #-} {-# LANGUAGE FlexibleInstances #-} {-# LANGUAGE LambdaCase #-}+{-# LANGUAGE DeriveGeneric #-}+{-# LANGUAGE CPP #-} module Sym where +import GHC.Generics import Test.Tasty import Test.Tasty.HUnit -import qualified Data.IntMap.Strict as IM import qualified Data.Set as S 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 +#if MIN_VERSION_base(4,18,0)+#else import Control.Applicative (liftA2)+#endif import Control.Monad (unless) +import Data.Function ((&)) import Data.Equality.Graph.Lens import Data.Equality.Graph import Data.Equality.Extraction@@ -38,8 +39,9 @@ | Const !Double | UnOp !UOp !a | BinOp !BOp !a !a- deriving ( Eq, Ord, Functor- , Foldable, Traversable+ deriving ( Eq, Ord, Show+ , Functor, Foldable, Traversable+ , Generic ) data BOp = Add | Sub@@ -48,19 +50,13 @@ | Pow | Diff | Integral- deriving (Eq, Ord, Show)+ deriving (Eq, Ord, Show, Generic) data UOp = Sin | Cos | Sqrt | Ln- deriving (Eq, Ord, Show)--deriveEq1 ''Expr-deriveOrd1 ''Expr-deriveShow1 ''Expr--instance Language Expr+ deriving (Eq, Ord, Show, Generic) instance IsString (Fix Expr) where fromString = Fix . Sym@@ -123,17 +119,15 @@ !_ <- unless (a == b || (a == 0 && b == (-0)) || (a == (-0) && b == 0)) (error "Merged non-equal constants!") return a - modifyA cl = case cl^._data of- Nothing -> (cl, [])- Just d -> ((_nodes %~ S.filter (F.null .unNode)) cl, [Fix (Const d)])-- -- -- Add constant as e-node- -- new_c <- represent (Fix $ Const d)- -- _ <- GM.merge i new_c-- -- -- Prune all except leaf e-nodes- -- modify (_class i._nodes %~ S.filter (F.null . unNode))-+ modifyA cl eg0+ = case eg0^._class cl._data of+ Nothing -> eg0+ Just d ->+ -- Add constant as e-node+ let (new_c,eg1) = represent (Fix (Const d)) eg0+ (rep, eg2) = merge cl new_c eg1+ -- Prune all except leaf e-nodes+ in eg2 & _class rep._nodes %~ S.filter (F.null .unNode) evalConstant :: Expr (Maybe Double) -> Maybe Double@@ -153,28 +147,29 @@ Sym _ -> Nothing Const x -> Just x -unsafeGetSubst :: Pattern Expr -> Subst -> ClassId-unsafeGetSubst (NonVariablePattern _) _ = error "unsafeGetSubst: NonVariablePattern; expecting VariablePattern"-unsafeGetSubst (VariablePattern v) subst = case IM.lookup v subst of+unsafeGetSubst :: Pattern Expr -> VarsState -> Subst -> ClassId+unsafeGetSubst (NonVariablePattern _) _ _ = error "unsafeGetSubst: NonVariablePattern; expecting VariablePattern"+unsafeGetSubst (VariablePattern v) vss subst = case lookupSubst (findVarName vss v) subst of Nothing -> error "Searching for non existent bound var in conditional" Just class_id -> class_id +-- | TODO: This condition is incorrect. See #35 is_not_zero :: Pattern Expr -> RewriteCondition (Maybe Double) Expr-is_not_zero v subst egr =- egr^._class (unsafeGetSubst v subst)._data /= Just 0+is_not_zero v vss subst egr =+ egr^._class (unsafeGetSubst v vss subst)._data /= Just 0 is_sym :: Pattern Expr -> RewriteCondition (Maybe Double) Expr-is_sym v subst egr =- any ((\case (Sym _) -> True; _ -> False) . unNode) (egr^._class (unsafeGetSubst v subst)._nodes)+is_sym v vss subst egr =+ any ((\case (Sym _) -> True; _ -> False) . unNode) (egr^._class (unsafeGetSubst v vss subst)._nodes) is_const :: Pattern Expr -> RewriteCondition (Maybe Double) Expr-is_const v subst egr =- isJust (egr^._class (unsafeGetSubst v subst)._data)+is_const v vss subst egr =+ isJust (egr^._class (unsafeGetSubst v vss subst)._data) is_const_or_distinct_var :: Pattern Expr -> Pattern Expr -> RewriteCondition (Maybe Double) Expr-is_const_or_distinct_var v w subst egr =- let v' = unsafeGetSubst v subst- w' = unsafeGetSubst w subst+is_const_or_distinct_var v w vss subst egr =+ let v' = unsafeGetSubst v vss subst+ w' = unsafeGetSubst w vss subst in (eClassId (egr^._class v') /= eClassId (egr^._class w')) && (isJust (egr^._class v'._data) || any ((\case (Sym _) -> True; _ -> False) . unNode) (egr^._class v'._nodes))
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"
+ test/T3.hs view
@@ -0,0 +1,47 @@+{-# LANGUAGE DeriveTraversable #-}+{-# LANGUAGE TypeFamilies #-}+{-# LANGUAGE OverloadedStrings #-}+{-# LANGUAGE TypeApplications #-}+module T3 where++-- Some e-graph unit tests++import Prelude hiding (not)++-- import Test.Tasty.HUnit+import Data.Equality.Graph+import Data.Equality.Utils+import qualified Data.Equality.Graph.Monad as EGM++data Lang a = And a a+ | Or a a+ | Not a+ | ToElim a+ | Sym Int+ deriving (Functor, Foldable, Traversable, Eq, Ord, Show)++main :: IO ()+main = do+ let _ = EGM.egraph @Lang @() $ do+ id1 <- EGM.represent (Fix (Sym 1))+ id2 <- EGM.represent (Fix (Not (Fix (Not (Fix (Sym 1))))))+ id3 <- EGM.represent (Fix (Sym 2))+ a1 <- EGM.add (Node (And id3 id1))+ a2 <- EGM.add (Node (And id3 id2))+ _ <- EGM.merge a1 a2+ EGM.rebuild -- even rebuilding this fails...+ return (id1,id2)++ -- The children should now be in the same e-class?+ --+ -- Turns out, they don't. So this test should actually be the other way+ -- around (we do not learn s1 ~= s2 from merging a1 and a2)+ --+ -- A counter example:+ -- Consider `f` that returns its second argument (f _ x = x):+ -- so f(a,c) = f(b,c), but a != b+ --+ -- find s1 eg @=? find s2 eg++ return ()+
+ test/T32.hs view
@@ -0,0 +1,57 @@+{-# LANGUAGE LambdaCase #-}+{-# LANGUAGE FlexibleInstances #-}+{-# LANGUAGE TypeFamilies #-}+{-# LANGUAGE MultiParamTypeClasses #-}+{-# LANGUAGE OverloadedStrings #-}+{-# LANGUAGE DeriveTraversable #-}+{-# LANGUAGE TypeApplications #-}+module T32 where++import Test.Tasty+import Test.Tasty.HUnit++import Data.Equality.Utils+import Data.Equality.Matching+import Data.Equality.Saturation+import Data.Equality.Graph+import qualified Data.Equality.Graph as EG+import Data.Equality.Saturation.Scheduler (defaultBackoffScheduler)+import Data.Equality.Graph.Monad++data SymExpr a = Symbol String+ | a :+: a+ deriving (Functor, Foldable, Traversable, Eq, Ord, Show)+infix 6 :+:++-- This test tests that using "VariablePattern 1, VariablePattern 2,+-- VariablePattern 3" in a rewrite rule succeeds, as opposed to using the+-- IsString instance for Pattern.+-- (Well, it tests that we can no longer screw up by writing the numbers+-- directly as well. Now the implementation transforms the strings into Ids+-- when compiling the pattern)+rewrites :: [Rewrite () SymExpr]+rewrites =+ [ pat (VariablePattern "1" :+: pat (VariablePattern "2" :+: VariablePattern "3")) := pat (pat (VariablePattern "1" :+: VariablePattern "2") :+: VariablePattern "3")+ ]++e1, e1' :: Fix SymExpr+e1 = Fix (Fix (Fix (Symbol "a") :+: Fix (Symbol "b")) :+: Fix (Symbol "c"))+e1' = Fix (Fix (Symbol "a") :+: Fix (Fix (Symbol "b") :+: Fix (Symbol "c")))++somePattern :: Pattern []+somePattern = NonVariablePattern [VariablePattern "0",VariablePattern "1"] ++-- Test basic associativity using pattern vars+testT32 :: TestTree+testT32 = testGroup "T32"+ [ testCase "basic associativity with VarPattern 1,2,3" $+ let+ (e1_id, eg0) = EG.represent @() e1 emptyEGraph+ (e1'_id, eg1) = EG.represent @() e1' eg0+ ((), eg2) = runEGraphM eg1 (runEqualitySaturation defaultBackoffScheduler rewrites)+ e1_canon = EG.find e1_id eg2+ e1'_canon = EG.find e1'_id eg2+ in e1_canon @?= e1'_canon+ -- , testCase "compiling pattern" $+ -- compileToQuery somePattern @?= Query ...+ ]
test/Test.hs view
@@ -10,9 +10,11 @@ import Sym import Lambda import SimpleSym+import T32 import qualified T1 import qualified T2+import qualified T3 tests :: TestTree tests = testGroup "Tests"@@ -22,6 +24,8 @@ , invariants , testCase "T1" (T1.main `catch` (\(e :: SomeException) -> assertFailure (show e))) , testCase "T2" (T2.main `catch` (\(e :: SomeException) -> assertFailure (show e)))+ , testCase "T3" (T3.main `catch` (\(e :: SomeException) -> assertFailure (show e)))+ , testT32 ] main :: IO ()