From 92baa4f6bf99efe229abf2ab895cec8fd7a8936d Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Fri, 17 May 2024 16:16:49 -0700 Subject: Fix case on datatype. --- elab.sml | 108 ++++++++++++++++++++++++++++++++++++++++++++++----------------- 1 file changed, 79 insertions(+), 29 deletions(-) (limited to 'elab.sml') diff --git a/elab.sml b/elab.sml index 3a361a6..eb168c5 100644 --- a/elab.sml +++ b/elab.sml @@ -12,12 +12,12 @@ struct | "div" => Syntax.PDiv | _ => raise Fail ("invalid op: " ^ s) - fun mapIdx (f : int * 'a -> 'b) (l : 'a list) : 'b list = - let - fun go _ [] = [] - | go i (x :: xs) = f (i, x) :: go (i + 1) xs - in go 0 l - end + fun enumerate (l : 'a list) : (int * 'a) list = + let + fun go _ [] = [] + | go i (x :: xs) = (i, x) :: go (i + 1) xs + in go 0 l + end fun hdstls ([] : 'a list list) : ('a list * 'a list list) option = SOME ([], []) | hdstls ([] :: _) = NONE @@ -32,7 +32,7 @@ struct NONE => [] | SOME (heads, tails) => heads :: transpose tails - type env = int StringMap.map * int StringMap.map + type env = int StringMap.map * (int * int) StringMap.map fun bindVar (name : string) (sym : int) ((vars, types) : env) : env = (StringMap.insert name sym vars, types) @@ -43,16 +43,23 @@ struct fun bindDataCons (cons : (string * Syntax.etype option) list) ((vars, types) : env) : env = let + val nCons = length cons val (_, vars, types) = foldl (fn ((name, _), (i, vars, types)) => - (i + 1, StringMap.insert name (Gensym.new ()) vars, StringMap.insert name i types)) + (i + 1, StringMap.insert name (Gensym.new ()) vars, StringMap.insert name (i, nCons) types)) (0, vars, types) cons in (vars, types) end - fun lookupCon (name : string) ((_, types) : env) : int option = StringMap.lookup name types + fun lookupCon (name : string) ((_, types) : env) : int option = + Option.map (fn (i, _) => i) (StringMap.lookup name types) + + fun nConstructors (name : string) ((_, types) : env) : int = + let val (_, n) = valOf (StringMap.lookup name types) + in n + end fun patternMatrix (arms : Syntax.pat list) : Syntax.pat list list = let val ts = @@ -92,12 +99,12 @@ struct in case cols of [] => [expr] - | _ => List.concat (mapIdx (fn (i, col) => occurrenceVector (Syntax.LSelect (i, expr)) col) cols) + | _ => 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 (mapIdx (fn (i, p) => patternBindings (Syntax.LSelect (i, expr)) p) t) + List.concat (map (fn (i, p) => patternBindings (Syntax.LSelect (i, expr)) p) (enumerate t)) | patternBindings expr (Syntax.PCon (_, arg)) = patternBindings (Syntax.LSelect (1, expr)) arg | patternBindings _ _ = [] @@ -113,6 +120,22 @@ struct fun specialize (env : env) (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 con1 = + List.find + (fn Syntax.PCon (name, _) => valOf (lookupCon name env) = n + | _ => false) + (map hd patterns) + val newTupleSize = + case con1 of + SOME (Syntax.PCon (_, Syntax.PTuple t)) => length t + | _ => 0 + val occHead = hd occurrences + val occRest = tl occurrences + val occurrences = + if newTupleSize = 0 + then Syntax.LSelect (1, occHead) :: occRest + else + List.tabulate (newTupleSize, fn i => Syntax.LSelect (i, Syntax.LSelect (1, occHead))) @ occRest fun specializeRow (Syntax.PInt i :: rest) = if i = n then SOME (Syntax.PWild :: rest) else NONE | specializeRow (Syntax.PWild :: rest) = SOME (Syntax.PWild :: rest) @@ -131,20 +154,20 @@ struct (List.mapPartial (fn (p, a) => Option.map (fn p => (p, a)) (specializeRow p)) - (ListPair.zip (patterns, actions))) - in (patterns, tl occurrences, actions) + (ListPair.zipEq (patterns, actions))) + in (patterns, 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) + (List.filter + (fn (Syntax.PWild :: _, _) => true + | (Syntax.PVar _ :: _, _) => true + | _ => false) + (ListPair.zipEq (patterns, actions))) + in (patterns, occurrences, actions) end (* There's ambiguity between pattern variables and constructors that can only @@ -165,8 +188,9 @@ struct let val refutablePattern = List.find (fn (_, Syntax.PInt _) => true + | (_, Syntax.PCon _) => true | _ => false) - (mapIdx (fn x => x) firstRow) + (enumerate firstRow) in case refutablePattern of NONE => hd actions @@ -176,21 +200,35 @@ struct if i = 0 then (patterns, occurrences) else swap i patterns occurrences + val firstCol = map hd patterns val signatures = map (fn (x, _) => x) (IntMap.toList (foldl (fn (Syntax.PInt i, acc) => IntMap.insert i true acc + | (Syntax.PCon (c, _), acc) => IntMap.insert (valOf (lookupCon c env)) true acc | (_, acc) => acc) IntMap.empty - (map hd patterns))) + firstCol)) + val nCons = + case List.find (fn (Syntax.PCon _) => true | _ => false) firstCol of + SOME (Syntax.PCon (name, _)) => nConstructors name env + | _ => ~1 + val defaultCase = + if length signatures = nCons + then NONE + else SOME (compilePatternMatching env (default (patterns, occurrences, actions))) + val switchOperand = + if nCons < 0 + then hd occurrences + else Syntax.LSelect (0, hd occurrences) in Syntax.LSwitch - ( hd occurrences + ( switchOperand , map (fn i => (i, compilePatternMatching env (specialize env i (patterns, occurrences, actions)))) signatures - , compilePatternMatching env (default (patterns, occurrences, actions)) + , defaultCase ) end end @@ -258,13 +296,25 @@ struct | Syntax.ELet (Syntax.DDatatype (name, cons) :: decls, body) => let val env = bindDataCons cons env - fun go _ [] = [] - | go i ((name, _) :: cons) = - let val v = Gensym.new () - in (lookupVar name env, v, Syntax.LRecord [Syntax.LInt i, Syntax.LVar v]) :: go (i + 1) cons - end + val funs = + List.mapPartial + (fn (_, (_, NONE)) => NONE + | (i, (name, _)) => + let val v = Gensym.new () + in SOME (lookupVar name env, v, Syntax.LRecord [Syntax.LInt i, Syntax.LVar v]) + end) + (enumerate cons) + val vals = + List.mapPartial + (fn (i, (name, NONE)) => SOME (lookupVar name env, Syntax.LRecord [Syntax.LInt i]) + | _ => NONE) + (enumerate cons) in - Syntax.LFix (go 0 cons, elab env (Syntax.ELet (decls, body))) + foldl + (fn ((v, x), acc) => + Syntax.LApp (Syntax.LFn (v, acc), x)) + (Syntax.LFix (funs, elab env (Syntax.ELet (decls, body)))) + vals end | Syntax.ELet (Syntax.DVal (pat, v) :: decls, body) => elab env (Syntax.ECase (v, [(pat, Syntax.ELet (decls, body))])) -- cgit v1.3.1