module Realization.Nvidia.SM86.RotaryTableSM86 import Accelerator.SM86.Control import Accelerator.SM86.Immediate import Accelerator.SM86.Instruction import Accelerator.SM86.InstructionEncoding import Accelerator.SM86.NumericSemantics import Accelerator.SM86.Program import Accelerator.SM86.Types import Std.Natural -- One column of the rotary tables: for the positions p of a block, the -- cosine and sine of p theta, stored at `cos + p stride` and `sin + p -- stride` -- a launch per frequency theta, a thread per position. -- -- The angle is carried in two words: theta = thetaHigh + thetaLow (the -- binary32 nearest theta, and the binary32 nearest what it misses), so -- p theta = p thetaHigh -- whose rounding error a fused multiply-add recovers -- exactly -- plus p thetaLow. It is reduced by whole turns, k = the nearest -- integer to p thetaHigh / (2 pi) (the magic-number rounding: adding 1.5 x -- 2^23 leaves the integer in the low bits), with 2 pi likewise in two words -- (Cody and Waite): r = ((a - k twoPiHigh) - k twoPiLow) + p thetaLow, which -- lies within a turn of zero. MUFU.COS and MUFU.SIN take their argument in -- turns (r / (2 pi)). -- -- Constant bank: 0x160 the cosine column's address, 0x168 the sine column's; -- scalars from 0x190: thetaHigh, thetaLow, 1 / (2 pi), twoPiHigh, twoPiLow, -- 1.5 x 2^23, the stride between positions in bytes. def rotaryTableSM86CosineArgument : Nat = 0 def rotaryTableSM86SineArgument : Nat = 1 def rotaryTableSM86FirstScalarPairArgument : Nat = 6 def rotaryTableSM86ScalarPairCount : Nat = 4 def rotaryTableSM86BlockThreads : Nat = 256 def rotaryTableRegister = (lambda unrestricted index : Byte . (sm86Register index)) def rotaryTableWord = (lambda unrestricted a : Byte . (lambda unrestricted b : Byte . (constructor SM86Unsigned32 SM86Unsigned32Value a b (byte 0) (byte 0)))) def rotaryTableCosine = (rotaryTableWord (byte 96) (byte 1)) def rotaryTableSine = (rotaryTableWord (byte 104) (byte 1)) def rotaryTableScalar = (lambda unrestricted k : Nat . (rotaryTableWord (nat-to-byte (naturalAdd 144 (naturalMultiply 4 k))) (byte 1))) def rotaryTableStackPointer = (rotaryTableWord (byte 40) (byte 0)) def rotaryTableZero = (rotaryTableWord (byte 0) (byte 0)) def rotaryTableSet0 = (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier0)) def rotaryTableWait0 = (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier0)) def rotaryTableSet1 = (sm86SetBarrierControl (constructor SM86Barrier SM86Barrier1)) def rotaryTableWait1 = (sm86WaitBarrierControl (constructor SM86WaitBarrier SM86WaitBarrier1)) def rotaryTableNext = (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) . (constructor SM86Program SM86ProgramNext (sm86Instruction body) tail))) def rotaryTableFloat = (lambda unrestricted d : Byte . (lambda unrestricted k : Nat . (constructor SM86InstructionBody SM86MoveConstant (rotaryTableRegister d) (byte 0) (rotaryTableScalar k) sm86SafeControl))) def rotaryTableFMA = (lambda unrestricted d : Byte . (lambda unrestricted a : Byte . (lambda unrestricted b : Byte . (lambda unrestricted c : Byte . (constructor SM86InstructionBody SM86FloatFusedMultiplyAdd (rotaryTableRegister d) (rotaryTableRegister a) (rotaryTableRegister b) (rotaryTableRegister c) sm86SafeControl))))) def rotaryTableNegate = (lambda unrestricted d : Byte . (lambda unrestricted a : Byte . (constructor SM86InstructionBody SM86FloatNegate (rotaryTableRegister d) (rotaryTableRegister a) sm86SafeControl))) def rotaryTableSM86Program : (family SM86Program) = -- R1 the stack pointer (the ABI's first word); R0 the position (rotaryTableNext (constructor SM86InstructionBody SM86MoveConstant (rotaryTableRegister (byte 1)) (byte 0) rotaryTableStackPointer sm86SafeControl) (rotaryTableNext (constructor SM86InstructionBody SM86SpecialToRegister (rotaryTableRegister (byte 0)) (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) rotaryTableSet0) (rotaryTableNext (constructor SM86InstructionBody SM86SpecialToRegister (rotaryTableRegister (byte 2)) (constructor SM86SpecialRegister SM86ThreadIdX) rotaryTableSet0) (rotaryTableNext (constructor SM86InstructionBody SM86IntegerMultiplyAddConstant (rotaryTableRegister (byte 0)) (rotaryTableRegister (byte 0)) (byte 0) rotaryTableZero (rotaryTableRegister (byte 2)) rotaryTableWait0) -- the scalars (rotaryTableNext (rotaryTableFloat (byte 5) 0) (rotaryTableNext (rotaryTableFloat (byte 6) 1) (rotaryTableNext (rotaryTableFloat (byte 7) 2) (rotaryTableNext (rotaryTableFloat (byte 8) 3) (rotaryTableNext (rotaryTableFloat (byte 9) 4) (rotaryTableNext (rotaryTableFloat (byte 10) 5) (rotaryTableNext (rotaryTableFloat (byte 3) 6) -- a = p thetaHigh, and the rounding error of that product, exactly (rotaryTableNext (constructor SM86InstructionBody SM86IntegerToFloat (rotaryTableRegister (byte 4)) (rotaryTableRegister (byte 0)) sm86SafeControl) (rotaryTableNext (constructor SM86InstructionBody SM86FloatMultiply (rotaryTableRegister (byte 11)) (rotaryTableRegister (byte 4)) (rotaryTableRegister (byte 5)) sm86SafeControl) (rotaryTableNext (rotaryTableNegate (byte 12) (byte 11)) (rotaryTableNext (rotaryTableFMA (byte 13) (byte 4) (byte 5) (byte 12)) -- the angle's low part: the product's error plus p thetaLow (rotaryTableNext (rotaryTableFMA (byte 14) (byte 4) (byte 6) (byte 13)) -- k, the nearest whole number of turns in a (rotaryTableNext (rotaryTableFMA (byte 15) (byte 11) (byte 7) (byte 10)) (rotaryTableNext (rotaryTableNegate (byte 17) (byte 10)) (rotaryTableNext (constructor SM86InstructionBody SM86FloatAdd (rotaryTableRegister (byte 16)) (rotaryTableRegister (byte 15)) (rotaryTableRegister (byte 17)) sm86SafeControl) (rotaryTableNext (rotaryTableNegate (byte 18) (byte 16)) -- r = ((a - k twoPiHigh) - k twoPiLow) + the low part (rotaryTableNext (rotaryTableFMA (byte 19) (byte 18) (byte 8) (byte 11)) (rotaryTableNext (rotaryTableFMA (byte 19) (byte 18) (byte 9) (byte 19)) (rotaryTableNext (constructor SM86InstructionBody SM86FloatAdd (rotaryTableRegister (byte 19)) (rotaryTableRegister (byte 19)) (rotaryTableRegister (byte 14)) sm86SafeControl) -- in turns, then the cosine and the sine (rotaryTableNext (constructor SM86InstructionBody SM86FloatMultiply (rotaryTableRegister (byte 20)) (rotaryTableRegister (byte 19)) (rotaryTableRegister (byte 7)) sm86SafeControl) (rotaryTableNext (constructor SM86InstructionBody SM86MultiFunctionUnitApproximation (rotaryTableRegister (byte 21)) (rotaryTableRegister (byte 20)) (constructor SM86MultiFunction SM86Cosine) rotaryTableSet1) (rotaryTableNext (constructor SM86InstructionBody SM86MultiFunctionUnitApproximation (rotaryTableRegister (byte 22)) (rotaryTableRegister (byte 20)) (constructor SM86MultiFunction SM86Sine) rotaryTableSet1) -- the column entries: base + p stride (rotaryTableNext (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (rotaryTableRegister (byte 24)) (rotaryTableRegister (byte 0)) (rotaryTableRegister (byte 3)) (byte 0) rotaryTableCosine sm86SafeControl) (rotaryTableNext (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (rotaryTableRegister (byte 26)) (rotaryTableRegister (byte 0)) (rotaryTableRegister (byte 3)) (byte 0) rotaryTableSine sm86SafeControl) (rotaryTableNext (constructor SM86InstructionBody SM86StoreGlobal (rotaryTableRegister (byte 24)) (rotaryTableRegister (byte 21)) rotaryTableZero rotaryTableWait1) (rotaryTableNext (constructor SM86InstructionBody SM86StoreGlobal (rotaryTableRegister (byte 26)) (rotaryTableRegister (byte 22)) rotaryTableZero sm86SafeControl) (rotaryTableNext (constructor SM86InstructionBody SM86Exit sm86SafeControl) (constructor SM86Program SM86ProgramEnd)))))))))))))))))))))))))))))))) def rotaryTableSM86RegisterCount : Nat = 32