summaryrefslogtreecommitdiffstats
path: root/cps.sml
diff options
context:
space:
mode:
authorRose Hogenson <rosehogenson@posteo.net>2024-04-27 10:19:54 -0700
committerRose Hogenson <rosehogenson@posteo.net>2024-05-04 15:23:59 -0700
commit8ea9fc2536aa0eb45092ed3654dba19b0c8e7c20 (patch)
tree4c11670239a6722836891f5750098c7582a44011 /cps.sml
parent5737e8430d43b3b5bd448f761cd2f3c35707a174 (diff)
downloadsml-8ea9fc2536aa0eb45092ed3654dba19b0c8e7c20.tar.zst
Implement recursive functions.
Diffstat (limited to 'cps.sml')
-rw-r--r--cps.sml63
1 files changed, 38 insertions, 25 deletions
diff --git a/cps.sml b/cps.sml
index ad3e0e8..cd61a7a 100644
--- a/cps.sml
+++ b/cps.sml
@@ -51,7 +51,7 @@ struct
let
fun go [] vars =
let val temp = Gensym.new ()
- in Syntax.CRecord (map (fn v => (v, [])) (rev vars), temp, cont (Syntax.VVar temp))
+ 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))
@@ -60,7 +60,8 @@ struct
| Syntax.LSwitch (expr, arms, otherwise) =>
let
val sortedArms = Sort.sort (fn ((x, _), (y, _)) => Int.compare (x, y)) arms
- val addr = Gensym.new ()
+ val contAddr = Gensym.new ()
+ val otherwiseAddr = Gensym.new ()
val arg = Gensym.new ()
fun go _ [] cont = toCPS otherwise cont
| go v [(x, arm)] cont =
@@ -74,7 +75,7 @@ struct
( Syntax.PIf
, [Syntax.VVar b]
, []
- , [toCPS arm cont, toCPS otherwise cont]
+ , [toCPS arm cont, Syntax.CApp (Syntax.VVar otherwiseAddr, [])]
)
]
)
@@ -100,27 +101,26 @@ struct
]
)
end
+ fun contFunc x = Syntax.CApp (Syntax.VVar contAddr, [x])
in
Syntax.CFix
- ( [(addr, [arg], cont (Syntax.VVar arg))]
- , toCPS expr
- (fn v =>
- go v sortedArms
- (fn x =>
- Syntax.CApp (Syntax.VVar addr, [x])))
+ ( [ (contAddr, [arg], cont (Syntax.VVar arg))
+ , (otherwiseAddr, [], toCPS otherwise contFunc)
+ ]
+ , toCPS expr (fn v => go v sortedArms contFunc)
)
end
| _ => raise Fail ("malformed expression " ^ Syntax.lexpToString e)
fun hoist (expr : Syntax.cexp) : Syntax.cexp =
let
- fun exprs (Syntax.CRecord (args, res, k)) = Syntax.CRecord (args, res, exprs k)
+ fun exprs (Syntax.CRecord (records, k)) = Syntax.CRecord (records, exprs k)
| exprs (Syntax.CSelect (i, arg, res, k)) = Syntax.CSelect (i, arg, res, exprs k)
| exprs (expr as Syntax.CApp _) = expr
| exprs (Syntax.CFix (_, body)) = exprs body
| exprs (Syntax.CPrimop (p, args, res, ks)) = Syntax.CPrimop (p, args, res, (map exprs ks))
- fun funs (Syntax.CRecord (_, _, k)) acc = funs k acc
+ 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 =
@@ -138,12 +138,23 @@ struct
fun freeVars (expr : Syntax.cexp) : unit VarMap.map =
case expr of
- Syntax.CRecord (args, res, k) =>
+ Syntax.CRecord (records, k) =>
let
- val argFreeVars = varSet (List.mapPartial (fn (Syntax.VVar v, _) => SOME v | _ => NONE) args)
- val kFreeVars = VarMap.delete res (freeVars k)
+ val kFreeVars = freeVars k
+ val argFreeVars =
+ foldl
+ (fn ((args, _), acc) =>
+ VarMap.union acc (varSet (List.mapPartial (fn (Syntax.VVar v, _) => SOME v | _ => NONE) args)))
+ VarMap.empty
+ records
+ val results =
+ foldl
+ (fn ((_, res), acc) =>
+ VarMap.insert res () acc)
+ VarMap.empty
+ records
in
- VarMap.union argFreeVars kFreeVars
+ VarMap.difference (VarMap.union kFreeVars argFreeVars) results
end
| Syntax.CSelect (_, arg, res, k) =>
let
@@ -182,7 +193,7 @@ struct
VarMap.union argFreeVars kFreeVars
end
- fun freeVarsClosure (name, args, body) = map (fn (x, _) => x) (VarMap.toList (VarMap.difference (freeVars body) (varSet (name :: args))))
+ fun freeVarsClosure (name, args, body) = map (fn (x, _) => x) (VarMap.toList (VarMap.difference (freeVars body) (varSet args)))
fun enumerate l = ListPair.zip (List.tabulate (length l, (fn x => x)), l)
@@ -193,8 +204,8 @@ struct
| 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.CRecord (records, k) =>
+ Syntax.CRecord (map (fn (args, res) => (map (fn (v, p) => (translateValue v, p)) args, res)) records, convertExpr varMap k)
| Syntax.CSelect (i, arg, res, k) =>
Syntax.CSelect (i, translateValue arg, res, convertExpr varMap k)
| Syntax.CApp (func, args) =>
@@ -223,13 +234,15 @@ struct
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)
- (convertExpr varMap body)
- (ListPair.zip (funcs, convertedFuncs))
+ Syntax.CRecord
+ ( map
+ (fn (old as (oldName, args, body), (newName, _, _)) =>
+ let val funcFreeVars = freeVarsClosure old
+ in ((Syntax.VLabel newName, []) :: map (fn v => (Syntax.VVar (translate v), [])) funcFreeVars, oldName)
+ end)
+ (ListPair.zip (funcs, convertedFuncs))
+ , convertExpr varMap body
+ )
in
Syntax.CFix (convertedFuncs, newBody)
end