module Realization.Nvidia.SM86.LinearStepLoopSM86 import Accelerator.SM86.Control import Accelerator.SM86.Immediate import Accelerator.SM86.Instruction import Accelerator.SM86.Types import Realization.Nvidia.SM86.LinearStepSM86 import Std.Natural -- The linear step with LOOPS: the same computation as -- Realization.Nvidia.SM86.LinearStepSM86 -- the same operations on the same -- operands in the same order, so the record is the specification's word for -- word -- in a program of constant length and constant register demand. -- Where the unrolled program keeps W, x, t, dW and W' in registers (848 of -- them for 16 x 16; the register file admits 38 x 1 .. 1 x 57), this one -- walks the outputs in an outer loop and the inputs in two inner loops, -- loading each element when it is used and storing each result when it is -- made, with 64-bit addresses from IMAD.WIDE (index x 4 + the parameter -- word's address). The loops close with ISETP.GT (an unsigned compare -- -- Checked.IntegerCompareProbe) and a backward BRA. -- -- prologue (tid, the lane guard, the parameter words; LinearStepSM86) -- R11 = -eta R12 = 0.5 R13 = loss = +0 R35 = 4 R14 = i = 0 R16 = e = 0 -- outer: R28 = y = +0 R15 = j = 0 -- first: W_e, x_j loaded; y = fma(W_e, x_j, y); e += 1; j += 1; again while j <= k-1 -- y_i stored; t_i loaded; d = y + (-t_i); loss = fma(d, d, loss); e -= k; j = 0 -- second: x_j, W_e loaded; dW = d x x_j; W' = fma(dW, -eta, W_e); both stored; -- e += 1; j += 1; again while j <= k-1 -- i += 1; again while i <= m-1 -- loss = 0.5 x loss, stored at the record's word m (i = m by then); exit -- -- Each backward branch's offset is its loop's length in bytes, negated, -- counted from the loop's own block (llLoop). def llR11 : Nat = 11 def llHalfRegister : Nat = 12 def llLoss : Nat = 13 def llI : Nat = 14 def llJ : Nat = 15 def llE : Nat = 16 def llWAddress : Nat = 18 def llXAddress : Nat = 20 def llYAddress : Nat = 22 def llTAddress : Nat = 24 def llOutAddress : Nat = 26 def llY : Nat = 28 def llW : Nat = 29 def llX : Nat = 30 def llT : Nat = 31 def llD : Nat = 32 def llDW : Nat = 33 def llWp : Nat = 34 def llFour : Nat = 35 def llP1 : (family SM86Predicate) = (constructor SM86Predicate SM86Predicate1) def llP2 : (family SM86Predicate) = (constructor SM86Predicate SM86Predicate2) def llMove = (lambda unrestricted destination : Nat . (lambda unrestricted value : Nat . (constructor SM86InstructionBody SM86MoveImmediate (lsR destination) (lsU value) lsControl))) def llAddImmediate = (lambda unrestricted destination : Nat . (lambda unrestricted value : Nat . (constructor SM86InstructionBody SM86IntegerAddThreeImmediate (lsR destination) (lsR destination) (lsU value) lsControl))) -- the pair at `destination` = index x 4 + the address in parameter word `parameter` def llAddressOf = (lambda unrestricted destination : Nat . (lambda unrestricted index : Nat . (lambda unrestricted parameter : Nat . (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (lsR destination) (lsR index) (lsR llFour) (byte 0) (lsU parameter) lsControl)))) def llGreater = (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted source : Nat . (lambda unrestricted bound : Nat . (constructor SM86InstructionBody SM86PredicateGreaterThanImmediate predicate (lsR source) (lsU bound) lsControl)))) -- @!P BRA back `instructions` (counted from the branch's successor) def llBackUnless = (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted instructions : Nat . (lambda unrestricted tail : (family SM86Program) . (constructor SM86Program SM86ProgramNext (sm86NegatedPredicatedInstruction predicate (constructor SM86InstructionBody SM86Branch (lsU (naturalSaturatingSubtract 4294967296 (naturalMultiply 16 instructions))) (sm86Unsigned32 (byte 255) (byte 255) (byte 131) (byte 3)) sm86BranchControl)) tail)))) -- the number of instructions a block prepends def llLength = (lambda unrestricted block : (pi unrestricted tail : (family SM86Program) . (family SM86Program)) . (eliminate SM86Program (lambda unrestricted current : (family SM86Program) . Nat) (block (constructor SM86Program SM86ProgramEnd)) (branch SM86ProgramEnd . zero) (branch SM86ProgramNext head rest induction . (succ induction)))) -- A loop: the block, then @!P BRA back to its first instruction -- the -- distance is the block's own length and the branch, so no count is written -- by hand. def llLoop = (lambda unrestricted predicate : (family SM86Predicate) . (lambda unrestricted block : (pi unrestricted tail : (family SM86Program) . (family SM86Program)) . (lambda unrestricted tail : (family SM86Program) . (block (llBackUnless predicate (succ (llLength block)) tail))))) -- the first inner loop: W_e, x_j loaded; y = fma(W_e, x_j, y); e, j advanced def llFirstLoop = (lambda unrestricted k : Nat . (llLoop llP1 (lambda unrestricted tail : (family SM86Program) . (lsNext (llAddressOf llWAddress llE 352) (lsNext (lsLoad llW llWAddress 0) (lsNext (llAddressOf llXAddress llJ 360) (lsNext (lsLoad llX llXAddress 0) (lsNext (lsFusedWith lsControlWaitSB1 llY llW llX llY) (lsNext (llAddImmediate llE 1) (lsNext (llAddImmediate llJ 1) (lsNext (llGreater llP1 llJ (naturalSaturatingSubtract k 1)) tail))))))))))) -- the second inner loop: x_j, W_e loaded; dW = d x x_j; W' = fma(dW, -eta, W_e); -- both stored; e, j advanced def llSecondLoop = (lambda unrestricted m : Nat . (lambda unrestricted k : Nat . (llLoop llP1 (lambda unrestricted tail : (family SM86Program) . (lsNext (llAddressOf llXAddress llJ 360) (lsNext (lsLoad llX llXAddress 0) (lsNext (llAddressOf llWAddress llE 352) (lsNext (lsLoad llW llWAddress 0) (lsNext (constructor SM86InstructionBody SM86FloatMultiply (lsR llDW) (lsR llD) (lsR llX) lsControlWaitSB1) (lsNext (lsFused llWp llDW llR11 llW) (lsNext (llAddressOf llOutAddress llE 376) (lsNext (lsStore llOutAddress llDW (lsGradientOffset m k 0 0)) (lsNext (lsStore llOutAddress llWp (lsUpdatedOffset m k 0 0)) (lsNext (llAddImmediate llE 1) (lsNext (llAddImmediate llJ 1) (lsNext (llGreater llP1 llJ (naturalSaturatingSubtract k 1)) tail)))))))))))))))) -- the outer loop: y = +0, j = 0, the first inner loop, y_i stored, d and the -- loss, e back to the row's first element, j = 0, the second inner loop, i -- advanced def llOuterLoop = (lambda unrestricted m : Nat . (lambda unrestricted k : Nat . (llLoop llP2 (lambda unrestricted tail : (family SM86Program) . (lsNext (lsZero llY) (lsNext (llMove llJ 0) (llFirstLoop k (lsNext (llAddressOf llYAddress llI 376) (lsNext (lsStore llYAddress llY 0) (lsNext (llAddressOf llTAddress llI 368) (lsNext (lsLoad llT llTAddress 0) (lsNext (constructor SM86InstructionBody SM86FloatNegate (lsR llD) (lsR llT) lsControlWaitSB1) (lsNext (lsAdd llD llY llD) (lsNext (lsFused llLoss llD llD llLoss) (lsNext (llAddImmediate llE (naturalSaturatingSubtract 4294967296 k)) (lsNext (llMove llJ 0) (llSecondLoop m k (lsNext (llAddImmediate llI 1) (lsNext (llGreater llP2 llI (naturalSaturatingSubtract m 1)) tail))))))))))))))))))) def linearStepLoopSM86ProgramFor = (lambda unrestricted m : Nat . (lambda unrestricted k : Nat . (lsPrologue k m (lsNext (lsNegate llR11 10) (lsNext (llMove llHalfRegister 1056964608) (lsNext (lsZero llLoss) (lsNext (llMove llFour 4) (lsNext (llMove llI 0) (lsNext (llMove llE 0) (llOuterLoop m k (lsNext (lsMultiply llLoss llHalfRegister llLoss) (lsNext (llAddressOf llYAddress llI 376) (lsNext (lsStore llYAddress llLoss 0) lsExit))))))))))))) -- instructions run: 18 before the outer loop; per output 2 + 9k + 9 + 13k + 3; -- 4 after; with a margin def linearStepLoopSM86FuelFor = (lambda unrestricted m : Nat . (lambda unrestricted k : Nat . (naturalAdd 32 (naturalMultiply m (naturalAdd 16 (naturalMultiply 22 k))))))