Renamed script to generate asm for multiplication. Added a new script to generate asm for squaring.

This commit is contained in:
Ken MacKay
2014-03-04 00:27:09 -08:00
parent 142b404f02
commit 630d418aca
3 changed files with 316 additions and 158 deletions
-158
View File
@@ -1,158 +0,0 @@
#!/usr/bin/env python
def rx(i):
return i + 2
def ry(i):
return i + 12
#### set up registers
print r'"adiw r30, 10\n\t"'
print r'"adiw r28, 10\n\t"'
for i in xrange(10):
print r'"ld r%s, x+ \n\t"' % (rx(i))
for i in xrange(10):
print r'"ld r%s, y+ \n\t"' % (ry(i))
#### first two multiplications of initial block (x = 0-9, y = 10-19)
print r'"ldi r25, 0 \n\t"'
print ""
print r'"ldi r23, 0 \n\t"'
print r'"mul r2, r12 \n\t"'
print r'"st z+, r0 \n\t"'
print r'"mov r22, r1 \n\t"'
print ""
print r'"ldi r24, 0 \n\t"'
print r'"mul r2, r13 \n\t"'
print r'"add r22, r0 \n\t"'
print r'"adc r23, r1 \n\t"'
print r'"mul r3, r12 \n\t"'
print r'"add r22, r0 \n\t"'
print r'"adc r23, r1 \n\t"'
print r'"adc r24, r25 \n\t"'
print r'"st z+, r22 \n\t"'
print ""
#### rest of initial block, with moving accumulator registers
acc = [23, 24, 22]
for r in xrange(2, 10):
print r'"ldi r%s, 0 \n\t"' % (acc[2])
for i in xrange(0, r+1):
print r'"mul r%s, r%s \n\t"' % (rx(i), ry(r - i))
print r'"add r%s, r0 \n\t"' % (acc[0])
print r'"adc r%s, r1 \n\t"' % (acc[1])
print r'"adc r%s, r25 \n\t"' % (acc[2])
print r'"st z+, r%s \n\t"' % (acc[0])
print ""
acc = acc[1:] + acc[:1]
for r in xrange(1, 9):
print r'"ldi r%s, 0 \n\t"' % (acc[2])
for i in xrange(0, 10-r):
print r'"mul r%s, r%s \n\t"' % (rx(r+i), ry(9 - i))
print r'"add r%s, r0 \n\t"' % (acc[0])
print r'"adc r%s, r1 \n\t"' % (acc[1])
print r'"adc r%s, r25 \n\t"' % (acc[2])
print r'"st z+, r%s \n\t"' % (acc[0])
print ""
acc = acc[1:] + acc[:1]
print r'"mul r%s, r%s \n\t"' % (rx(9), ry(9))
print r'"add r%s, r0 \n\t"' % (acc[0])
print r'"adc r%s, r1 \n\t"' % (acc[1])
print r'"st z+, r%s \n\t"' % (acc[0])
print r'"st z+, r%s \n\t"' % (acc[1])
print ""
#### reset y and z pointers (x still points to left + 10)
print r'"sbiw r30, 30\n\t"'
print r'"sbiw r28, 20\n\t"'
#### load y registers
for i in xrange(10):
print r'"ld r%s, y+ \n\t"' % (ry(i))
print ""
#### do x = 0-9, y = 0-9 multiplications
print r'"ldi r23, 0 \n\t"'
print r'"mul r2, r12 \n\t"'
print r'"st z+, r0 \n\t"'
print r'"mov r22, r1 \n\t"'
print ""
print r'"ldi r24, 0 \n\t"'
print r'"mul r2, r13 \n\t"'
print r'"add r22, r0 \n\t"'
print r'"adc r23, r1 \n\t"'
print r'"mul r3, r12 \n\t"'
print r'"add r22, r0 \n\t"'
print r'"adc r23, r1 \n\t"'
print r'"adc r24, r25 \n\t"'
print r'"st z+, r22 \n\t"'
print ""
acc = [23, 24, 22]
for r in xrange(2, 10):
print r'"ldi r%s, 0 \n\t"' % (acc[2])
for i in xrange(0, r+1):
print r'"mul r%s, r%s \n\t"' % (rx(i), ry(r - i))
print r'"add r%s, r0 \n\t"' % (acc[0])
print r'"adc r%s, r1 \n\t"' % (acc[1])
print r'"adc r%s, r25 \n\t"' % (acc[2])
print r'"st z+, r%s \n\t"' % (acc[0])
print ""
acc = acc[1:] + acc[:1]
#### now we need to start shifting x and loading from z
x_regs = [2, 3, 4, 5, 6, 7, 8, 9, 10, 11]
for r in xrange(0, 10):
x_regs = x_regs[1:] + x_regs[:1]
print r'"ld r%s, x+ \n\t"' % (x_regs[9]) # load next byte of left
print r'"ldi r%s, 0 \n\t"' % (acc[2])
for i in xrange(0, 10):
print r'"mul r%s, r%s \n\t"' % (x_regs[i], ry(9 - i))
print r'"add r%s, r0 \n\t"' % (acc[0])
print r'"adc r%s, r1 \n\t"' % (acc[1])
print r'"adc r%s, r25 \n\t"' % (acc[2])
print r'"ld r0, z \n\t"' # load stored value from initial block, and add to accumulator (note z does not increment)
print r'"add r%s, r0 \n\t"' % (acc[0])
print r'"adc r%s, r25 \n\t"' % (acc[1])
print r'"adc r%s, r25 \n\t"' % (acc[2])
print r'"st z+, r%s \n\t"' % (acc[0]) # store next byte (z increments)
print ""
acc = acc[1:] + acc[:1]
# done shifting x, start shifting y
y_regs = [12, 13, 14, 15, 16, 17, 18, 19, 20, 21]
for r in xrange(0, 10):
y_regs = y_regs[1:] + y_regs[:1]
print r'"ld r%s, y+ \n\t"' % (y_regs[9]) # load next byte of right
print r'"ldi r%s, 0 \n\t"' % (acc[2])
for i in xrange(0, 10):
print r'"mul r%s, r%s \n\t"' % (x_regs[i], y_regs[9 -i])
print r'"add r%s, r0 \n\t"' % (acc[0])
print r'"adc r%s, r1 \n\t"' % (acc[1])
print r'"adc r%s, r25 \n\t"' % (acc[2])
print r'"ld r0, z \n\t"' # load stored value from initial block, and add to accumulator (note z does not increment)
print r'"add r%s, r0 \n\t"' % (acc[0])
print r'"adc r%s, r25 \n\t"' % (acc[1])
print r'"adc r%s, r25 \n\t"' % (acc[2])
print r'"st z+, r%s \n\t"' % (acc[0]) # store next byte (z increments)
print ""
acc = acc[1:] + acc[:1]
# done both shifts, do remaining corner
for r in xrange(1, 9):
print r'"ldi r%s, 0 \n\t"' % (acc[2])
for i in xrange(0, 10-r):
print r'"mul r%s, r%s \n\t"' % (x_regs[r+i], y_regs[9 - i])
print r'"add r%s, r0 \n\t"' % (acc[0])
print r'"adc r%s, r1 \n\t"' % (acc[1])
print r'"adc r%s, r25 \n\t"' % (acc[2])
print r'"st z+, r%s \n\t"' % (acc[0])
print ""
acc = acc[1:] + acc[:1]
print r'"mul r%s, r%s \n\t"' % (x_regs[9], y_regs[9])
print r'"add r%s, r0 \n\t"' % (acc[0])
print r'"adc r%s, r1 \n\t"' % (acc[1])
print r'"st z+, r%s \n\t"' % (acc[0])
print r'"st z+, r%s \n\t"' % (acc[1])
print r'"eor r1, r1 \n\t"'
Executable
+162
View File
@@ -0,0 +1,162 @@
#!/usr/bin/env python
def rx(i):
return i + 2
def ry(i):
return i + 12
def emit(line, *args):
s = '"' + line + r' \n\t"'
print s % args
#### set up registers
emit("adiw r30, 10")
emit("adiw r28, 10")
for i in xrange(10):
emit("ld r%s, x+", rx(i))
for i in xrange(10):
emit("ld r%s, y+", ry(i))
#### first two multiplications of initial block (x = 0-9, y = 10-19)
emit("ldi r25, 0")
print ""
emit("ldi r23, 0")
emit("mul r2, r12")
emit("st z+, r0")
emit("mov r22, r1")
print ""
emit("ldi r24, 0")
emit("mul r2, r13")
emit("add r22, r0")
emit("adc r23, r1")
emit("mul r3, r12")
emit("add r22, r0")
emit("adc r23, r1")
emit("adc r24, r25")
emit("st z+, r22")
print ""
#### rest of initial block, with moving accumulator registers
acc = [23, 24, 22]
for r in xrange(2, 10):
emit("ldi r%s, 0", acc[2])
for i in xrange(0, r+1):
emit("mul r%s, r%s", rx(i), ry(r - i))
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("adc r%s, r25", acc[2])
emit("st z+, r%s", acc[0])
print ""
acc = acc[1:] + acc[:1]
for r in xrange(1, 9):
emit("ldi r%s, 0", acc[2])
for i in xrange(0, 10-r):
emit("mul r%s, r%s", rx(r+i), ry(9 - i))
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("adc r%s, r25", acc[2])
emit("st z+, r%s", acc[0])
print ""
acc = acc[1:] + acc[:1]
emit("mul r%s, r%s", rx(9), ry(9))
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("st z+, r%s", acc[0])
emit("st z+, r%s", acc[1])
print ""
#### reset y and z pointers (x still points to left + 10)
emit("sbiw r30, 30")
emit("sbiw r28, 20")
#### load y registers
for i in xrange(10):
emit("ld r%s, y+", ry(i))
print ""
#### do x = 0-9, y = 0-9 multiplications
emit("ldi r23, 0")
emit("mul r2, r12")
emit("st z+, r0")
emit("mov r22, r1")
print ""
emit("ldi r24, 0")
emit("mul r2, r13")
emit("add r22, r0")
emit("adc r23, r1")
emit("mul r3, r12")
emit("add r22, r0")
emit("adc r23, r1")
emit("adc r24, r25")
emit("st z+, r22")
print ""
acc = [23, 24, 22]
for r in xrange(2, 10):
emit("ldi r%s, 0", acc[2])
for i in xrange(0, r+1):
emit("mul r%s, r%s", rx(i), ry(r - i))
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("adc r%s, r25", acc[2])
emit("st z+, r%s", acc[0])
print ""
acc = acc[1:] + acc[:1]
#### now we need to start shifting x and loading from z
x_regs = [2, 3, 4, 5, 6, 7, 8, 9, 10, 11]
for r in xrange(0, 10):
x_regs = x_regs[1:] + x_regs[:1]
emit("ld r%s, x+", x_regs[9]) # load next byte of left
emit("ldi r%s, 0", acc[2])
for i in xrange(0, 10):
emit("mul r%s, r%s", x_regs[i], ry(9 - i))
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("adc r%s, r25", acc[2])
emit("ld r0, z") # load stored value from initial block, and add to accumulator (note z does not increment)
emit("add r%s, r0", acc[0])
emit("adc r%s, r25", acc[1])
emit("adc r%s, r25", acc[2])
emit("st z+, r%s", acc[0]) # store next byte (z increments)
print ""
acc = acc[1:] + acc[:1]
# done shifting x, start shifting y
y_regs = [12, 13, 14, 15, 16, 17, 18, 19, 20, 21]
for r in xrange(0, 10):
y_regs = y_regs[1:] + y_regs[:1]
emit("ld r%s, y+", y_regs[9]) # load next byte of right
emit("ldi r%s, 0", acc[2])
for i in xrange(0, 10):
emit("mul r%s, r%s", x_regs[i], y_regs[9 -i])
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("adc r%s, r25", acc[2])
emit("ld r0, z") # load stored value from initial block, and add to accumulator (note z does not increment)
emit("add r%s, r0", acc[0])
emit("adc r%s, r25", acc[1])
emit("adc r%s, r25", acc[2])
emit("st z+, r%s", acc[0]) # store next byte (z increments)
print ""
acc = acc[1:] + acc[:1]
# done both shifts, do remaining corner
for r in xrange(1, 9):
emit("ldi r%s, 0", acc[2])
for i in xrange(0, 10-r):
emit("mul r%s, r%s", x_regs[r+i], y_regs[9 - i])
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("adc r%s, r25", acc[2])
emit("st z+, r%s", acc[0])
print ""
acc = acc[1:] + acc[:1]
emit("mul r%s, r%s", x_regs[9], y_regs[9])
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("st z+, r%s", acc[0])
emit("st z+, r%s", acc[1])
emit("eor r1, r1")
Executable
+154
View File
@@ -0,0 +1,154 @@
#!/usr/bin/env python
def r(i):
return i + 2
def emit(line, *args):
s = '"' + line + r' \n\t"'
print s % args
#### set up registers
for i in xrange(20):
emit("ld r%s, x+", r(i))
#### first few columns
emit("ldi r27, 0") # zero register
print ""
emit("ldi r23, 0")
emit("mul r2, r2")
emit("st z+, r0")
emit("mov r22, r1")
print ""
emit("ldi r24, 0")
emit("mul r2, r3")
emit("lsl r0")
emit("rol r1")
emit("adc r24, r27") # put carry bit in r24
emit("add r22, r0")
emit("adc r23, r1")
emit("adc r24, r27")
emit("st z+, r22")
print ""
emit("ldi r22, 0")
emit("mul r2, r4")
emit("lsl r0")
emit("rol r1")
emit("adc r22, r27") # put carry bit in r22
emit("add r23, r0")
emit("adc r24, r1")
emit("adc r22, r27")
emit("mul r3, r3")
emit("add r23, r0")
emit("adc r24, r1")
emit("adc r22, r27")
emit("st z+, r23")
print ""
acc = [23, 24, 22]
old_acc = [25, 26]
for i in xrange(3, 20):
emit("ldi r%s, 0", acc[0])
emit("ldi r%s, 0", old_acc[0])
emit("ldi r%s, 0", old_acc[1])
tmp = [acc[1], acc[2]]
acc = [acc[0], old_acc[0], old_acc[1]]
old_acc = tmp
# gather non-equal words
for j in xrange(0, (i+1)//2):
emit("mul r%s, r%s", r(j), r(i-j))
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("adc r%s, r27", acc[2])
# multiply by 2
emit("lsl r%s", acc[0])
emit("rol r%s", acc[1])
emit("rol r%s", acc[2])
# add equal word (if any)
if ((i+1) % 2) != 0:
emit("mul r%s, r%s", r(i//2), r(i//2))
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("adc r%s, r27", acc[2])
# add old accumulator
emit("add r%s, r%s", acc[0], old_acc[0])
emit("adc r%s, r%s", acc[1], old_acc[1])
emit("adc r%s, r27", acc[2])
# store
emit("st z+, r%s", acc[0])
print ""
for i in xrange(1, 17):
emit("ldi r%s, 0", acc[0])
emit("ldi r%s, 0", old_acc[0])
emit("ldi r%s, 0", old_acc[1])
tmp = [acc[1], acc[2]]
acc = [acc[0], old_acc[0], old_acc[1]]
old_acc = tmp
# gather non-equal words
for j in xrange(0, (20-i)//2):
emit("mul r%s, r%s", r(i+j), r(19-j))
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("adc r%s, r27", acc[2])
# multiply by 2
emit("lsl r%s", acc[0])
emit("rol r%s", acc[1])
emit("rol r%s", acc[2])
# add equal word (if any)
if ((20-i) % 2) != 0:
emit("mul r%s, r%s", r(i + (20-i)//2), r(i + (20-i)//2))
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("adc r%s, r27", acc[2])
# add old accumulator
emit("add r%s, r%s", acc[0], old_acc[0])
emit("adc r%s, r%s", acc[1], old_acc[1])
emit("adc r%s, r27", acc[2])
# store
emit("st z+, r%s", acc[0])
print ""
acc = acc[1:] + acc[:1]
emit("ldi r%s, 0", acc[2])
emit("mul r19, r21")
emit("lsl r0")
emit("rol r1")
emit("adc r%s, r27", acc[2])
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("adc r%s, r27", acc[2])
emit("mul r20, r20")
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("adc r%s, r27", acc[2])
emit("st z+, r%s", acc[0])
print ""
acc = acc[1:] + acc[:1]
emit("ldi r%s, 0", acc[2])
emit("mul r20, r21")
emit("lsl r0")
emit("rol r1")
emit("adc r%s, r27", acc[2])
emit("add r%s, r0", acc[0])
emit("adc r%s, r1", acc[1])
emit("adc r%s, r27", acc[2])
emit("st z+, r%s", acc[0])
print ""
emit("mul r21, r21")
emit("add r%s, r0", acc[1])
emit("adc r%s, r1", acc[2])
emit("st z+, r%s", acc[1])
emit("st z+, r%s", acc[2])
emit("eor r1, r1")