Codegen for unary, ifexpr bugfix, mandelbrot sample
This commit is contained in:
parent
74aab348c3
commit
f11cd1aecc
1 changed files with 95 additions and 24 deletions
119
chapter6.py
119
chapter6.py
|
|
@ -126,7 +126,7 @@ class UnaryExprAST(ExprAST):
|
||||||
def __init__(self, op, operand):
|
def __init__(self, op, operand):
|
||||||
self.op = op
|
self.op = op
|
||||||
self.operand = operand
|
self.operand = operand
|
||||||
|
|
||||||
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.op)
|
' ' * indent, self.__class__.__name__, self.op)
|
||||||
|
|
@ -245,7 +245,7 @@ class FunctionAST(ASTNode):
|
||||||
|
|
||||||
def is_anonymous(self):
|
def is_anonymous(self):
|
||||||
return self.proto.name.startswith('_anon')
|
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())
|
||||||
|
|
@ -268,7 +268,7 @@ class Parser(object):
|
||||||
self.token_generator = Lexer(buf).tokens()
|
self.token_generator = Lexer(buf).tokens()
|
||||||
self.cur_tok = None
|
self.cur_tok = None
|
||||||
self._get_next_token()
|
self._get_next_token()
|
||||||
|
|
||||||
if self.cur_tok.kind == TokenKind.EXTERN:
|
if self.cur_tok.kind == TokenKind.EXTERN:
|
||||||
return self._parse_external()
|
return self._parse_external()
|
||||||
elif self.cur_tok.kind == TokenKind.DEF:
|
elif self.cur_tok.kind == TokenKind.DEF:
|
||||||
|
|
@ -372,7 +372,7 @@ class Parser(object):
|
||||||
self._match(TokenKind.ELSE)
|
self._match(TokenKind.ELSE)
|
||||||
else_expr = self._parse_expression()
|
else_expr = self._parse_expression()
|
||||||
return IfExprAST(cond_expr, then_expr, else_expr)
|
return IfExprAST(cond_expr, then_expr, else_expr)
|
||||||
|
|
||||||
# forexpr ::= 'for' identifier '=' expr ',' expr (',' expr)? 'in' expr
|
# forexpr ::= 'for' identifier '=' expr ',' expr (',' expr)? 'in' expr
|
||||||
def _parse_for_expr(self):
|
def _parse_for_expr(self):
|
||||||
self._get_next_token() # consume the 'for'
|
self._get_next_token() # consume the 'for'
|
||||||
|
|
@ -392,7 +392,7 @@ class Parser(object):
|
||||||
self._match(TokenKind.IN)
|
self._match(TokenKind.IN)
|
||||||
body = self._parse_expression()
|
body = self._parse_expression()
|
||||||
return ForExprAST(id_name, start_expr, end_expr, step_expr, body)
|
return ForExprAST(id_name, start_expr, end_expr, step_expr, body)
|
||||||
|
|
||||||
# unary
|
# unary
|
||||||
# ::= primary
|
# ::= primary
|
||||||
# ::= <op> unary
|
# ::= <op> unary
|
||||||
|
|
@ -406,7 +406,7 @@ class Parser(object):
|
||||||
op = self.cur_tok.value
|
op = self.cur_tok.value
|
||||||
self._get_next_token()
|
self._get_next_token()
|
||||||
return UnaryExprAST(op, self._parse_unary())
|
return UnaryExprAST(op, self._parse_unary())
|
||||||
|
|
||||||
# binoprhs ::= (<binop> primary)*
|
# binoprhs ::= (<binop> primary)*
|
||||||
def _parse_binop_rhs(self, expr_prec, lhs):
|
def _parse_binop_rhs(self, expr_prec, lhs):
|
||||||
"""Parse the right-hand-side of a binary expression.
|
"""Parse the right-hand-side of a binary expression.
|
||||||
|
|
@ -449,7 +449,7 @@ class Parser(object):
|
||||||
|
|
||||||
# prototype
|
# prototype
|
||||||
# ::= id '(' id* ')'
|
# ::= id '(' id* ')'
|
||||||
# ::= 'binary' LETTER number? '(' id id ')'
|
# ::= 'binary' LETTER number? '(' id id ')'
|
||||||
def _parse_prototype(self):
|
def _parse_prototype(self):
|
||||||
prec = 30
|
prec = 30
|
||||||
if self.cur_tok.kind == TokenKind.IDENTIFIER:
|
if self.cur_tok.kind == TokenKind.IDENTIFIER:
|
||||||
|
|
@ -555,6 +555,11 @@ class LLVMCodeGenerator(object):
|
||||||
def _codegen_VariableExprAST(self, node):
|
def _codegen_VariableExprAST(self, node):
|
||||||
return self.func_symtab[node.name]
|
return self.func_symtab[node.name]
|
||||||
|
|
||||||
|
def _codegen_UnaryExprAST(self, node):
|
||||||
|
operand = self._codegen(node.operand)
|
||||||
|
func = self.module.get_global('unary{0}'.format(node.op))
|
||||||
|
return self.builder.call(func, [operand], 'unop')
|
||||||
|
|
||||||
def _codegen_BinaryExprAST(self, node):
|
def _codegen_BinaryExprAST(self, node):
|
||||||
lhs = self._codegen(node.lhs)
|
lhs = self._codegen(node.lhs)
|
||||||
rhs = self._codegen(node.rhs)
|
rhs = self._codegen(node.rhs)
|
||||||
|
|
@ -594,7 +599,7 @@ class LLVMCodeGenerator(object):
|
||||||
self.builder.position_at_start(then_bb)
|
self.builder.position_at_start(then_bb)
|
||||||
then_val = self._codegen(node.then_expr)
|
then_val = self._codegen(node.then_expr)
|
||||||
self.builder.branch(merge_bb)
|
self.builder.branch(merge_bb)
|
||||||
|
|
||||||
# Emission of then_val could have modified the current basic block. To
|
# Emission of then_val could have modified the current basic block. To
|
||||||
# properly set up the PHI, remember which block the 'then' part ends in.
|
# properly set up the PHI, remember which block the 'then' part ends in.
|
||||||
then_bb = self.builder.block
|
then_bb = self.builder.block
|
||||||
|
|
@ -603,6 +608,9 @@ class LLVMCodeGenerator(object):
|
||||||
self.builder.function.basic_blocks.append(else_bb)
|
self.builder.function.basic_blocks.append(else_bb)
|
||||||
self.builder.position_at_start(else_bb)
|
self.builder.position_at_start(else_bb)
|
||||||
else_val = self._codegen(node.else_expr)
|
else_val = self._codegen(node.else_expr)
|
||||||
|
|
||||||
|
# Emission of else_val could have modified the current basic block.
|
||||||
|
else_bb = self.builder.block
|
||||||
self.builder.branch(merge_bb)
|
self.builder.branch(merge_bb)
|
||||||
|
|
||||||
# Emit the merge ('ifcnt') block
|
# Emit the merge ('ifcnt') block
|
||||||
|
|
@ -612,7 +620,7 @@ class LLVMCodeGenerator(object):
|
||||||
phi.add_incoming(then_val, then_bb)
|
phi.add_incoming(then_val, then_bb)
|
||||||
phi.add_incoming(else_val, else_bb)
|
phi.add_incoming(else_val, else_bb)
|
||||||
return phi
|
return phi
|
||||||
|
|
||||||
def _codegen_ForExprAST(self, node):
|
def _codegen_ForExprAST(self, node):
|
||||||
# Output this as:
|
# Output this as:
|
||||||
# ...
|
# ...
|
||||||
|
|
@ -695,7 +703,7 @@ class LLVMCodeGenerator(object):
|
||||||
raise CodegenError('Call argument length mismatch', node.callee)
|
raise CodegenError('Call argument length mismatch', node.callee)
|
||||||
call_args = [self._codegen(arg) for arg in node.args]
|
call_args = [self._codegen(arg) for arg in node.args]
|
||||||
return self.builder.call(callee_func, call_args, 'calltmp')
|
return self.builder.call(callee_func, call_args, 'calltmp')
|
||||||
|
|
||||||
def _codegen_PrototypeAST(self, node):
|
def _codegen_PrototypeAST(self, node):
|
||||||
funcname = node.name
|
funcname = node.name
|
||||||
# Create a function type
|
# Create a function type
|
||||||
|
|
@ -731,7 +739,7 @@ class LLVMCodeGenerator(object):
|
||||||
# Create the function skeleton from the prototype.
|
# Create the function skeleton from the prototype.
|
||||||
func = self._codegen(node.proto)
|
func = self._codegen(node.proto)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
# Create the entry BB in the function and set the builder to it.
|
# Create the entry BB in the function and set the builder to it.
|
||||||
bb_entry = func.append_basic_block('entry')
|
bb_entry = func.append_basic_block('entry')
|
||||||
|
|
@ -759,7 +767,7 @@ class KaleidoscopeEvaluator(object):
|
||||||
self._add_builtins(self.codegen.module)
|
self._add_builtins(self.codegen.module)
|
||||||
|
|
||||||
self.target = llvm.Target.from_default_triple()
|
self.target = llvm.Target.from_default_triple()
|
||||||
|
|
||||||
def evaluate(self, codestr, optimize=True, llvmdump=False):
|
def evaluate(self, codestr, optimize=True, llvmdump=False):
|
||||||
"""Evaluate code in codestr.
|
"""Evaluate code in codestr.
|
||||||
|
|
||||||
|
|
@ -769,7 +777,7 @@ class KaleidoscopeEvaluator(object):
|
||||||
# Parse the given code and generate code from it
|
# Parse the given code and generate code from it
|
||||||
ast = self.parser.parse_toplevel(codestr)
|
ast = self.parser.parse_toplevel(codestr)
|
||||||
self.codegen.generate_code(ast)
|
self.codegen.generate_code(ast)
|
||||||
|
|
||||||
if llvmdump:
|
if llvmdump:
|
||||||
print('======== Unoptimized LLVM IR')
|
print('======== Unoptimized LLVM IR')
|
||||||
print(str(self.codegen.module))
|
print(str(self.codegen.module))
|
||||||
|
|
@ -897,12 +905,12 @@ class TestParser(unittest.TestCase):
|
||||||
p.parse_toplevel('def binary% 77(a b) a + b')
|
p.parse_toplevel('def binary% 77(a b) a + b')
|
||||||
ast = p.parse_toplevel('a * 10 % 5 * 10')
|
ast = p.parse_toplevel('a * 10 % 5 * 10')
|
||||||
self._assert_body(ast,
|
self._assert_body(ast,
|
||||||
['Binop', '*',
|
['Binop', '*',
|
||||||
['Binop', '*',
|
['Binop', '*',
|
||||||
['Variable', 'a'],
|
['Variable', 'a'],
|
||||||
['Binop', '%', ['Number', '10'], ['Number', '5']]],
|
['Binop', '%', ['Number', '10'], ['Number', '5']]],
|
||||||
['Number', '10']])
|
['Number', '10']])
|
||||||
|
|
||||||
ast = p.parse_toplevel('a % 20 * 5')
|
ast = p.parse_toplevel('a % 20 * 5')
|
||||||
self._assert_body(ast,
|
self._assert_body(ast,
|
||||||
['Binop', '*',
|
['Binop', '*',
|
||||||
|
|
@ -926,14 +934,77 @@ class TestEvaluator(unittest.TestCase):
|
||||||
self.assertEqual(e.evaluate('10 % 5'), 5)
|
self.assertEqual(e.evaluate('10 % 5'), 5)
|
||||||
self.assertEqual(e.evaluate('100 % 5.5'), 94.5)
|
self.assertEqual(e.evaluate('100 % 5.5'), 94.5)
|
||||||
|
|
||||||
|
def test_custom_unop(self):
|
||||||
|
e = KaleidoscopeEvaluator()
|
||||||
|
e.evaluate('def unary!(a) 0 - a')
|
||||||
|
e.evaluate('def unary^(a) a * a')
|
||||||
|
self.assertEqual(e.evaluate('!10'), -10)
|
||||||
|
self.assertEqual(e.evaluate('^10'), 100)
|
||||||
|
self.assertEqual(e.evaluate('!^10'), -100)
|
||||||
|
self.assertEqual(e.evaluate('^!10'), 100)
|
||||||
|
|
||||||
|
def test_mixed_ops(self):
|
||||||
|
e = KaleidoscopeEvaluator()
|
||||||
|
e.evaluate('def unary!(a) 0 - a')
|
||||||
|
e.evaluate('def unary^(a) a * a')
|
||||||
|
e.evaluate('def binary %(a b) a - b')
|
||||||
|
self.assertEqual(e.evaluate('!10 % !20'), 10)
|
||||||
|
self.assertEqual(e.evaluate('^(!10 % !20)'), 100)
|
||||||
|
|
||||||
|
|
||||||
|
def generate_mandelbrot():
|
||||||
|
e = KaleidoscopeEvaluator()
|
||||||
|
e.evaluate('def unary- (v) 0 - v')
|
||||||
|
e.evaluate('def binary> 10 (lhs rhs) rhs < lhs')
|
||||||
|
e.evaluate('def binary: 1 (x y) y')
|
||||||
|
e.evaluate('''
|
||||||
|
def binary| 5 (lhs rhs)
|
||||||
|
if lhs then 1 else if rhs then 1 else 0
|
||||||
|
''')
|
||||||
|
e.evaluate('''
|
||||||
|
def printdensity(d)
|
||||||
|
if d > 8 then
|
||||||
|
putchard(32) # ' '
|
||||||
|
else if d > 4 then
|
||||||
|
putchard(46) # '.'
|
||||||
|
else if d > 2 then
|
||||||
|
putchard(43) # '+'
|
||||||
|
else
|
||||||
|
putchard(42) # '*'
|
||||||
|
''')
|
||||||
|
e.evaluate('''
|
||||||
|
def mandelconverger(real imag iters creal cimag)
|
||||||
|
if iters > 255 | (real*real + imag*imag > 4) then
|
||||||
|
iters
|
||||||
|
else
|
||||||
|
mandelconverger(real*real - imag*imag + creal,
|
||||||
|
2*real*imag + cimag,
|
||||||
|
iters+1, creal, cimag)
|
||||||
|
''')
|
||||||
|
e.evaluate('''
|
||||||
|
def mandelconverge(real imag)
|
||||||
|
mandelconverger(real, imag, 0, real, imag)
|
||||||
|
''')
|
||||||
|
e.evaluate('''
|
||||||
|
def mandelhelp(xmin xmax xstep ymin ymax ystep)
|
||||||
|
for y = ymin, y < ymax, ystep in (
|
||||||
|
(for x = xmin, x < xmax, xstep in
|
||||||
|
printdensity(mandelconverge(x, y)))
|
||||||
|
: putchard(10))
|
||||||
|
''')
|
||||||
|
e.evaluate('''
|
||||||
|
def mandel(realstart imagstart realmag imagmag)
|
||||||
|
mandelhelp(realstart, realstart+realmag*78, realmag,
|
||||||
|
imagstart, imagstart+imagmag*48, imagmag)
|
||||||
|
''')
|
||||||
|
e.evaluate('mandel(-2.3, -1.3, 0.05, 0.07)')
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
p = Parser()
|
generate_mandelbrot()
|
||||||
print(p.parse_toplevel('def binary% 77(a b) a + b').dump())
|
#p = Parser()
|
||||||
print(p.parse_toplevel('def fra(x t) x % t').dump())
|
#print(p.parse_toplevel('def binary% 77(a b) a + b').dump())
|
||||||
kalei = KaleidoscopeEvaluator()
|
#print(p.parse_toplevel('def fra(x t) x % t').dump())
|
||||||
kalei.evaluate('def binary% 77(a b) a + b')
|
#kalei = KaleidoscopeEvaluator()
|
||||||
print(kalei.evaluate('5 % 10', optimize=False, llvmdump=True))
|
#kalei.evaluate('def binary% 77(a b) a + b')
|
||||||
|
#print(kalei.evaluate('5 % 10', optimize=False, llvmdump=True))
|
||||||
#kalei.evaluate('def foo(a b) for x = 65, x < a, b in putchard(x)')
|
|
||||||
#print(kalei.evaluate('foo(79, 1)', optimize=True, llvmdump=True))
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue