module Realization.Nvidia.SM86.TargetGatherSM86 import Accelerator.SM86.Control import Accelerator.SM86.Immediate import Accelerator.SM86.Instruction import Accelerator.SM86.InstructionEncoding import Accelerator.SM86.Operands import Accelerator.SM86.Program import Accelerator.SM86.Types import Std.Natural -- The SM86 realization of Learning.Checked.TargetGather: one thread per row, -- row = block x block width + thread; a thread past the last row exits; the -- others load the row's target t, the logit at row x vocabulary + t, and -- store it to the row's slot. No arithmetic touches the word. -- -- Preconditions (Realization contract): -- the launch covers every row (blocks x block width >= rows); -- the parameter block carries, at the three offsets given, the 64-bit -- addresses of the output (4 bytes per row), the logits (row-major, 4 -- bytes each) and the targets (a 32-bit token per row); -- rows x vocabulary x 4 fits 32 bits (the index is a 32-bit product); -- every target is below the vocabulary. -- -- The same generator makes Coppelius's image (output 0x160, logits 0x168, -- targets 0x170, vocabulary 12288, 1024 rows in 4 blocks of 256) and the -- checked path's (the linear step's layout: Checked.TargetGatherProbe). def targetGatherSM86CoalescedBlockThreads : Nat = 256 def targetGatherSM86OutputArgument : Nat = 0 def targetGatherSM86LogitsArgument : Nat = 1 def targetGatherSM86TargetsArgument : Nat = 2 def tgR = (lambda unrestricted index : Nat . (sm86Register (nat-to-byte index))) def tgU = (lambda unrestricted value : Nat . (sm86Unsigned32 (nat-to-byte (nat-modulo value 256)) (nat-to-byte (nat-modulo (nat-divide value 256) 256)) (nat-to-byte (nat-modulo (nat-divide value 65536) 256)) (nat-to-byte (nat-modulo (nat-divide value 16777216) 256)))) -- stall 15; the two S2R on SB0, each LDG on SB1 def tgControlWith = (lambda unrestricted write : (family SM86Barrier) . (lambda unrestricted wait : Nat . (constructor SM86Control SM86ControlValue (byte 15) (constructor SM86YieldMode SM86Continue) write (constructor SM86Barrier SM86BarrierNone) (nat-to-byte wait) (byte 0)))) def tgControl : (family SM86Control) = (tgControlWith (constructor SM86Barrier SM86BarrierNone) 0) def tgSetSB0 : (family SM86Control) = (tgControlWith (constructor SM86Barrier SM86Barrier0) 0) def tgSetSB1 : (family SM86Control) = (tgControlWith (constructor SM86Barrier SM86Barrier1) 0) def tgWaitSB0 : (family SM86Control) = (tgControlWith (constructor SM86Barrier SM86BarrierNone) 1) def tgWaitSB1 : (family SM86Control) = (tgControlWith (constructor SM86Barrier SM86BarrierNone) 2) def tgNext = (lambda unrestricted body : (family SM86InstructionBody) . (lambda unrestricted tail : (family SM86Program) . (constructor SM86Program SM86ProgramNext (sm86Instruction body) tail))) -- the pair destination := left x R2 (= 4) + the 64-bit parameter at offset def tgAddress = (lambda unrestricted destination : Nat . (lambda unrestricted left : Nat . (lambda unrestricted offset : Nat . (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (tgR destination) (tgR left) (tgR 2) (byte 0) (tgU offset) tgControl)))) -- Registers: R0 row; R1 block; R2 4; R4:R5 ⌖ R6 target; R8:R9 -- &output; R10 row x vocabulary + target; R12:R13 &logit; R14 the logit. def targetGatherSM86Program = (lambda unrestricted outputOffset : Nat . (lambda unrestricted logitsOffset : Nat . (lambda unrestricted targetsOffset : Nat . (lambda unrestricted vocabulary : Nat . (lambda unrestricted rows : Nat . (lambda unrestricted blockWidth : Nat . (tgNext (constructor SM86InstructionBody SM86SpecialToRegister (tgR 0) (constructor SM86SpecialRegister SM86ThreadIdX) tgSetSB0) (tgNext (constructor SM86InstructionBody SM86SpecialToRegister (tgR 1) (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) tgSetSB0) (tgNext (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate (tgR 0) (tgR 1) (tgU blockWidth) (tgR 0) tgWaitSB0) (tgNext (constructor SM86InstructionBody SM86PredicateGreaterThanImmediate (constructor SM86Predicate SM86Predicate0) (tgR 0) (tgU (naturalSaturatingSubtract rows 1)) tgControl) (constructor SM86Program SM86ProgramNext (sm86PredicatedInstruction (constructor SM86Predicate SM86Predicate0) (constructor SM86InstructionBody SM86Exit tgControl)) (tgNext (constructor SM86InstructionBody SM86MoveImmediate (tgR 2) (tgU 4) tgControl) (tgNext (tgAddress 4 0 targetsOffset) (tgNext (constructor SM86InstructionBody SM86LoadGlobal (tgR 6) (tgR 4) (tgU 0) tgSetSB1) (tgNext (tgAddress 8 0 outputOffset) (tgNext (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate (tgR 10) (tgR 0) (tgU vocabulary) (tgR 6) tgWaitSB1) (tgNext (tgAddress 12 10 logitsOffset) (tgNext (constructor SM86InstructionBody SM86LoadGlobal (tgR 14) (tgR 12) (tgU 0) tgSetSB1) (tgNext (constructor SM86InstructionBody SM86StoreGlobal (tgR 8) (tgR 14) (tgU 0) tgWaitSB1) (tgNext (constructor SM86InstructionBody SM86Exit tgControl) (constructor SM86Program SM86ProgramEnd))))))))))))))))))))) -- Coppelius's instance: `correct` at 0x160, the logits at 0x168, the -- targets at 0x170; 12288 tokens; 1024 rows in blocks of 256 def targetGatherSM86Coppelius : (family SM86Program) = (targetGatherSM86Program 0x160 0x168 0x170 12288 1024 256) def targetGatherSM86Encode = (lambda unrestricted program : (family SM86Program) . (eliminate SM86ProgramEncodingResult (lambda unrestricted current : (family SM86ProgramEncodingResult) . Bytes) (sm86EncodeProgram program) (branch SM86ProgramEncodingSucceeded bytes telemetry . bytes) (branch SM86ProgramEncodingFailed index failure telemetry . b"")))