summaryrefslogtreecommitdiffstats
path: root/cps.sml
diff options
context:
space:
mode:
Diffstat (limited to 'cps.sml')
-rw-r--r--cps.sml161
1 files changed, 161 insertions, 0 deletions
diff --git a/cps.sml b/cps.sml
new file mode 100644
index 0000000..948b8c8
--- /dev/null
+++ b/cps.sml
@@ -0,0 +1,161 @@
+structure CPS =
+struct
+ fun toCPS (e : Syntax.lexp) (cont : Syntax.value -> Syntax.cexp) : Syntax.cexp =
+ case e of
+ Syntax.LApp (Syntax.LPrim primop, Syntax.LRecord args) =>
+ let
+ val temp = Gensym.new ()
+ fun go [] acc = Syntax.CPrimop (primop, rev acc, [temp], [cont (Syntax.VVar temp)])
+ | go (arg :: args) acc = toCPS arg (fn arg' => go args (arg' :: acc))
+ in go args []
+ end
+ | Syntax.LApp (Syntax.LPrim primop, arg) => toCPS (Syntax.LApp (Syntax.LPrim primop, Syntax.LRecord [arg])) cont
+ | Syntax.LApp (f, x) =>
+ let
+ val addr = Gensym.new ()
+ val arg = Gensym.new ()
+ in Syntax.CFix
+ ([(addr, [arg], cont (Syntax.VVar arg))],
+ toCPS f (fn f' =>
+ toCPS x (fn x' =>
+ Syntax.CApp (f', [x', Syntax.VVar addr]))))
+ end
+ | Syntax.LInt i => cont (Syntax.VInt i)
+ | Syntax.LString s => cont (Syntax.VString s)
+ | Syntax.LRecord [] => cont (Syntax.VInt 0)
+ | Syntax.LRecord exprs =>
+ let
+ fun go [] vars =
+ let val temp = Gensym.new ()
+ in Syntax.CRecord (map (fn v => (v, [])) (rev vars), temp, cont (Syntax.VVar temp))
+ end
+ | go (expr :: exprs) vars =
+ toCPS expr (fn v => go exprs (v :: vars))
+ in go exprs []
+ end
+ | _ => raise Fail ("malformed expression " ^ Syntax.lexpToString e)
+
+ fun hoist (expr : Syntax.cexp) : Syntax.cexp =
+ let
+ fun funs (Syntax.CRecord (_, _, k)) acc = funs k acc
+ | funs (Syntax.CSelect (_, _, _, k)) acc = funs k acc
+ | funs (Syntax.CApp _) acc = acc
+ | funs (Syntax.CFix (fs, body)) acc = funs body (fs @ acc)
+ | funs (Syntax.CPrimop (_, _, _, ks)) acc = foldl (fn (x, acc) => funs x acc) acc ks
+
+ fun exprs (Syntax.CRecord (args, res, k)) = Syntax.CRecord (args, res, exprs k)
+ | exprs (Syntax.CSelect (i, arg, res, k)) = Syntax.CSelect (i, arg, res, exprs k)
+ | exprs (expr as Syntax.CApp (func, args)) = expr
+ | exprs (Syntax.CFix (_, body)) = body
+ | exprs (Syntax.CPrimop (p, args, res, ks)) = Syntax.CPrimop (p, args, res, (map exprs ks))
+ val entryPoint = Gensym.new ()
+ in
+ Syntax.CFix ((entryPoint, [], exprs expr) :: funs expr [], Syntax.CApp (Syntax.VLabel entryPoint, []))
+ end
+
+ structure VarMap = Map (type k = Syntax.var
+ val cmp = Int.compare)
+
+ fun varSet (l : Syntax.var list) : unit VarMap.map = VarMap.fromList (map (fn x => (x, ())) l)
+
+ fun freeVars (expr : Syntax.cexp) : unit VarMap.map =
+ case expr of
+ Syntax.CRecord (args, res, k) =>
+ let
+ val argFreeVars = varSet (List.mapPartial (fn (Syntax.VVar v, _) => SOME v | _ => NONE) args)
+ val kFreeVars = VarMap.delete res (freeVars k)
+ in
+ VarMap.union argFreeVars kFreeVars
+ end
+ | Syntax.CSelect (_, arg, res, k) =>
+ let
+ val argFreeVars =
+ case arg of
+ Syntax.VVar v => varSet [v]
+ | _ => VarMap.empty
+ val kFreeVars = VarMap.delete res (freeVars k)
+ in
+ VarMap.union argFreeVars kFreeVars
+ end
+ | Syntax.CApp (func, args) =>
+ let
+ val funcFreeVars =
+ case func of
+ Syntax.VVar v => varSet [v]
+ | _ => VarMap.empty
+ val argFreeVars = varSet (List.mapPartial (fn Syntax.VVar v => SOME v | _ => NONE) args)
+ in
+ VarMap.union funcFreeVars argFreeVars
+ end
+ | Syntax.CFix (funs, body) =>
+ let
+ val names = varSet (map (fn (name, _, _) => name) funs)
+ val funsFreeVars = map (fn (name, args, fixBody) => VarMap.difference (freeVars fixBody) (varSet args)) funs
+ val bodyFreeVars = freeVars body
+ in
+ foldl (fn (x, acc) => VarMap.union acc (VarMap.difference x names)) VarMap.empty (bodyFreeVars :: funsFreeVars)
+ end
+ | Syntax.CPrimop (_, args, res, ks) =>
+ let
+ val argFreeVars = varSet (List.mapPartial (fn Syntax.VVar v => SOME v | _ => NONE) args)
+ val boundVars = varSet res
+ val kFreeVars = foldl (fn (x, acc) => VarMap.union acc x) VarMap.empty (map (fn k => VarMap.difference (freeVars k) boundVars) ks)
+ in
+ VarMap.union argFreeVars kFreeVars
+ end
+
+ fun freeVarsClosure (name, args, body) = map (fn (x, _) => x) (VarMap.toList (VarMap.difference (freeVars body) (varSet (name :: args))))
+
+ fun enumerate l = ListPair.zip (List.tabulate (length l, (fn x => x)), l)
+
+ fun convertExpr varMap expr =
+ let
+ fun translate var = getOpt (VarMap.lookup var varMap, var)
+ fun translateValue (Syntax.VVar v) = Syntax.VVar (translate v)
+ | translateValue v = v
+ in
+ case expr of
+ Syntax.CRecord (args, res, k) =>
+ Syntax.CRecord (map (fn (v, p) => (translateValue v, p)) args, res, convertExpr varMap k)
+ | Syntax.CSelect (i, arg, res, k) =>
+ Syntax.CSelect (i, translateValue arg, res, convertExpr varMap k)
+ | Syntax.CApp (func, args) =>
+ let val temp = Gensym.new ()
+ in Syntax.CSelect (0, translateValue func, temp,
+ Syntax.CApp (Syntax.VVar temp, map translateValue args))
+ end
+ | Syntax.CFix (funcs, body) =>
+ let
+ val convertedFuncs =
+ map
+ (fn this as (name, args, body) =>
+ let
+ val funcFreeVars = freeVarsClosure this
+ val varMap' = VarMap.union varMap (VarMap.fromList (map (fn v => (v, Gensym.new ())) funcFreeVars))
+ val closure = Gensym.new ()
+ val newBody =
+ foldl
+ (fn ((i, x), acc) => Syntax.CSelect (i + 1, Syntax.VVar closure, valOf (VarMap.lookup x varMap'), acc))
+ (convertExpr varMap' body)
+ (enumerate funcFreeVars)
+ in
+ (Gensym.new (), closure :: args, newBody)
+ end)
+ funcs
+ val newBody =
+ foldl
+ (fn ((old as (oldName, args, body), (newName, _, _)), acc) =>
+ let val funcFreeVars = freeVarsClosure old
+ in Syntax.CRecord ((Syntax.VLabel newName, []) :: map (fn v => (Syntax.VVar v, [])) funcFreeVars, oldName, acc)
+ end)
+ body
+ (ListPair.zip (funcs, convertedFuncs))
+ in
+ Syntax.CFix (convertedFuncs, newBody)
+ end
+ | Syntax.CPrimop (p, args, res, ks) =>
+ Syntax.CPrimop (p, map translateValue args, res, map (convertExpr varMap) ks)
+ end
+
+ fun convertClosures (expr : Syntax.cexp) : Syntax.cexp = hoist (convertExpr VarMap.empty expr)
+end