diff options
| author | Jan Tuomi <jans.tuomi@gmail.com> | 2023-12-01 15:10:41 +0200 |
|---|---|---|
| committer | Jan Tuomi <jans.tuomi@gmail.com> | 2023-12-01 15:10:41 +0200 |
| commit | 7f0614990a6485dccf24e160b491c93c56771b1e (patch) | |
| tree | 79826befff2f7843c63686e84bc276593ad4d52e | |
| parent | e2fd1c7ee759ac8870f21afe5fc4e36c8b808d25 (diff) | |
WIP
| -rw-r--r-- | asm/atk16.py | 4 | ||||
| -rw-r--r-- | asm/test_py_src.py | 2 | ||||
| -rw-r--r-- | src/ast_compiler.py | 199 | ||||
| -rw-r--r-- | src/optimizer.py | 189 |
4 files changed, 205 insertions, 189 deletions
diff --git a/asm/atk16.py b/asm/atk16.py index cdd9f64..2a1c0a1 100644 --- a/asm/atk16.py +++ b/asm/atk16.py @@ -23,3 +23,7 @@ def asm(asm: str): def ord(char: Char) -> int: """Convert char to int""" raise NotImplementedError + +def const(word: Word16) -> ConstWord16: + """When used in an assignment such as `A = const(0xFF)`, stores the value as a globally accessible constant.""" + raise NotImplementedError
\ No newline at end of file diff --git a/asm/test_py_src.py b/asm/test_py_src.py index 676f8cd..619d22e 100644 --- a/asm/test_py_src.py +++ b/asm/test_py_src.py @@ -8,6 +8,8 @@ GRAPHICS_SPRITE_MODE: ConstWord16 = 2 GRAPHICS_MODE_PP: ConstWord16 = 0x17 TEXT_MEM_PP: ConstWord16 = 0x19 +TEXT_MEM_PP: ConstWord16 = 0x19 + # asm( # "@label keyboard_isr" # " spu RA" diff --git a/src/ast_compiler.py b/src/ast_compiler.py index f0d0c79..99b4176 100644 --- a/src/ast_compiler.py +++ b/src/ast_compiler.py @@ -7,7 +7,7 @@ import ast from dataclasses import dataclass from typing import Literal, Set, cast, Any, TypeAlias, TypeVar from collections import OrderedDict -from tokenizer import tokenize + if len(sys.argv) != 3: print("usage: ast_compiler.py <infile.py> <outfile.atk16>") @@ -217,10 +217,6 @@ class Compiler(ast.NodeVisitor): self.emit("hlt") def visit_Expr(self, expr: ast.Expr): - # if type(expr.value) != ast.Call: - # self.emit("; NOP top-level expression") - # return - self.visit(expr.value) with self.allocated_reg() as reg: @@ -679,185 +675,15 @@ class Compiler(ast.NodeVisitor): raise NotImplementedError("Unhandled AnnAssign:\n" + ast.dump(node, indent=4)) -class Optimizer: - def __init__(self): - pass - - def optimize(self, asm_str: str): - asm = asm_str.split("\n") - asm = [row.strip() for row in asm] - asm = [self.strip_comment(row) for row in asm] - asm = [row for row in asm if not len(row) == 0] - asm = [tokenize(row, retain_curlies=True) for row in asm] - asm = self.compact_spu_spo_pattern(asm) - - # TODO: not safe when setting loading SP and FP - # ldr RA SP - # mov RA FP - # gets optimized to - # ldr RA FP - #asm = self.compact_target_mov_pattern(asm) - #asm = self.compact_target_mov_pattern(asm) - - # TODO: not safe at all. E.g. breaks a while True: pass loop - #asm = self.compact_mov_source_pattern(asm) - #asm = self.compact_mov_source_pattern(asm) - - asm = self.compact_spu_load_spo_pattern(asm) - # asm = self.convert_alr_to_ali(asm) - result = "\n".join([format_asm_row(" ".join(row)) for row in asm]) - - return result - - def strip_comment(self, row: str): - ret: str = "" - for c in row: - if c == ";": break - ret += c - - return ret - - def compact_spu_spo_pattern(self, asm: list[list[str]]) -> list[list[str]]: - i = 0 - result: list[list[str]] = [] - while i < len(asm): - current = asm[i] - next = asm[i + 1] if i + 1 < len(asm) else None - i += 1 - if current[0].startswith("@") or next is None: - result.append(current) - continue - - if current[0] == "spu" and next[0] == "spo": - print("=== compact_spu_spo_pattern") - print(current) - print(next) - - arg_current = current[1] - arg_next = next[1] - - print("=== arg_current:", arg_current) - print("=== arg_next:", arg_next) - - if arg_current == arg_next: - print("=== pass") - pass # remove both spu and spo - else: - print("=== mov") - mov = ["mov", arg_current, arg_next] - result.append(mov) - print("=== mov:", mov) - - print("=== last of result:", result[len(result) - 1]) - print() - i += 1 - continue - - result.append(current) - - return result - - def compact_target_mov_pattern(self, asm: list[list[str]]) -> list[list[str]]: - ops_with_target = ["ldi", "ldr", "add", "sub", "addi", "subi", "and", "or", "xor", "sll", "slr", "sar", "slli", "slri", "sari", "inc", "dec", "mov", "ali", "alr", "lpc"] - i = 0 - result: list[list[str]] = [] - while i < len(asm): - current = asm[i] - next = asm[i + 1] if i + 1 < len(asm) else None - i += 1 - if current[0].startswith("@") or next is None: - result.append(current) - continue - - if current[0] in ops_with_target and next[0] == "mov": - - op_op, op_operands, op_target_reg = current[0], current[1:len(current) - 1], current[len(current) - 1] - mov_from_reg, mov_to_reg = next[1], next[2] - - if op_target_reg == mov_from_reg: - ret = [op_op, *op_operands, mov_to_reg] - result.append(ret) - i += 1 - continue - - result.append(current) - - return result - - def compact_mov_source_pattern(self, asm: list[list[str]]) -> list[list[str]]: - ops_with_source = ["str", "ldr", "add", "sub", "addi", "subi", "and", "or", "xor", "sll", "slr", "sar", "slli", "slri", "sari", "mov", "ali", "alr"] - i = 0 - result: list[list[str]] = [] - while i < len(asm): - current = asm[i] - next = asm[i + 1] if i + 1 < len(asm) else None - i += 1 - if current[0].startswith("@") or next is None: - result.append(current) - continue - - if current[0] == "mov" and next[0] in ops_with_source: - mov_from_reg, mov_to_reg = current[1], current[2] - op_op, op_source_reg, op_operands = next[0], next[1], next[2:] - - if op_source_reg == mov_to_reg: - ret = [op_op, op_source_reg, *op_operands] - result.append(ret) - i += 1 - continue - - result.append(current) - - return result - - def compact_spu_load_spo_pattern(self, asm: list[list[str]]) -> list[list[str]]: - # spu RA - # ldi int_7 RA - # ldr RA RB - # spo RA - # OR - # spu RA - # ldi 3 RB - # spo RA - - i = 0 - result: list[list[str]] = [] - while i < len(asm): - instr0 = asm[i] - instr1 = asm[i + 1] if i + 1 < len(asm) else None - instr2 = asm[i + 2] if i + 2 < len(asm) else None - instr3 = asm[i + 3] if i + 3 < len(asm) else None - i += 1 - if instr0[0].startswith("@") or instr1 is None or instr2 is None or instr3 is None: - result.append(instr0) - continue - - if instr0[0] == "spu" and instr1[0] == "ldi" and instr2[0] == "ldr" and instr3[0] == "spo": - spu_op, spu_reg = instr0 - ldi_op, ldi_imm, ldi_target_reg = instr1 - ldr_op, ldr_from_reg, ldr_to_reg = instr2 - spo_op, spo_reg = instr3 - - if spu_reg == spo_reg and ldi_target_reg == ldr_from_reg and ldr_from_reg != ldr_to_reg: - ret0 = f"ldi {ldi_imm} {ldr_to_reg}".split() - ret1 = f"ldr {ldr_to_reg} {ldr_to_reg}".split() - result.append(ret0) - result.append(ret1) - i += 3 - continue - - elif instr0[0] == "spu" and instr1[0] == "ldi" and instr2[0] == "spo": - spu_op, spu_reg = instr0 - ldi_op, ldi_imm, ldi_target_reg = instr1 - spo_op, spo_reg = instr2 - - if spu_reg == spo_reg and ldi_target_reg != spu_reg: - ret0 = f"ldi {ldi_imm} {ldi_target_reg}".split() - result.append(ret0) - i += 2 - continue - - result.append(instr0) + def get_module_exports(self, node: ast.Module): + result: list[str] = [] + for stmt in node.body: + match stmt: + case ast.FunctionDef(name): + result.append(name) + case _: + # only function exported since recognizing which assignments are constants is kinda hard + pass return result @@ -873,11 +699,6 @@ asm_out = compiler.compile( source_py, ) -# optimizer = Optimizer() -# asm_out_optimized = optimizer.optimize(asm_out) - with open(outfile_path, "w") as f: f.write(asm_out) -# with open(f"{outfile_path}_optimized", "w") as f: -# f.write(asm_out_optimized) diff --git a/src/optimizer.py b/src/optimizer.py new file mode 100644 index 0000000..0ee2837 --- /dev/null +++ b/src/optimizer.py @@ -0,0 +1,189 @@ +from tokenizer import tokenize + +def format_asm_row(asm: str) -> str: + if not (asm.startswith("@") or asm.startswith(";")) and not asm.startswith(" ") and len(asm) > 0: + return " " + asm + else: + return asm + +class Optimizer: + def __init__(self): + pass + + def optimize(self, asm_str: str): + asm = asm_str.split("\n") + asm = [row.strip() for row in asm] + asm = [self.strip_comment(row) for row in asm] + asm = [row for row in asm if not len(row) == 0] + asm = [tokenize(row, retain_curlies=True) for row in asm] + asm = self.compact_spu_spo_pattern(asm) + + # TODO: not safe when setting loading SP and FP + # ldr RA SP + # mov RA FP + # gets optimized to + # ldr RA FP + #asm = self.compact_target_mov_pattern(asm) + #asm = self.compact_target_mov_pattern(asm) + + # TODO: not safe at all. E.g. breaks a while True: pass loop + #asm = self.compact_mov_source_pattern(asm) + #asm = self.compact_mov_source_pattern(asm) + + asm = self.compact_spu_load_spo_pattern(asm) + # asm = self.convert_alr_to_ali(asm) + result = "\n".join([format_asm_row(" ".join(row)) for row in asm]) + + return result + + def strip_comment(self, row: str): + ret: str = "" + for c in row: + if c == ";": break + ret += c + + return ret + + def compact_spu_spo_pattern(self, asm: list[list[str]]) -> list[list[str]]: + i = 0 + result: list[list[str]] = [] + while i < len(asm): + current = asm[i] + next = asm[i + 1] if i + 1 < len(asm) else None + i += 1 + if current[0].startswith("@") or next is None: + result.append(current) + continue + + if current[0] == "spu" and next[0] == "spo": + print("=== compact_spu_spo_pattern") + print(current) + print(next) + + arg_current = current[1] + arg_next = next[1] + + print("=== arg_current:", arg_current) + print("=== arg_next:", arg_next) + + if arg_current == arg_next: + print("=== pass") + pass # remove both spu and spo + else: + print("=== mov") + mov = ["mov", arg_current, arg_next] + result.append(mov) + print("=== mov:", mov) + + print("=== last of result:", result[len(result) - 1]) + print() + i += 1 + continue + + result.append(current) + + return result + + def compact_target_mov_pattern(self, asm: list[list[str]]) -> list[list[str]]: + ops_with_target = ["ldi", "ldr", "add", "sub", "addi", "subi", "and", "or", "xor", "sll", "slr", "sar", "slli", "slri", "sari", "inc", "dec", "mov", "ali", "alr", "lpc"] + i = 0 + result: list[list[str]] = [] + while i < len(asm): + current = asm[i] + next = asm[i + 1] if i + 1 < len(asm) else None + i += 1 + if current[0].startswith("@") or next is None: + result.append(current) + continue + + if current[0] in ops_with_target and next[0] == "mov": + + op_op, op_operands, op_target_reg = current[0], current[1:len(current) - 1], current[len(current) - 1] + mov_from_reg, mov_to_reg = next[1], next[2] + + if op_target_reg == mov_from_reg: + ret = [op_op, *op_operands, mov_to_reg] + result.append(ret) + i += 1 + continue + + result.append(current) + + return result + + def compact_mov_source_pattern(self, asm: list[list[str]]) -> list[list[str]]: + ops_with_source = ["str", "ldr", "add", "sub", "addi", "subi", "and", "or", "xor", "sll", "slr", "sar", "slli", "slri", "sari", "mov", "ali", "alr"] + i = 0 + result: list[list[str]] = [] + while i < len(asm): + current = asm[i] + next = asm[i + 1] if i + 1 < len(asm) else None + i += 1 + if current[0].startswith("@") or next is None: + result.append(current) + continue + + if current[0] == "mov" and next[0] in ops_with_source: + mov_from_reg, mov_to_reg = current[1], current[2] + op_op, op_source_reg, op_operands = next[0], next[1], next[2:] + + if op_source_reg == mov_to_reg: + ret = [op_op, op_source_reg, *op_operands] + result.append(ret) + i += 1 + continue + + result.append(current) + + return result + + def compact_spu_load_spo_pattern(self, asm: list[list[str]]) -> list[list[str]]: + # spu RA + # ldi int_7 RA + # ldr RA RB + # spo RA + # OR + # spu RA + # ldi 3 RB + # spo RA + + i = 0 + result: list[list[str]] = [] + while i < len(asm): + instr0 = asm[i] + instr1 = asm[i + 1] if i + 1 < len(asm) else None + instr2 = asm[i + 2] if i + 2 < len(asm) else None + instr3 = asm[i + 3] if i + 3 < len(asm) else None + i += 1 + if instr0[0].startswith("@") or instr1 is None or instr2 is None or instr3 is None: + result.append(instr0) + continue + + if instr0[0] == "spu" and instr1[0] == "ldi" and instr2[0] == "ldr" and instr3[0] == "spo": + spu_op, spu_reg = instr0 + ldi_op, ldi_imm, ldi_target_reg = instr1 + ldr_op, ldr_from_reg, ldr_to_reg = instr2 + spo_op, spo_reg = instr3 + + if spu_reg == spo_reg and ldi_target_reg == ldr_from_reg and ldr_from_reg != ldr_to_reg: + ret0 = f"ldi {ldi_imm} {ldr_to_reg}".split() + ret1 = f"ldr {ldr_to_reg} {ldr_to_reg}".split() + result.append(ret0) + result.append(ret1) + i += 3 + continue + + elif instr0[0] == "spu" and instr1[0] == "ldi" and instr2[0] == "spo": + spu_op, spu_reg = instr0 + ldi_op, ldi_imm, ldi_target_reg = instr1 + spo_op, spo_reg = instr2 + + if spu_reg == spo_reg and ldi_target_reg != spu_reg: + ret0 = f"ldi {ldi_imm} {ldi_target_reg}".split() + result.append(ret0) + i += 2 + continue + + result.append(instr0) + + return result |
