diff options
| author | Rose Hogenson <rosehogenson@posteo.net> | 2024-04-27 10:19:54 -0700 |
|---|---|---|
| committer | Rose Hogenson <rosehogenson@posteo.net> | 2024-05-04 15:23:59 -0700 |
| commit | 8ea9fc2536aa0eb45092ed3654dba19b0c8e7c20 (patch) | |
| tree | 4c11670239a6722836891f5750098c7582a44011 | |
| parent | 5737e8430d43b3b5bd448f761cd2f3c35707a174 (diff) | |
| download | sml-8ea9fc2536aa0eb45092ed3654dba19b0c8e7c20.tar.zst | |
Implement recursive functions.
| -rw-r--r-- | codegen.sml | 37 | ||||
| -rw-r--r-- | cps.sml | 63 | ||||
| -rw-r--r-- | elab.sml | 12 | ||||
| -rw-r--r-- | parser.sml | 6 | ||||
| -rw-r--r-- | sort.sml | 2 | ||||
| -rw-r--r-- | syntax.sml | 9 | ||||
| -rw-r--r-- | tests/13-fibonacci.sml | 7 |
7 files changed, 92 insertions, 44 deletions
diff --git a/codegen.sml b/codegen.sml index a4fef6c..9dbee51 100644 --- a/codegen.sml +++ b/codegen.sml @@ -62,7 +62,11 @@ struct end fun go expr = case expr of - Syntax.CRecord (_, res, k) => insert res (go k) + Syntax.CRecord (records, k) => + foldl + (fn ((_, res), acc) => insert res acc) + (go k) + records | Syntax.CSelect (_, _, res, k) => insert res (go k) | Syntax.CApp _ => VarMap.empty | Syntax.CFix (funcs, body) => @@ -115,20 +119,25 @@ struct | translateVal x = x fun go expr = case expr of - Syntax.CRecord (args, res, k) => - Syntax.OAlloc (translate res, Syntax.VInt (length args)) - :: List.concat + Syntax.CRecord (records, k) => + map (fn (args, res) => Syntax.OAlloc (translate res, Syntax.VInt (length args))) + records + @ List.concat (map - (fn (i, (arg, path)) => - let val (temp, ops) = - foldl - (fn (off, (arg, ops)) => - (Syntax.VVar tempReg, Syntax.OPeek (tempReg, off, translateVal arg) :: ops)) - (translateVal arg, []) - path - in rev (Syntax.OPoke (i, translate res, temp) :: ops) - end) - (enumerate args)) + (fn (args, res) => + List.concat + (map + (fn (i, (arg, path)) => + let val (temp, ops) = + foldl + (fn (off, (arg, ops)) => + (Syntax.VVar tempReg, Syntax.OPeek (tempReg, off, arg) :: ops)) + (translateVal arg, []) + path + in rev (Syntax.OPoke (i, translate res, temp) :: ops) + end) + (enumerate args))) + records) @ go k | Syntax.CSelect (i, arg, res, k) => Syntax.OPeek (translate res, i, translateVal arg) :: go k | Syntax.CApp (func, args) => shuffle (map translateVal (func :: args)) @ [Syntax.OCall] @@ -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 @@ -34,6 +34,18 @@ struct | Syntax.ELet ([], body) => elab env body | Syntax.ELet (Syntax.DVal (pat, v) :: decls, body) => elab env (Syntax.ECase (v, [(pat, Syntax.ELet (decls, body))])) + | Syntax.ELet (Syntax.DValRec (Syntax.PVar name, f as Syntax.ELambda _) :: decls, body) => + let + val n = Gensym.new () + val env' = StringMap.insert name n env + val (arg, fnBody) = + case elab env' f of + Syntax.LFn x => x + | _ => raise Fail "Syntax.ELambda should expand to Syntax.LFn" + in + Syntax.LFix ([(n, arg, fnBody)], elab env' (Syntax.ELet (decls, body))) + end + | Syntax.ELet (Syntax.DValRec _ :: _, _) => raise Fail "invalid val rec" | Syntax.ELambda (pat, body) => let val v = Gensym.new () @@ -459,10 +459,14 @@ struct <|> orelseExpr) st and dec : Syntax.dec parser = fn st => (reserved "val" >> + bind (true <$ reserved "rec" <|> const false) (fn isRec => bind pat (fn p => reserved "=" >> bind expr (fn e => - const (Syntax.DVal (p, e))))) st + const + (if isRec + then Syntax.DValRec (p, e) + else Syntax.DVal (p, e)))))) st val program : Syntax.expr parser = (fn decs => Syntax.ELet (decs, Syntax.EInt 0)) <$> many dec fun parse (f : string) : (string, Syntax.expr) Result.either = runParser program f @@ -12,7 +12,7 @@ struct fun merge (_ : 'a * 'a -> order) ([] : 'a list) (l2 : 'a list) : 'a list = l2 | merge _ l1 [] = l1 - | merge cmp (xl as (x :: xs)) (yl as (y :: ys)) = + | merge cmp (xl as x :: xs) (yl as y :: ys) = (case cmp (x, y) of GREATER => y :: merge cmp xl ys | _ => x :: merge cmp xs yl) @@ -27,7 +27,9 @@ struct | ELambda of pat * expr | ECase of expr * (pat * expr) list - and dec = DVal of pat * expr + and dec = + DVal of pat * expr + | DValRec of pat * expr (* Lambda language *) type var = int @@ -61,7 +63,7 @@ struct | VString of string datatype cexp = - CRecord of (value * int list) list * var * cexp + CRecord of ((value * int list) list * var) list * cexp | CSelect of int * value * var * cexp | CApp of value * value list | CFix of (var * var list * cexp) list * cexp @@ -131,6 +133,7 @@ struct and decToStringI (indent : string) (x : dec) : string = case x of DVal (p, e) => "DVal (" ^ patToString p ^ ", " ^ exprToStringI indent e ^ ")" + | DValRec (p, e) => "DValRec (" ^ patToString p ^ ", " ^ exprToStringI indent e ^ ")" val exprToString : expr -> string = exprToStringI "" @@ -171,7 +174,7 @@ struct val self = cexpToStringI indent val newIndent = indent ^ "\t" in case x of - CRecord (a, b, c) => "CRecord (" ^ listToString (fn (x, y) => "(" ^ valueToString x ^ ", " ^ listToString Int.toString y ^ ")") a ^ ", " ^ Int.toString b ^ ",\n" ^ indent ^ self c ^ ")" + CRecord (records, c) => "CRecord (" ^ listToString (fn (a, b) => listToString (fn (x, y) => "(" ^ valueToString x ^ ", " ^ listToString Int.toString y ^ ")") a ^ ", " ^ Int.toString b ^ ")") records ^ ",\n" ^ indent ^ self c ^ ")" | CSelect (a, b, c, d) => "CSelect (" ^ Int.toString a ^ ", " ^ valueToString b ^ ", " ^ Int.toString c ^ ",\n" ^ indent ^ self d ^ ")" | CApp (a, b) => "CApp (" ^ valueToString a ^ ", " ^ listToString valueToString b ^ ")" | CFix (a, b) => "CFix (" ^ multilineListToString (fn indent' => fn (x, y, z) => "(" ^ Int.toString x ^ ", " ^ listToString Int.toString y ^ ",\n" ^ indent' ^ "\t" ^ cexpToStringI (indent' ^ "\t") z) newIndent a ^ ",\n" ^ newIndent ^ cexpToStringI newIndent b ^ ")" diff --git a/tests/13-fibonacci.sml b/tests/13-fibonacci.sml new file mode 100644 index 0000000..ec89f70 --- /dev/null +++ b/tests/13-fibonacci.sml @@ -0,0 +1,7 @@ +val rec fib = fn n => + case n of + 0 => 0 + | 1 => 1 + | _ => __builtin "add" (fib (__builtin "sub" (n, 1)), fib (__builtin "sub" (n, 2))) + +val _ = __builtin "exit" (__builtin "add" (fib 9, 8)) |
