Added part 4 into chapter3and4.py
This commit is contained in:
parent
92a2a95560
commit
861d8e9c84
1 changed files with 97 additions and 25 deletions
|
|
@ -1,6 +1,7 @@
|
||||||
# Chapter 3 - Code generation to LLVM IR
|
# Chapter 3 - Code generation to LLVM IR
|
||||||
|
|
||||||
from collections import namedtuple
|
from collections import namedtuple
|
||||||
|
from ctypes import CFUNCTYPE, c_double
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
|
|
||||||
import llvmlite.ir as ir
|
import llvmlite.ir as ir
|
||||||
|
|
@ -146,6 +147,20 @@ class FunctionAST(ASTNode):
|
||||||
self.proto = proto
|
self.proto = proto
|
||||||
self.body = body
|
self.body = body
|
||||||
|
|
||||||
|
_anonymous_function_counter = 0
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create_anonymous(klass, expr):
|
||||||
|
"""Create an anonymous function to hold an expression."""
|
||||||
|
klass._anonymous_function_counter += 1
|
||||||
|
return klass(
|
||||||
|
PrototypeAST('_anon{0}'.format(klass._anonymous_function_counter),
|
||||||
|
[]),
|
||||||
|
expr)
|
||||||
|
|
||||||
|
def is_anonymous(self):
|
||||||
|
return self.proto.name.startswith('_anon')
|
||||||
|
|
||||||
def dump(self, indent=0):
|
def dump(self, indent=0):
|
||||||
s = '{0}{1}[{2}]\n'.format(
|
s = '{0}{1}[{2}]\n'.format(
|
||||||
' ' * indent, self.__class__.__name__, self.proto.dump())
|
' ' * indent, self.__class__.__name__, self.proto.dump())
|
||||||
|
|
@ -317,8 +332,7 @@ class Parser(object):
|
||||||
# toplevel ::= expression
|
# toplevel ::= expression
|
||||||
def _parse_toplevel_expression(self):
|
def _parse_toplevel_expression(self):
|
||||||
expr = self._parse_expression()
|
expr = self._parse_expression()
|
||||||
# Anonymous function
|
return FunctionAST.create_anonymous(expr)
|
||||||
return FunctionAST(PrototypeAST('', []), expr)
|
|
||||||
|
|
||||||
|
|
||||||
class CodegenError(Exception): pass
|
class CodegenError(Exception): pass
|
||||||
|
|
@ -330,7 +344,11 @@ class LLVMCodeGenerator(object):
|
||||||
|
|
||||||
This creates a new LLVM module into which code is generated. The
|
This creates a new LLVM module into which code is generated. The
|
||||||
generate_code() method can be called multiple times. It adds the code
|
generate_code() method can be called multiple times. It adds the code
|
||||||
generated for this node into the module, and returns the module.
|
generated for this node into the module, and returns the IR value for
|
||||||
|
the node.
|
||||||
|
|
||||||
|
At any time, the current LLVM module being constructed can be obtained
|
||||||
|
from the module attribute.
|
||||||
"""
|
"""
|
||||||
self.module = ir.Module()
|
self.module = ir.Module()
|
||||||
|
|
||||||
|
|
@ -341,14 +359,9 @@ class LLVMCodeGenerator(object):
|
||||||
# names to ir.Value.
|
# names to ir.Value.
|
||||||
self.func_symtab = {}
|
self.func_symtab = {}
|
||||||
|
|
||||||
# Counter to disambiguate anonymous functions created from toplevel
|
|
||||||
# expressions.
|
|
||||||
self.anon_function_counter = 0
|
|
||||||
|
|
||||||
def generate_code(self, node):
|
def generate_code(self, node):
|
||||||
assert isinstance(node, (PrototypeAST, FunctionAST))
|
assert isinstance(node, (PrototypeAST, FunctionAST))
|
||||||
self._codegen(node)
|
return self._codegen(node)
|
||||||
return self.module
|
|
||||||
|
|
||||||
def _codegen(self, node):
|
def _codegen(self, node):
|
||||||
"""Node visitor. Dispathces upon node type.
|
"""Node visitor. Dispathces upon node type.
|
||||||
|
|
@ -392,10 +405,6 @@ class LLVMCodeGenerator(object):
|
||||||
|
|
||||||
def _codegen_PrototypeAST(self, node):
|
def _codegen_PrototypeAST(self, node):
|
||||||
funcname = node.name
|
funcname = node.name
|
||||||
if not funcname:
|
|
||||||
funcname = 'anon{0}'.format(self.anon_function_counter)
|
|
||||||
self.anon_function_counter += 1
|
|
||||||
|
|
||||||
# Create a function type
|
# Create a function type
|
||||||
func_ty = ir.FunctionType(ir.DoubleType(),
|
func_ty = ir.FunctionType(ir.DoubleType(),
|
||||||
[ir.DoubleType()] * len(node.argnames))
|
[ir.DoubleType()] * len(node.argnames))
|
||||||
|
|
@ -436,20 +445,83 @@ class LLVMCodeGenerator(object):
|
||||||
return func
|
return func
|
||||||
|
|
||||||
|
|
||||||
|
class KaleidoscopeEvaluator(object):
|
||||||
|
def __init__(self):
|
||||||
|
llvm.initialize()
|
||||||
|
llvm.initialize_native_target()
|
||||||
|
llvm.initialize_native_asmprinter()
|
||||||
|
|
||||||
|
self.codegen = LLVMCodeGenerator()
|
||||||
|
|
||||||
|
self.target = llvm.Target.from_default_triple()
|
||||||
|
self.target_machine = self.target.create_target_machine()
|
||||||
|
|
||||||
|
def evaluate(self, codestr, optimize=True, llvmdump=False):
|
||||||
|
# Parse the given code and generate code from it
|
||||||
|
ast = Parser(codestr).parse_toplevel()
|
||||||
|
self.codegen.generate_code(ast)
|
||||||
|
|
||||||
|
if llvmdump:
|
||||||
|
print('======== Unoptimized LLVM IR')
|
||||||
|
print(str(self.codegen.module))
|
||||||
|
|
||||||
|
# If we're evaluating a definition or extern declaration, don't do
|
||||||
|
# anything else. If we're evaluating an anonymous wrapper for a toplevel
|
||||||
|
# expression, JIT-compile the module and run the function to get its
|
||||||
|
# result.
|
||||||
|
if not (isinstance(ast, FunctionAST) and ast.is_anonymous()):
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Convert LLVM IR into in-memory representation
|
||||||
|
llvmmod = llvm.parse_assembly(str(self.codegen.module))
|
||||||
|
|
||||||
|
# Optimize the module
|
||||||
|
if optimize:
|
||||||
|
pmb = llvm.create_pass_manager_builder()
|
||||||
|
pmb.opt_level = 2
|
||||||
|
pm = llvm.create_module_pass_manager()
|
||||||
|
pmb.populate(pm)
|
||||||
|
pm.run(llvmmod)
|
||||||
|
|
||||||
|
if llvmdump:
|
||||||
|
print('======== Optimized LLVM IR')
|
||||||
|
print(str(llvmmod))
|
||||||
|
|
||||||
|
with llvm.create_mcjit_compiler(llvmmod, self.target_machine) as ee:
|
||||||
|
ee.finalize_object()
|
||||||
|
|
||||||
|
if llvmdump:
|
||||||
|
print('======== Machine code')
|
||||||
|
print(self.target_machine.emit_assembly(llvmmod))
|
||||||
|
|
||||||
|
func = llvmmod.get_function(ast.proto.name)
|
||||||
|
fptr = CFUNCTYPE(c_double)(ee.get_pointer_to_function(func))
|
||||||
|
|
||||||
|
result = fptr()
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
def parse(s):
|
#def parse(s):
|
||||||
return Parser(s).parse_toplevel()
|
#return Parser(s).parse_toplevel()
|
||||||
|
|
||||||
dfoo = parse('def foo(a b) a + 4.1414 * (b < a)')
|
#dfoo = parse('def foo(a b) a + 4.1414 * (b < a)')
|
||||||
dcos = parse('extern cos(a)')
|
#dcos = parse('extern cos(a)')
|
||||||
dcl = parse('2 + foo(3.14, cos(9.91))')
|
#dcl = parse('2 + foo(3.14, cos(9.91))')
|
||||||
|
#dcl2 = parse('3 + foo(3.14, cos(9.91))')
|
||||||
|
|
||||||
cg = LLVMCodeGenerator()
|
#cg = LLVMCodeGenerator()
|
||||||
mod = cg.generate_code(dfoo)
|
#cg.generate_code(dfoo)
|
||||||
mod = cg.generate_code(dcos)
|
#cg.generate_code(dcos)
|
||||||
mod = cg.generate_code(dcl)
|
#cg.generate_code(dcl)
|
||||||
|
#cg.generate_code(dcl2)
|
||||||
|
|
||||||
print(mod)
|
#print(cg.module)
|
||||||
|
|
||||||
llvmmod = llvm.parse_assembly(str(mod))
|
#llvmmod = llvm.parse_assembly(str(cg.module))
|
||||||
llvmmod.verify()
|
#llvmmod.verify()
|
||||||
|
|
||||||
|
kalei = KaleidoscopeEvaluator()
|
||||||
|
print(kalei.evaluate('def adder(a b) a + b'))
|
||||||
|
print(kalei.evaluate('def foo(x) (1+2+x)*(x+(1+2))'))
|
||||||
|
print(kalei.evaluate('adder(foo(4), 5)', optimize=True, llvmdump=True))
|
||||||
Loading…
Add table
Add a link
Reference in a new issue