Added part 4 into chapter3and4.py

This commit is contained in:
Eli Bendersky 2015-01-28 16:54:53 -08:00
commit 861d8e9c84

View file

@ -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))