From bc599df65fc521364ecf4bce40d3f784d9cf8027 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Madeleine=20Sydney=20=C5=9Alaga?= Date: Wed, 2 Sep 2026 16:04:45 -0600 Subject: [PATCH] shared closures maybe --- .dir-locals.el | 3 + doc/closure-conversion.org | 31 ++++ src/Gyehoek/CPS/Close.hs | 37 +++-- src/Gyehoek/CPS/Hoist.hs | 4 +- src/Gyehoek/CPS/Stackify.hs | 320 ++++++++++++------------------------ src/Gyehoek/CPS/Syntax.hs | 26 ++- src/Gyehoek/Stack/Syntax.hs | 4 +- 7 files changed, 183 insertions(+), 242 deletions(-) diff --git a/.dir-locals.el b/.dir-locals.el index fc1236c..6062139 100644 --- a/.dir-locals.el +++ b/.dir-locals.el @@ -9,6 +9,9 @@ . (progn (defun apply-cabal-fmt-h () (haskell-mode-buffer-apply-command "cabal-fmt")) (add-hook 'before-save-hook #'apply-cabal-fmt-h nil t))))) + (scheme-mode + . ((eval . (dolist (s '(kappa κ prim)) + (put s 'scheme-indent-function 1))))) (nil . ((eval . (progn (defun display-ansi () diff --git a/doc/closure-conversion.org b/doc/closure-conversion.org index 56a7f1d..993a5be 100644 --- a/doc/closure-conversion.org +++ b/doc/closure-conversion.org @@ -132,3 +132,34 @@ multiple ~env-ref~ calls could probably be replaced with a primitive that loads $code) 1)))) #+end_src + +** example + +#+begin_src scheme + (λ (n m ktail) + (letrec ((f (λ (x ktail-0) (+ x n ktail-0))) + (g (λ (y ktail-1) (+ y g ktail-1)))) + (prim (cons f g) ktail))) +#+end_src + +#+begin_src scheme + (λ (n m ktail) + (letrec ((f-code (λ (x ktail-0) + (prim (env-get 2) + (κ (n) + (+ x n ktail-0))))) + (g-code (λ (y ktail-1) + (prim (env-get 3) + (κ (m) + (+ y m ktail-1)))))) + (letrec ((with-closure-code + (κ (f g) + (prim (get-env 0) + (κ (ktail) + (prim cons f g ktail)))))) + (prim (make-shared-closure (with-closure-code) + ktail) + (κ (with-closure) + (prim (make-shared-closure (f-code g-code) n m) + with-closure)))))) +#+end_src diff --git a/src/Gyehoek/CPS/Close.hs b/src/Gyehoek/CPS/Close.hs index d715c2c..2080d05 100644 --- a/src/Gyehoek/CPS/Close.hs +++ b/src/Gyehoek/CPS/Close.hs @@ -16,38 +16,41 @@ import Data.Traversable genCodeName :: GenSym :> es => Name -> Eff es Name genCodeName f = gensym' @Name $ f ^. _Wrapped' . to (<> "-code") -bindEnv :: Name -> List Name -> Exp -> Exp -bindEnv l frees m = [cps| - (letrec ((#{l} (κ #{frees} #{m}))) - (prim (get-env) #{l})) +bindEnv :: List Name -> Exp -> Exp +bindEnv frees m = [cps| + (prim (get-env) (κ #{frees} #{m})) |] -close :: forall es. GenSym :> es => Exp -> Eff es Exp -close = transformM \case - ExpLetRec bs e -> do +close1 :: forall es. GenSym :> es => Exp -> Eff es Exp +close1 = \case + lr@(ExpLetRec bs e) -> do let boundNames = bs ^.. each . _1 + let boundNames' = setOf each boundNames let frees = bs & foldMapOf - (each . _2 . absBody) - (freeWithBound' $ setOf (each . _1) bs) + (each . _2) + (freeWithBound' boundNames') & nub + pTraceShowM frees env_cont_l <- gensym' @Name "env-cont" e_l <- gensym' @Name "letrec-body-cont" - -- let ab' = ab & absBody .~ m' bs' <- for bs \(f,ab) -> do f_code_l <- genCodeName f - pure - ( f_code_l - , ab & absBody %~ bindEnv f_code_l (boundNames ++ frees) - ) + pure ( f_code_l + , ab & absBody %~ bindEnv (boundNames ++ frees) + ) + let codes = bs' ^.. each . _1 . to MkLabel pure [cps| (letrec #{bs'} - (letrec ((#{e_l} (κ #{boundNames} #{e}))) - (prim (make-shared-closure #{boundNames} #{frees}) - #{e_l}))) + (prim (make-shared-closure #{codes} #{frees}) + (κ #{boundNames} + #{e}))) |] e -> pure e +close :: forall es. GenSym :> es => Exp -> Eff es Exp +close = transformM close1 + closeProgram :: GenSym :> es => Program -> Eff es Program closeProgram = traverseOf (#body . #body) close diff --git a/src/Gyehoek/CPS/Hoist.hs b/src/Gyehoek/CPS/Hoist.hs index 0031132..1b8b265 100644 --- a/src/Gyehoek/CPS/Hoist.hs +++ b/src/Gyehoek/CPS/Hoist.hs @@ -9,12 +9,12 @@ import Effectful.Writer.Static.Local import Data.Foldable -type Hoist = Writer (HashMap Name Abs) +type Hoist = Writer (HashMap Label Abs) hoist :: Hoist :> es => Exp -> Eff es Exp hoist = transformM \case ExpLetRec bs m -> do - traverse_ (\(k,v) -> tell $ H.singleton k v) bs + traverse_ (\(k,v) -> tell $ H.singleton (MkLabel k) v) bs pure m e -> pure e diff --git a/src/Gyehoek/CPS/Stackify.hs b/src/Gyehoek/CPS/Stackify.hs index 4896873..e21fa68 100644 --- a/src/Gyehoek/CPS/Stackify.hs +++ b/src/Gyehoek/CPS/Stackify.hs @@ -17,29 +17,9 @@ import Data.Text qualified as T import Gyehoek.Prelude import Debug.Pretty.Simple import qualified Gyehoek.Sexp as S +import Data.Monoid -type Stackify = Writer Stk.Program - -runStackify :: Eff (Stackify : es) a -> Eff es (a, Stk.Program) -runStackify = runWriter - -live :: Free a => Env -> a -> List Name --- TODO: free' should return an OSet lol -live g e = nub (free' e) & filter \x -> - x `elem` g.bound - -- && not (x `elem` g.contStack) - --- | The expression @load g e r n@ emits a 'Stk.Load' instruction if --- stack variable @n@ is live-out in expression @e@. Otherwise, a --- 'Stk.Pop' instruction is emitted. -load :: Free a => Env -> a -> Reg -> Int -> Stk.Instr -load g e r 0 - | Just x <- g ^? #bound . _head - , x `elem` free e - = Stk.Pop r -load g e r n = Stk.Load r n - data BlockBuilder = Code (List Stk.Instr) BlockBuilder | Tail Stk.Tail @@ -50,231 +30,137 @@ buildBlock = go [] where go acc (Code xs bb) = go (acc ++ xs) bb go acc (Tail t) = Stk.MkBlock acc t -emitRoutine :: Stackify :> es => Stk.Routine -> Eff es () -emitRoutine rt = tell [rt] - -stackify - :: forall es. (GenSym :> es, Stackify :> es) - => Env -> Exp -> Eff es BlockBuilder - -stackify g (ExpLetRec bs e) = do - for_ bs \(f,a) -> - emitRoutine =<< case a of - AbsKappa kap -> stackifyKappa g (MkLabel f) kap - AbsLambda lam -> stackifyLambda g (MkLabel f) lam - stackify g e - -stackify g (ExpIf c t f) = do - let c' = stackifyVal g c - let jump l = - Stk.MkBlock - [Stk.Push . Stk.ValLabel . MkLabel $ l] - (Stk.TailCall 0) - pure . Tail $ Stk.If c' (jump t) (jump f) - -stackify g (ExpApply f xs ktail) = do - pure $ - Code [ Stk.Push $ stackifyVal g (ValVar $ ktail ^?! #KexpVar) - , Stk.Push $ stackifyVal g f - ] $ - Code (pushArgs g xs) $ - Tail (Stk.Call (length xs)) - -stackify g e@(ExpContinue k xs) - | isn't (#_ValVar . only g.tail) k = pure $ - Code [ Stk.Push (stackifyVal g k) ] $ - Code (pushArgs g xs) $ - Tail $ Stk.TailCall (length xs) - | otherwise = pure $ - Code (pushArgs g xs) $ - Tail (Stk.Return (length xs)) - -stackify g (ExpPrim (PrimCallCC withcc) cc) = do - let cc' = cc ^?! #KexpVar . to MkLabel - cc_l <- gensym' @Label "cc" - pure $ - Code [ Stk.Push $ stackifyVal g withcc - , Stk.Push $ stackifyVal g (ValLabel cc') - ] $ - Tail Stk.CallCC - -stackify g (ExpPrim p cc) = pure $ - Code [ Stk.Push (Stk.ValLabel . MkLabel $ cc ^?! #KexpVar) - , Stk.Prim (stackifyVal g <$> p) - ] $ - Tail $ Stk.TailCall 1 - -stackify _ e = error [i|unimplemented exp: #{e}|] - -loadArgs :: List Name -> List Stk.Instr -loadArgs = imapOf itraversed \n x -> Stk.Load (MkReg x) n - -pushArgs :: Env -> List Val -> List Stk.Instr -pushArgs g args = [ Stk.Push $ stackifyVal g x | x <- reverse args ] - -- affine _ValName :: Traversal' Val Name _ValName = failing #_ValVar (#_ValImm . #_ImmLabel . #_MkLabel) -stackifyKappa - :: (Stackify :> es, GenSym :> es) - => Env -> Label -> Kappa - -> Eff es Stk.Routine -stackifyKappa g kname (MkKappa xs m) = do - let g' = g & #bound <>:~ xs - m' <- stackify g' m - pure $ - Stk.MkRoutine kname . buildBlock $ - Code (loadArgs xs) $ - Code [ Stk.Load (MkReg r) j - | v <- g ^.. #liveness . ix kname . each - , (j,r) <- itoListOf (#bound . itraversed) g - , r == v - ] $ - m' +stackify + :: forall es. (GenSym :> es) + => Env -> Exp -> Eff es BlockBuilder -stackifyLambda - :: (Stackify :> es, GenSym :> es) - => Env -> Label -> Lambda - -> Eff es Stk.Routine -stackifyLambda g name (MkLambda xs k m) = do - m' <- stackify (g & #bound .~ xs & #tail .~ k) m - pure $ - Stk.MkRoutine name . buildBlock $ - Code (loadArgs xs) $ - -- Code [Stk.Load (MkReg k) (length xs + 1)] $ - m' +stackify _ (ExpContinue (ValVar k) xs) = + Code [ ] _ -stackifyVal :: Env -> Val -> Stk.Val -stackifyVal g = \case - ValImm imm -> Stk.ValImm imm - ValVar v -> case regOf g v of - Just r -> Stk.ValReg r - Nothing -> Stk.ValLabel (MkLabel v) - v -> error [i|unimplemented val: #{v}|] +stackify _ (ExpPrim p k) = _ -regOf :: Env -> Name -> Maybe Reg -regOf g x - | x `elem` g.bound || x == g.tail = Just . MkReg $ x - | otherwise = Nothing +stackify _ e = error [i|unimplemented exp: #{e}|] + +stackifyAbs :: (GenSym :> es) => Env -> Label -> Abs -> Eff es Stk.Routine + +stackifyAbs g lbl (MkAbs xs mtail e) = + Stk.MkRoutine lbl . buildBlock . preamble <$> stackify g e + where + preamble = Code (popArgs $ (mtail ^.. _Just) ++ xs) + +popArgs :: List Name -> List Stk.Instr +popArgs = fmap (Stk.Pop . MkReg) . reverse + +pushArgs :: List Name -> List Stk.Instr +pushArgs = _ data Env = MkEnv - -- | `bound` tracks the stack lifetime of bound variables. - { bound :: List Name - -- | for each locally-bound continuation @k@, @liveness@ has an - -- entry @(k,ls)@ where @ls@ is the sequence of registers @k@ - -- expects to find saved on the stack. - , liveness :: HashMap Label (List Name) - , tail :: Name + { } deriving (Show, Generic) emptyEnv :: Env emptyEnv = MkEnv - { bound = mempty - , liveness = mempty - , tail = "halt" + { } -stackifyProgram :: GenSym :> es => HoistedProgram -> Eff es Stk.Program -stackifyProgram p = do - let liveness = p & foldMapOf - (#bindings . itraversed . withIndex . aside #AbsKappa) - \(kname,kap) -> H.singleton - (MkLabel kname) - (nub $ freeWithBound' (H.keysSet p.bindings) kap) - let g = MkEnv - { bound = mempty - , liveness - , tail = p.body.ktail } - let e = p.body & #body %~ ExpLetRec (H.toList p.bindings) - (_,p') <- runStackify $ emitRoutine =<< stackifyLambda g "start" e - pure p' - -blah :: HoistedProgram -blah = [cps| -(letrec ((prim-k3 (κ (r2) (if r2 truthy-cont4 falsey-cont5))) - (prim-k7 (κ (r6) (fac r6 r8))) - (make-closure-cont15 (κ (fac) (fac 20 r12))) - (falsey-cont5 (κ () (prim (- n 1) prim-k7))) - (r12 (κ (x13) (continue start-ktail0 x13))) - (truthy-cont4 (κ () (continue lambda-tail1 1))) - (prim-k11 (κ (r10) (continue lambda-tail1 r10))) - (r8 (κ (x9) (prim (* n x9) prim-k11))) - (fac-code14 (λ (n lambda-tail1) (prim (zero? n) prim-k3)))) - (λ (start-ktail0) - (prim (make-closure $fac-code14) make-closure-cont15))) -|] +stackifyProgram + :: forall es. GenSym :> es + => HoistedProgram -> Eff es Stk.Program +stackifyProgram p = p + & ifoldMapOf + ((#bindings . itraversed) + <> (#body . to (H.singleton "start" . AbsLambda) . itraversed)) + (\l -> Ap . stackifyBinding l) + & getAp + where + g = emptyEnv + stackifyBinding lbl ab = + Stk.MkProgram . H.singleton lbl <$> stackifyAbs @es g lbl ab p :: HoistedProgram p = [cps| -(letrec ((r12-code32 (κ (r12 start-ktail0 x13) (continue start-ktail0 x13))) - (prim-k7-code22 - (κ (prim-k7 fac r6 r8 n x9 prim-k11 lambda-tail1 r10) +(letrec (($r12-code32 + (κ (x13) (prim - (make-shared-closure (r8) (n x9 prim-k11 lambda-tail1 r10)) - letrec-body-cont18))) - (letrec-body-cont24 - (κ (truthy-cont4 falsey-cont5) - (if r2 - truthy-cont4 - falsey-cont5))) - (letrec-body-cont18 (κ (r8) (fac r6 r8))) - (prim-k11-code16 - (κ (prim-k11 lambda-tail1 r10) - (continue lambda-tail1 r10))) - (falsey-cont5-code26 - (κ (truthy-cont4 falsey-cont5 lambda-tail1 n prim-k7 - fac r6 r8 x9 prim-k11 r10) + (get-env) + (κ (r12 start-ktail0) + (continue start-ktail0 x13))))) + ($prim-k7-code22 + (κ (r6) (prim - (make-shared-closure - (prim-k7) - (fac r6 r8 n x9 prim-k11 lambda-tail1 r10)) - letrec-body-cont21))) - (letrec-body-cont31 (κ (r12) (fac 20 r12))) - (r8-code19 - (κ (r8 n x9 prim-k11 lambda-tail1 r10) + (get-env) + (κ (prim-k7 lambda-tail1 n fac) + (prim + (make-shared-closure ($r8-code19) (lambda-tail1 n)) + (κ (r8) + (fac r6 r8))))))) + ($prim-k11-code16 + (κ (r10) (prim - (make-shared-closure (prim-k11) (lambda-tail1 r10)) - letrec-body-cont15))) - (letrec-body-cont28 (κ (prim-k3) (prim (zero? n) prim-k3))) - (truthy-cont4-code25 - (κ (truthy-cont4 falsey-cont5 lambda-tail1 n prim-k7 fac - r6 r8 x9 prim-k11 r10) - (continue lambda-tail1 1))) - (fac-code35 - (κ (fac n prim-k3 r2 truthy-cont4 falsey-cont5 lambda-tail1 - prim-k7 r6 r8 x9 prim-k11 r10) + (get-env) + (κ (prim-k11 lambda-tail1) + (continue lambda-tail1 r10))))) + ($falsey-cont5-code26 + (κ () (prim - (make-shared-closure - (prim-k3) - (r2 truthy-cont4 falsey-cont5 lambda-tail1 n prim-k7 - fac r6 r8 x9 prim-k11 r10)) - letrec-body-cont28))) - (prim-k3-code29 - (κ (prim-k3 r2 truthy-cont4 falsey-cont5 lambda-tail1 n prim-k7 - fac r6 r8 x9 prim-k11 r10) + (get-env) + (κ (truthy-cont4 falsey-cont5 lambda-tail1 n fac) + (prim + (make-shared-closure ($prim-k7-code22) (lambda-tail1 n fac)) + (κ (prim-k7) + (prim (- n 1) prim-k7))))))) + ($r8-code19 + (κ (x9) (prim - (make-shared-closure - (truthy-cont4 falsey-cont5) - (lambda-tail1 n prim-k7 fac r6 r8 x9 prim-k11 r10)) - letrec-body-cont24))) - (letrec-body-cont21 (κ (prim-k7) (prim (- n 1) prim-k7))) - (letrec-body-cont15 (κ (prim-k11) (prim (* n x9) prim-k11))) - (letrec-body-cont34 - (κ (fac) + (get-env) + (κ (r8 lambda-tail1 n) + (prim + (make-shared-closure ($prim-k11-code16) (lambda-tail1)) + (κ (prim-k11) + (prim (* n x9) prim-k11))))))) + ($truthy-cont4-code25 + (κ () (prim - (make-shared-closure (r12) (start-ktail0 x13)) - letrec-body-cont31)))) + (get-env) + (κ (truthy-cont4 falsey-cont5 lambda-tail1 n fac) + (continue lambda-tail1 1))))) + ($fac-code35 + (λ (n lambda-tail1) + (prim + (get-env) + (κ (fac) + (prim + (make-shared-closure ($prim-k3-code29) (lambda-tail1 n fac)) + (κ (prim-k3) + (prim (zero? n) prim-k3))))))) + ($prim-k3-code29 + (κ (r2) + (prim + (get-env) + (κ (prim-k3 lambda-tail1 n fac) + (prim + (make-shared-closure + ($truthy-cont4-code25 $falsey-cont5-code26) + (lambda-tail1 n fac)) + (κ (truthy-cont4 falsey-cont5) + (if r2 + truthy-cont4 + falsey-cont5)))))))) (λ (start-ktail0) (prim - (make-shared-closure - (fac) - (n prim-k3 r2 truthy-cont4 falsey-cont5 lambda-tail1 - prim-k7 r6 r8 x9 prim-k11 r10)) - letrec-body-cont34))) + (make-shared-closure ($fac-code35) ()) + (κ (fac) + (prim + (make-shared-closure ($r12-code32) (start-ktail0)) + (κ (r12) + (fac 20 r12))))))) |] diff --git a/src/Gyehoek/CPS/Syntax.hs b/src/Gyehoek/CPS/Syntax.hs index 409853e..09e596c 100644 --- a/src/Gyehoek/CPS/Syntax.hs +++ b/src/Gyehoek/CPS/Syntax.hs @@ -42,6 +42,8 @@ module Gyehoek.CPS.Syntax , pattern ValLabel , pattern ObjLabel , absBody + , pattern MkAbs + , _MkAbs ) where @@ -125,6 +127,23 @@ pattern AbsKappa' xs e = AbsKappa (MkKappa xs e) pattern AbsLambda' :: List Name -> Name -> Exp -> Abs pattern AbsLambda' xs e ktail = AbsLambda (MkLambda xs e ktail) +{-# COMPLETE AbsKappa', AbsLambda' #-} + +_MkAbs :: Iso' Abs (List Name, Maybe Name, Exp) +_MkAbs = iso + (\case + AbsKappa' xs e -> (xs,Nothing,e) + AbsLambda' xs ktail e -> (xs,Just ktail,e)) + (\(xs,ktail,e) -> case ktail of + Just k -> AbsLambda' xs k e + Nothing -> AbsKappa' xs e) + +pattern MkAbs :: List Name -> Maybe Name -> Exp -> Abs +pattern MkAbs xs ktail body <- (view _Abs' -> (xs,ktail,body)) + where MkAbs xs ktail body = review _Abs' (xs,ktail,body) + +{-# COMPLETE MkAbs #-} + data Exp = ExpPrim (Prim Val) Kexp | ExpLetRec { binders :: List (Name, Abs), body :: Exp } @@ -158,12 +177,12 @@ data Program = MkProgram deriving (Show, Generic, Data) data HoistedProgram = MkHoistedProgram - { bindings :: HashMap Name Abs + { bindings :: HashMap Label Abs , body :: Lambda } deriving stock (Show, Generic, Data) -type instance Index HoistedProgram = Name +type instance Index HoistedProgram = Label type instance IxValue HoistedProgram = Abs instance Ixed HoistedProgram where ix j = #bindings . ix j @@ -196,7 +215,6 @@ _AbsLambda' = prism' instance Plated Exp where plate = uniplate - absBody :: Lens' Abs Exp absBody = lens (\case @@ -350,7 +368,7 @@ instance S.DatumIso Program where instance S.DatumIso HoistedProgram where datumIso = S.with \prog -> S.letLike "letrec" - (S.datumIso @Name) (S.datumIso @Abs) (S.datumIso @Lambda) + (S.datumIso @Label) (S.datumIso @Abs) (S.datumIso @Lambda) >>> S.onTail (S.iso H.fromList H.toList) >>> prog diff --git a/src/Gyehoek/Stack/Syntax.hs b/src/Gyehoek/Stack/Syntax.hs index 08199de..79ec768 100644 --- a/src/Gyehoek/Stack/Syntax.hs +++ b/src/Gyehoek/Stack/Syntax.hs @@ -76,7 +76,7 @@ data Instr = Pop Reg | Push Val | Load Reg Int - | Prim (Prim Val) + | Prim Reg (Prim Val) deriving stock (Show, Generic, Data) deriving anyclass (NFData) @@ -99,7 +99,7 @@ instance S.DatumIso Instr where $ S.With (S.headTagged1 "pop!" S.datumIso >>>) $ S.With (S.headTagged1 "push!" S.datumIso >>>) $ S.With (S.headTagged2 "load" S.datumIso S.datumIso >>>) - $ S.With (S.headTagged1 "prim" S.datumIso >>>) + $ S.With (S.headTagged2 "prim" S.datumIso S.datumIso >>>) $ S.End where