summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--codegen.sml37
-rw-r--r--cps.sml63
-rw-r--r--elab.sml12
-rw-r--r--parser.sml6
-rw-r--r--sort.sml2
-rw-r--r--syntax.sml9
-rw-r--r--tests/13-fibonacci.sml7
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]
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
diff --git a/elab.sml b/elab.sml
index f6f2e38..af6bec2 100644
--- a/elab.sml
+++ b/elab.sml
@@ -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 ()
diff --git a/parser.sml b/parser.sml
index 3c044af..e4c7359 100644
--- a/parser.sml
+++ b/parser.sml
@@ -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
diff --git a/sort.sml b/sort.sml
index 97e4a8b..407619f 100644
--- a/sort.sml
+++ b/sort.sml
@@ -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)
diff --git a/syntax.sml b/syntax.sml
index f85e6f2..8a16e91 100644
--- a/syntax.sml
+++ b/syntax.sml
@@ -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))