From 49b3e26f69d8911ea3ef04ff89cb198f3846aaac Mon Sep 17 00:00:00 2001 From: Tom Ellis Date: Fri, 19 Mar 2021 12:37:45 +0000 Subject: [PATCH] Checkpointing works --- src/ksc/Ksc/AD.hs | 1 + src/ksc/Ksc/ANF.hs | 2 + src/ksc/Ksc/Annotate.hs | 4 ++ src/ksc/Ksc/CSE.hs | 3 ++ src/ksc/Ksc/CatLang.hs | 1 + src/ksc/Ksc/Cgen.hs | 2 + src/ksc/Ksc/Lang.hs | 14 +++++++ src/ksc/Ksc/LangUtils.hs | 4 ++ src/ksc/Ksc/Opt.hs | 1 + src/ksc/Ksc/Opt/Shape.hs | 1 + src/ksc/Ksc/OptLet.hs | 4 ++ src/ksc/Ksc/Parse.hs | 9 ++++- src/ksc/Ksc/SUF.hs | 2 + src/ksc/Ksc/SUF/AD.hs | 13 ++++++ test/ksc/checkpoint.ks | 85 ++++++++++++++++++++++++++++++++++++++++ 15 files changed, 145 insertions(+), 1 deletion(-) create mode 100644 test/ksc/checkpoint.ks diff --git a/src/ksc/Ksc/AD.hs b/src/ksc/Ksc/AD.hs index 270839e1c..33ef17460 100644 --- a/src/ksc/Ksc/AD.hs +++ b/src/ksc/Ksc/AD.hs @@ -103,6 +103,7 @@ gradE _ _ e@(Let (TupPat _) _ _) = -- Let] pprPanic "gradE: TupPat encountered. This should not occur." (ppr e) gradE _ _ (App{}) = error "gradE of App not yet supported" +gradE _ _ Checkpoint{} = error "gradE of checkpoint not supported" -- Currently ignoring $inline when gradding. Perhaps we should -- perform the inlining before gradding. diff --git a/src/ksc/Ksc/ANF.hs b/src/ksc/Ksc/ANF.hs index 51281b52d..6da649c89 100644 --- a/src/ksc/Ksc/ANF.hs +++ b/src/ksc/Ksc/ANF.hs @@ -61,6 +61,8 @@ anfE subst (Lam v e) = do { e' <- anfExpr subst e anfE subst (Assert e1 e2) = do { e1' <- anfE subst e1 ; e2' <- anfExpr subst e2 ; return (Assert e1' e2') } +anfE subst (Checkpoint e) = do { e' <- anfExpr subst e + ; return (Checkpoint e') } -- anfE1 :: GenBndr p => ExprX p -> AnfM p (ExprX p) anfE1 :: Monad m => Subst -> TExpr -> AnfMT Typed m TExpr diff --git a/src/ksc/Ksc/Annotate.hs b/src/ksc/Ksc/Annotate.hs index 96b80cb01..e882aacf9 100644 --- a/src/ksc/Ksc/Annotate.hs +++ b/src/ksc/Ksc/Annotate.hs @@ -269,6 +269,10 @@ tcExpr (Assert cond body) text "Predicate of 'assert' has non-boolean type" ; return (TE (Assert acond abody) tybody) } +tcExpr (Checkpoint e) + = do { TE ae tye <- tcExpr e + ; return (TE (Checkpoint ae) tye) } + tcExpr (App fun arg) = do { TE afun fun_ty <- tcExpr fun ; TE aarg arg_ty <- tcExpr arg diff --git a/src/ksc/Ksc/CSE.hs b/src/ksc/Ksc/CSE.hs index 2c1c3f368..7f9a36b8d 100644 --- a/src/ksc/Ksc/CSE.hs +++ b/src/ksc/Ksc/CSE.hs @@ -160,6 +160,9 @@ cseE cse_env@(CS { cs_map = rev_map }) (Assert cond body) where cond' = cseE cse_env cond +cseE cse_env (Checkpoint e) + = Checkpoint (cseE cse_env e) + cseE cse_env (If e1 e2 e3) = If (cseE_check cse_env e1) (cseE_check cse_env e2) diff --git a/src/ksc/Ksc/CatLang.hs b/src/ksc/Ksc/CatLang.hs index a086d078d..efc5faeed 100644 --- a/src/ksc/Ksc/CatLang.hs +++ b/src/ksc/Ksc/CatLang.hs @@ -150,6 +150,7 @@ to_cl_expr Pruned _ e@(Let (TupPat _) _ _) = pprPanic "toCLExpr Let TupPat" (ppr to_cl_expr _ _ e@(Lam {}) = pprPanic "toCLExpr Lam" (ppr e) to_cl_expr _ _ e@(App {}) = pprPanic "toCLExpr App" (ppr e) to_cl_expr _ _ e@(Dummy {}) = pprPanic "toCLExpr Dummy" (ppr e) +to_cl_expr _ _ e@(Checkpoint {}) = pprPanic "toCLExpr Checkpoint" (ppr e) to_cl_expr NotPruned env e = prune env e diff --git a/src/ksc/Ksc/Cgen.hs b/src/ksc/Ksc/Cgen.hs index d39e0aa2c..eb78fc513 100644 --- a/src/ksc/Ksc/Cgen.hs +++ b/src/ksc/Ksc/Cgen.hs @@ -492,6 +492,8 @@ cgenExprWithoutResettingAlloc env = \case tybody (allocusagee1 <> allocusagebody) + Checkpoint e -> cgenExprR env e + Tuple vs -> do cgvs <- mapM (cgenExprR env) vs let cdecls = map getDecl cgvs diff --git a/src/ksc/Ksc/Lang.hs b/src/ksc/Ksc/Lang.hs index 13cfd8c69..74fae40cd 100644 --- a/src/ksc/Ksc/Lang.hs +++ b/src/ksc/Ksc/Lang.hs @@ -190,6 +190,7 @@ data ExprX p | If (ExprX p) (ExprX p) (ExprX p) | Assert (ExprX p) (ExprX p) | Dummy Type + | Checkpoint (ExprX p) deriving instance Eq (ExprX Parsed) deriving instance Eq (ExprX OccAnald) @@ -680,6 +681,7 @@ instance HasType TExpr where typeof (Let _ _ e2) = typeof e2 typeof (Assert _ e) = typeof e typeof (If _ t f) = makeIfType (typeof t) (typeof f) + typeof (Checkpoint e) = typeof e instance HasType Type where typeof t = t @@ -1138,6 +1140,7 @@ pprExpr p (Assert e1 e2) = pprExpr _ (App e1 e2) = parens (text "App" <+> sep [pprParendExpr e1, pprParendExpr e2]) -- We aren't expecting Apps, so I'm making them very visible +pprExpr _ (Checkpoint e) = parens (text "checkpoint" <+> ppr e) pprCall :: forall p. InPhase p => Prec -> FunX p -> ExprX p -> SDoc pprCall prec f e = mode @@ -1306,19 +1309,26 @@ cmpExpr e1 = go e1 M.empty where go :: TExpr -> M.Map Var TVar -> TExpr -> Ordering + go (Checkpoint e1) _ e2 + = case e2 of + Checkpoint e2' -> e1 `compare` e2' + _ -> LT go (Dummy t1) _ e2 = case e2 of + Checkpoint{} -> GT Dummy t2 -> t1 `compare` t2 _ -> LT go (Konst k1) _ e2 = case e2 of + Checkpoint{} -> GT Dummy {} -> GT Konst k2 -> k1 `compare` k2 _ -> LT go (Var v1) subst e2 = case e2 of + Checkpoint{} -> GT Dummy {} -> GT Konst {} -> GT Var v2 -> v1 `compare` M.findWithDefault v2 (tVarVar v2) subst @@ -1326,6 +1336,7 @@ cmpExpr e1 go (Call f1 e1) subst e2 = case e2 of + Checkpoint{} -> GT Dummy {} -> GT Konst {} -> GT Var {} -> GT @@ -1334,6 +1345,7 @@ cmpExpr e1 go (Tuple es1) subst e2 = case e2 of + Checkpoint{} -> GT Dummy {} -> GT Konst {} -> GT Var {} -> GT @@ -1343,6 +1355,7 @@ cmpExpr e1 go (Lam b1 e1) subst e2 = case e2 of + Checkpoint{} -> GT Dummy {} -> GT Konst {} -> GT Var {} -> GT @@ -1354,6 +1367,7 @@ cmpExpr e1 go (App e1a e1b) subst e2 = case e2 of + Checkpoint{} -> GT Dummy {} -> GT Konst {} -> GT Var {} -> GT diff --git a/src/ksc/Ksc/LangUtils.hs b/src/ksc/Ksc/LangUtils.hs index 2817e0b8f..7040156cc 100644 --- a/src/ksc/Ksc/LangUtils.hs +++ b/src/ksc/Ksc/LangUtils.hs @@ -99,6 +99,7 @@ substEMayCapture subst (Let v r b) = Let v (substEMayCapture subst r) $ substEMayCapture (subst M.\\ bindersAsMap v) b where bindersAsMap :: PatG TVar -> M.Map TVar () bindersAsMap = M.fromList . map (\x -> (x, ())) . patVars +substEMayCapture subst (Checkpoint e) = Checkpoint (substEMayCapture subst e) ----------------------------------------------- -- Free variables @@ -118,6 +119,7 @@ freeVarsOf = go go (Let v r b) = go r `S.union` (go b S.\\ S.fromList (patVars v)) go (Lam v e) = S.delete v $ go e go (Assert e1 e2) = go e1 `S.union` go e2 + go (Checkpoint e) = go e notFreeIn :: TVar -> TExpr -> Bool notFreeIn = go @@ -133,6 +135,7 @@ notFreeIn = go go v (Let v2 r b) = go v r && (v `elem` patVars v2 || go v b) go v (Lam v2 e) = v == v2 || go v e go v (Assert e1 e2) = go v e1 && go v e2 + go v (Checkpoint e) = go v e ----------------- @@ -403,3 +406,4 @@ noTupPatifyExpr in_scope = \case Konst k -> Konst k Var v -> Var v Dummy d -> Dummy d + Checkpoint e -> Checkpoint (noTupPatifyExpr in_scope e) diff --git a/src/ksc/Ksc/Opt.hs b/src/ksc/Ksc/Opt.hs index e5d78f785..928e581f0 100644 --- a/src/ksc/Ksc/Opt.hs +++ b/src/ksc/Ksc/Opt.hs @@ -131,6 +131,7 @@ optE env go e@(Dummy _) = e go (App e1 e2) = optApp env (go e1) (go e2) go (Assert e1 e2) = Assert (go e1) (go e2) + go (Checkpoint e) = Checkpoint (go e) go (Lam tv e) = Lam tv' (optE env' e) where (tv', env') = optSubstBndr tv env diff --git a/src/ksc/Ksc/Opt/Shape.hs b/src/ksc/Ksc/Opt/Shape.hs index 3613aea58..f3f12f94f 100644 --- a/src/ksc/Ksc/Opt/Shape.hs +++ b/src/ksc/Ksc/Opt/Shape.hs @@ -10,6 +10,7 @@ optShape :: TExpr -> TExpr optShape (Dummy ty) | Just s_ty <- shapeType ty = Dummy s_ty +optShape (Checkpoint e) = pShape e optShape (Assert e1 e2) = Assert e1 (pShape e2) optShape (If b t e) = If b (pShape t) (pShape e) optShape (Let v e1 e2) = Let v e1 (pShape e2) diff --git a/src/ksc/Ksc/OptLet.hs b/src/ksc/Ksc/OptLet.hs index 8e7743a10..dab6bbc40 100644 --- a/src/ksc/Ksc/OptLet.hs +++ b/src/ksc/Ksc/OptLet.hs @@ -39,6 +39,8 @@ occAnalE :: TExpr -> (ExprX OccAnald, OccMap) occAnalE (Var v) = (Var v, M.singleton v 1) occAnalE (Konst k) = (Konst k, M.empty) occAnalE (Dummy ty) = (Dummy ty, M.empty) +occAnalE (Checkpoint e) = (Checkpoint e', vs) + where (e', vs) = occAnalE e occAnalE (App e1 e2) = (App e1' e2', unionOccMap vs1 vs2) @@ -206,6 +208,7 @@ substExpr subst e go (Var tv) = substVar subst tv go (Dummy ty) = Dummy ty go (Konst k) = Konst k + go (Checkpoint e) = Checkpoint (go e) go (Call f es) = Call f (go es) go (If b t e) = If (go b) (go t) (go e) go (Tuple es) = Tuple (map go es) @@ -275,6 +278,7 @@ optLetsE = go go subst (Var tv) = substVar subst tv go _ubst (Dummy ty) = Dummy ty + go subst (Checkpoint e) = Checkpoint (go subst e) go _ubst (Konst k) = Konst k go subst (Call f es) = Call (coerceTFun f) (go subst es) go subst (If b t e) = If (go subst b) (go subst t) (go subst e) diff --git a/src/ksc/Ksc/Parse.hs b/src/ksc/Ksc/Parse.hs index 3c92444b5..1d02bfa7b 100644 --- a/src/ksc/Ksc/Parse.hs +++ b/src/ksc/Ksc/Parse.hs @@ -158,7 +158,7 @@ langDef = Tok.LanguageDef , Tok.opStart = mzero , Tok.opLetter = mzero , Tok.reservedNames = [ "def", "edef", "rule" - , "let", "if", "assert", "call", "tuple", ":", "$dummy" + , "let", "checkpoint", "if", "assert", "call", "tuple", ":", "$dummy" , "Integer", "Float", "Vec", "Lam", "String", "true", "false" ] , Tok.reservedOpNames = [] @@ -240,6 +240,7 @@ pKExpr = pIfThenElse <|> pCall <|> pTuple <|> pDummy + <|> pCheckpoint pType :: Parser TypeX pType = (pReserved "Integer" >> return TypeInteger) @@ -323,6 +324,12 @@ pLet = do { pReserved "let" ; e <- pExpr ; return $ foldr (\(v,r) e -> Let v r e) e pairs } +pCheckpoint :: Parser (ExprX Parsed) +pCheckpoint = do { pReserved "checkpoint" + ; e <- pExpr + ; return (Checkpoint e) + } + pIsUserFun :: InPhase p => Fun p -> Parser (UserFun p) pIsUserFun fun = case maybeUserFun fun of Nothing -> unexpected ("Unexpected non-UserFun in Def: " ++ render (ppr fun)) diff --git a/src/ksc/Ksc/SUF.hs b/src/ksc/Ksc/SUF.hs index 75d0359b2..a93f462b6 100644 --- a/src/ksc/Ksc/SUF.hs +++ b/src/ksc/Ksc/SUF.hs @@ -139,6 +139,8 @@ sufE avoid = \case -- SUF{dummy T} -> dummy T Dummy ty -> suf_many_and_dup L0 avoid (\L0 -> Dummy ty) + Checkpoint e -> suf_many_and_dup (L1 e) avoid (\(L1 e') -> Checkpoint e') + -- SUF{assert cond e} -> SUF{e} Assert e1 e2 -> easyVersion where -- TODO: The easy version is just to ignore the assertion. diff --git a/src/ksc/Ksc/SUF/AD.hs b/src/ksc/Ksc/SUF/AD.hs index 741a5058e..b6d0135f2 100644 --- a/src/ksc/Ksc/SUF/AD.hs +++ b/src/ksc/Ksc/SUF/AD.hs @@ -108,6 +108,19 @@ sufFwdRevPass gst subst = \case dp = patToExpr (fmap deltaOfSimple p) + Checkpoint e -> + let vs = mkTuple (map Var (S.toList (freeVarsOf e))) + vsPat = mkPat @Typed (S.toList (freeVarsOf e)) + bog = vs + + sufRevPass_ avoid' dt b = + let (fwdpass, _, revpass) = sufFwdRevPass gst avoid' e + (avoid'2, revpass_lets) = revpass avoid' dt (Let vsPat b (pSnd fwdpass)) + + in (avoid'2, revpass_lets) + + in (Tuple [e, vs], typeof bog, sufRevPass_) + -- { TODO: We currently just ignore $inline and $trace. We should -- decide what we do with them. Call f e | f `isThePrimFun` P_inline -> sufFwdRevPass gst subst e diff --git a/test/ksc/checkpoint.ks b/test/ksc/checkpoint.ks new file mode 100644 index 000000000..a86030f3e --- /dev/null +++ b/test/ksc/checkpoint.ks @@ -0,0 +1,85 @@ +;; The examples in this file are taken from: +;; +;; Jeffrey Mark Siskind & Barak A. Pearlmutter (2018) +;; Divide-and-conquer checkpointing for arbitrary programs with no +;; user annotation, Optimization Methods and Software,33:4-6, +;; 1288-1330, DOI: 10.1080/10556788.2018.1459621 +;; +;; https://engineering.purdue.edu/~qobi/papers/oms2018.pdf +;; +;; The following is another introduction to checkpointing: +;; +;; Benjamin Dauvergne and Laurent Hascoet +;; The Data-Flow Equations of Checkpointing inreverse Automatic +;; Differentiation +;; +;; https://www-sop.inria.fr/tropics/papers/DauvergneHascoet06.pdf + +(def e0 Float (x : Float) (sin x)) +(def e1 Float (x : Float) (sin x)) +(def e2 Float (x : Float) (sin x)) +(def ev Float (x : Float) (sin x)) + +(def has_a_big_bog Float (x : Float) (sin x)) + +(gdef suffwdpass [e0 Float]) +(gdef sufrevpass [e0 Float]) +(gdef suffwdpass [e1 Float]) +(gdef sufrevpass [e1 Float]) +(gdef suffwdpass [e2 Float]) +(gdef sufrevpass [e2 Float]) +(gdef suffwdpass [ev Float]) +(gdef sufrevpass [ev Float]) + +(gdef suffwdpass [has_a_big_bog Float]) +(gdef sufrevpass [has_a_big_bog Float]) + +(def without_checkpointing Float (x : Float) + (let (p0 (has_a_big_bog x)) + (let (p1 (e1 p0)) + p1))) + +(def with_checkpointing Float (x : Float) + (let (p0 (checkpoint (has_a_big_bog x))) + (let (p1 (e1 p0)) + p1))) + +(def figure2b Float (u : Float) + (let (p (checkpoint (e0 u))) + (ev p))) + +(def figure2c Float (u : Float) + (let (p2 (checkpoint + (let (p1 (checkpoint + (let (p0 (checkpoint (e0 u))) + (e1 p0)))) + (e2 p1)))) + (ev p2))) + +(def figure2d Float (u : Float) + (let (p0 (checkpoint (e0 u))) + (let (p1 (checkpoint (e1 p0))) + (let (p2 (checkpoint (e2 p1))) + (ev p2))))) + +(def figure2e Float (u : Float) + (let (p1 + (checkpoint (let (p0 (checkpoint (e0 u))) + (e1 p0)))) + (let (p2 (checkpoint (e2 p1))) + (ev p2)))) + +(gdef suffwdpass [without_checkpointing Float]) +(gdef sufrevpass [without_checkpointing Float]) +(gdef suffwdpass [with_checkpointing Float]) +(gdef sufrevpass [with_checkpointing Float]) +(gdef suffwdpass [figure2b Float]) +(gdef sufrevpass [figure2b Float]) +(gdef suffwdpass [figure2c Float]) +(gdef sufrevpass [figure2c Float]) +(gdef suffwdpass [figure2d Float]) +(gdef sufrevpass [figure2d Float]) +(gdef suffwdpass [figure2e Float]) +(gdef sufrevpass [figure2e Float]) + +(def main Integer () 0)