diff options
Diffstat (limited to 'elab.sml')
| -rw-r--r-- | elab.sml | 204 |
1 files changed, 182 insertions, 22 deletions
@@ -1,6 +1,7 @@ structure Elab = struct - structure StringMap = Map(type k = string val cmp = String.compare); + structure StringMap = Map(type k = string val cmp = String.compare) + structure IntMap = Map(type k = int val cmp = Int.compare) fun primop (s : string) : Syntax.primop = case s of @@ -11,7 +12,181 @@ struct | "div" => Syntax.PDiv | _ => raise Fail ("invalid op: " ^ s) - fun elab (env : int StringMap.map) (p : Syntax.expr) : Syntax.lexp = + fun enumerate l = ListPair.zip (List.tabulate (length l, (fn x => x)), l) + + fun hdstls ([] : 'a list list) : ('a list * 'a list list) option = SOME ([], []) + | hdstls ([] :: _) = NONE + | hdstls ((x :: xs) :: ls) = + case hdstls ls of + NONE => NONE + | SOME (heads, tails) => SOME (x :: heads, xs :: tails) + + fun transpose ([] : 'a list list) : 'a list list = [] + | transpose l = + case hdstls l of + NONE => [] + | SOME (heads, tails) => heads :: transpose tails + + fun patternMatrix (arms : Syntax.pat list) : Syntax.pat list list = + let val ts = + foldl + (fn (Syntax.PWild, tupleSize) => tupleSize + | (Syntax.PVar _, tupleSize) => tupleSize + | (Syntax.PInt _, _) => ~2 + | (Syntax.PTuple t, ~1) => length t + | (Syntax.PTuple t, tupleSize) => + if tupleSize = length t + then tupleSize + else raise Fail "a type checker would have caught this") + ~1 + arms + in if ts < 0 + then [arms] + else let + val expandedArms = + map + (fn Syntax.PTuple t => t + | _ => List.tabulate (ts, fn _ => Syntax.PWild)) + arms + val cols = transpose expandedArms + in List.concat (map patternMatrix cols) + end + end + + fun occurrenceVector (expr : Syntax.lexp) (arms : Syntax.pat list) : Syntax.lexp list = + let + val expandedArms = + List.mapPartial + (fn Syntax.PTuple t => SOME t + | _ => NONE) + arms + val cols = transpose expandedArms + in + case cols of + [] => [expr] + | _ => List.concat (map (fn (i, col) => occurrenceVector (Syntax.LSelect (i, expr)) col) (enumerate cols)) + end + + fun patternBindings (expr : Syntax.lexp) (Syntax.PVar v) : (string * Syntax.lexp) list = [(v, expr)] + | patternBindings expr (Syntax.PTuple t) = + List.concat (map (fn (i, p) => patternBindings (Syntax.LSelect (i, expr)) p) (enumerate t)) + | patternBindings _ _ = [] + + fun swap1 0 (l : 'a list) : 'a list = l + | swap1 n (first :: rest) = + (case swap1 (n - 1) rest of + x :: rest => x :: first :: rest + | _ => raise Fail "swap1: index out of bounds") + | swap1 _ _ = raise Fail "swap1: index out of bounds" + + fun swap (n : int) (patterns : Syntax.pat list list) (occurrences : Syntax.lexp list) : Syntax.pat list list * Syntax.lexp list = + (map (swap1 n) patterns, swap1 n occurrences) + + fun specialize (n : int) (patterns : Syntax.pat list list, occurrences : Syntax.lexp list, actions : Syntax.lexp list) : Syntax.pat list list * Syntax.lexp list * Syntax.lexp list = + let + val (patterns, actions) = + ListPair.unzip + (List.mapPartial + (fn (Syntax.PInt i :: rest, a) => if i = n then SOME (rest, a) else NONE + | (Syntax.PWild :: rest, a) => SOME (rest, a) + | (Syntax.PVar _ :: rest, a) => SOME (rest, a) + | _ => raise Fail "unexpected pattern in the matrix") + (ListPair.zip (patterns, actions))) + in (patterns, tl occurrences, actions) + end + + fun default (patterns : Syntax.pat list list, occurrences : Syntax.lexp list, actions : Syntax.lexp list) : Syntax.pat list list * Syntax.lexp list * Syntax.lexp list = + let + val (patterns, actions) = + ListPair.unzip + (List.mapPartial + (fn (Syntax.PWild :: rest, a) => SOME (rest, a) + | (Syntax.PVar _ :: rest, a) => SOME (rest, a) + | _ => NONE) + (ListPair.zip (patterns, actions))) + in (patterns, tl occurrences, actions) + end + + fun compilePatternMatching ([] : Syntax.pat list list, _ : Syntax.lexp list, _ : Syntax.lexp list) : Syntax.lexp = + raise Fail "nonexhaustive match" + | compilePatternMatching (patterns as firstRow :: rows, occurrences, actions) = + let val refutablePattern = + List.find + (fn (_, Syntax.PInt _) => true + | _ => false) + (enumerate firstRow) + in + case refutablePattern of + NONE => + if null rows + then hd actions + else raise Fail "redundant match" + | SOME (i, _) => + let + val (patterns, occurrences) = + if i = 0 + then (patterns, occurrences) + else swap i patterns occurrences + val signatures = + map (fn (x, _) => x) + (IntMap.toList + (foldl + (fn (Syntax.PInt i, acc) => IntMap.insert i true acc + | (_, acc) => acc) + IntMap.empty + (map hd patterns))) + in + Syntax.LSwitch + ( hd occurrences + , map + (fn i => (i, compilePatternMatching (specialize i (patterns, occurrences, actions)))) + signatures + , compilePatternMatching (default (patterns, occurrences, actions)) + ) + end + end + + fun actionVector (env : int StringMap.map) (expr : Syntax.lexp) (arms : (Syntax.pat * Syntax.expr) list) : Syntax.lexp list = + map + (fn (p, body) => + let + val bindings = patternBindings expr p + val env = + foldl + (fn ((name, _), env) => + StringMap.insert name (Gensym.new ()) env) + env + bindings + in + foldl + (fn ((name, binding), acc) => + Syntax.LApp (Syntax.LFn (valOf (StringMap.lookup name env), acc), binding)) + (elab env body) + bindings + end) + arms + + and elabCase (env : int StringMap.map) (expr : Syntax.lexp) (arms : (Syntax.pat * Syntax.expr) list) = + let + val patterns = transpose (patternMatrix (map (fn (x, _) => x) arms)) + val occurrences = occurrenceVector expr (map (fn (x, _) => x) arms) + val actions = actionVector env expr arms + val actionFns = + map + (fn a => (Gensym.new (), Gensym.new (), a)) + actions + val smallActions = + map + (fn (f, _, _) => Syntax.LApp (Syntax.LVar f, Syntax.LInt 0)) + actionFns + in + Syntax.LFix + ( actionFns + , compilePatternMatching (patterns, occurrences, smallActions) + ) + end + + and elab (env : int StringMap.map) (p : Syntax.expr) : Syntax.lexp = case p of Syntax.EIdent [i] => (case StringMap.lookup i env of NONE => raise Fail ("unbound identifier " ^ i) @@ -46,28 +221,13 @@ struct 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 () - val env' = - case pat of - Syntax.PVar name => StringMap.insert name v env - | _ => env - in Syntax.LFn (v, elab env' body) + | Syntax.ELambda body => + let val v = Gensym.new () + in Syntax.LFn (v, elabCase env (Syntax.LVar v) [body]) end - | Syntax.ECase (expr, []) => raise Fail "nonexhaustive match" | Syntax.ECase (expr, arms) => - let - fun go [] _ = raise Fail "nonexhaustive match" - | go ((Syntax.PWild, body) :: []) acc = Syntax.LSwitch (elab env expr, rev acc, elab env body) - | go ((Syntax.PVar name, body) :: []) acc = - let val v = Gensym.new () in - Syntax.LApp (Syntax.LFn (v, Syntax.LSwitch (Syntax.LVar v, rev acc, elab (StringMap.insert name v env) body)), elab env expr) - end - | go ((Syntax.PInt i, body) :: rest) acc = - go rest ((i, elab env body) :: acc) - | go _ _ = raise Fail "redundant match" - in go arms [] + let val v = Gensym.new () + in Syntax.LApp (Syntax.LFn (v, elabCase env (Syntax.LVar v) arms), elab env expr) end fun elaborate (p : Syntax.expr) : Syntax.lexp = elab StringMap.empty p |
