diff options
Diffstat (limited to 'codegen.sml')
| -rw-r--r-- | codegen.sml | 148 |
1 files changed, 148 insertions, 0 deletions
diff --git a/codegen.sml b/codegen.sml new file mode 100644 index 0000000..1f6a2a8 --- /dev/null +++ b/codegen.sml @@ -0,0 +1,148 @@ +structure CodeGen = +struct + fun enumerate l = ListPair.zip (List.tabulate (length l, (fn x => x)), l) + + (* There are 8 registers *) + val tempReg = 7 + + structure VarMap = Map (type k = Syntax.var + val cmp = Int.compare) + + fun cycle (outputs : Syntax.var VarMap.map) (output : Syntax.var) : Syntax.opcode list = + case VarMap.lookup output outputs of + NONE => Syntax.OShuf (output, Syntax.VVar tempReg) :: cycles outputs + | SOME input => + Syntax.OShuf (output, Syntax.VVar input) :: cycle (VarMap.delete output outputs) input + + and cycles (outputs : Syntax.var VarMap.map) : Syntax.opcode list = + case VarMap.lookupMin outputs of + NONE => [] + | SOME (output, input) => + Syntax.OShuf (tempReg, Syntax.VVar input) :: cycle (VarMap.delete output outputs) input + + fun shuffle' (inputs : unit VarMap.map) (outputs : Syntax.var VarMap.map) : Syntax.opcode list = + case VarMap.lookupMin (VarMap.difference outputs inputs) of + NONE => cycles outputs + | SOME (output, input) => + Syntax.OShuf (output, Syntax.VVar input) :: shuffle' (VarMap.delete input inputs) (VarMap.delete output outputs) + + fun shuffle (args : Syntax.value list) : Syntax.opcode list = + let + val outputMap = + VarMap.fromList + (List.mapPartial + (fn (i, Syntax.VVar v) => + if i = v + then NONE + else SOME (i, v) + | _ => NONE) + (enumerate args)) + val inputMap = + VarMap.fromList + (map + (fn (_, input) => (input, ())) + (VarMap.toList outputMap)) + val constants = + List.mapPartial + (fn (_, Syntax.VVar _) => NONE + | (i, constArg) => SOME (Syntax.OShuf (i, constArg))) + (enumerate args) + in + shuffle' inputMap outputMap @ constants + end + + fun buildVarMap (expr : Syntax.cexp) : Syntax.var VarMap.map = + let + val next = ref 0 + fun insert v m = + let val this = !next + in + next := this + 1; + VarMap.insert v this m + end + fun go expr = + case expr of + Syntax.CRecord (_, res, k) => insert res (go k) + | Syntax.CSelect (_, _, res, k) => insert res (go k) + | Syntax.CApp _ => VarMap.empty + | Syntax.CFix (funcs, body) => + let + val funcsVars = + foldl + (fn ((_, args, body), acc) => + let + val argsVars = + foldl + (fn ((i, arg), acc) => VarMap.insert arg i acc) + VarMap.empty + (ListPair.zip + (List.tabulate (length args, fn x => x + 1), + args)) + val bodyVars = buildVarMap body + in VarMap.union (VarMap.union acc argsVars) bodyVars + end) + VarMap.empty + funcs + val bodyVars = go body + in VarMap.union funcsVars bodyVars + end + | Syntax.CPrimop (_, _, res, k) => + let + val resVars = + foldl + (fn (v, acc) => insert v acc) + VarMap.empty + res + val kVars = + foldl + (fn (expr, acc) => VarMap.union acc (go expr)) + VarMap.empty + k + in + VarMap.union resVars kVars + end + in go expr + end + + fun toASM (expr : Syntax.cexp) : Syntax.opcode list = + let + val varMap = buildVarMap expr + fun translate v = valOf (VarMap.lookup v varMap) + fun translateVal (Syntax.VVar v) = Syntax.VVar (translate v) + | translateVal x = x + fun go expr = + case expr of + Syntax.CRecord (args, res, k) => + Syntax.OAlloc (translate res, Syntax.VInt (length args)) + :: 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)) + (arg, []) + path + in rev (Syntax.OPoke (i, translate res, temp) :: ops) + end) + (enumerate args)) + @ toASM k + | Syntax.CSelect (i, arg, res, k) => Syntax.OPeek (translate res, i, translateVal arg) :: toASM k + | Syntax.CApp (func, args) => shuffle (func :: args) @ [Syntax.OCall] + | Syntax.CFix (funcs, body) => + let + val bodyASM = toASM body + val funcsASM = + foldl + (fn ((name, _, body), acc) => + Syntax.OLabel name :: go body @ acc) + [] + funcs + in + bodyASM @ funcsASM + end + | Syntax.CPrimop (Syntax.PExit, [arg], _, _)=> [Syntax.OExit arg] + | _ => raise Fail ("malformed CPS:\n" ^ Syntax.cexpToString expr) + in go expr + end +end |
