Source/Packages

Realization.Nvidia.SM86.TargetGatherSM86

packages/realizations/cooperative/nvidia-sm86/src/Realization/Nvidia/SM86/TargetGatherSM86.alpha

102 lines17 declarations5.7 KiBSHA-256 6321b25295ef

Complete file · line 28

TargetGatherSM86.alpha

Definition view
1module Realization.Nvidia.SM86.TargetGatherSM86
2
3import Accelerator.SM86.Control
4import Accelerator.SM86.Immediate
5import Accelerator.SM86.Instruction
6import Accelerator.SM86.InstructionEncoding
7import Accelerator.SM86.Operands
8import Accelerator.SM86.Program
9import Accelerator.SM86.Types
10import Std.Natural
11
12-- The SM86 realization of Learning.Checked.TargetGather: one thread per row,
13-- row = block x block width + thread; a thread past the last row exits; the
14-- others load the row's target t, the logit at row x vocabulary + t, and
15-- store it to the row's slot.  No arithmetic touches the word.
16--
17-- Preconditions (Realization contract):
18--   the launch covers every row (blocks x block width >= rows);
19--   the parameter block carries, at the three offsets given, the 64-bit
20--   addresses of the output (4 bytes per row), the logits (row-major, 4
21--   bytes each) and the targets (a 32-bit token per row);
22--   rows x vocabulary x 4 fits 32 bits (the index is a 32-bit product);
23--   every target is below the vocabulary.
24--
25-- The same generator makes Coppelius's image (output 0x160, logits 0x168,
26-- targets 0x170, vocabulary 12288, 1024 rows in 4 blocks of 256) and the
27-- checked path's (the linear step's layout: Checked.TargetGatherProbe).
28def targetGatherSM86CoalescedBlockThreads : Nat = 256
29def targetGatherSM86OutputArgument : Nat = 0
30def targetGatherSM86LogitsArgument : Nat = 1
31def targetGatherSM86TargetsArgument : Nat = 2
32
33def tgR = (lambda unrestricted index : Nat . (sm86Register (nat-to-byte index)))
34def tgU =
35  (lambda unrestricted value : Nat .
36    (sm86Unsigned32
37      (nat-to-byte (nat-modulo value 256))
38      (nat-to-byte (nat-modulo (nat-divide value 256) 256))
39      (nat-to-byte (nat-modulo (nat-divide value 65536) 256))
40      (nat-to-byte (nat-modulo (nat-divide value 16777216) 256))))
41
42-- stall 15; the two S2R on SB0, each LDG on SB1
43def tgControlWith =
44  (lambda unrestricted write : (family SM86Barrier) . (lambda unrestricted wait : Nat .
45    (constructor SM86Control SM86ControlValue
46      (byte 15)
47      (constructor SM86YieldMode SM86Continue)
48      write
49      (constructor SM86Barrier SM86BarrierNone)
50      (nat-to-byte wait)
51      (byte 0))))
52def tgControl : (family SM86Control) = (tgControlWith (constructor SM86Barrier SM86BarrierNone) 0)
53def tgSetSB0 : (family SM86Control) = (tgControlWith (constructor SM86Barrier SM86Barrier0) 0)
54def tgSetSB1 : (family SM86Control) = (tgControlWith (constructor SM86Barrier SM86Barrier1) 0)
55def tgWaitSB0 : (family SM86Control) = (tgControlWith (constructor SM86Barrier SM86BarrierNone) 1)
56def tgWaitSB1 : (family SM86Control) = (tgControlWith (constructor SM86Barrier SM86BarrierNone) 2)
57
58def tgNext =
59  (lambda unrestricted body : (family SM86InstructionBody) .
60    (lambda unrestricted tail : (family SM86Program) .
61      (constructor SM86Program SM86ProgramNext (sm86Instruction body) tail)))
62
63-- the pair destination := left x R2 (= 4) + the 64-bit parameter at offset
64def tgAddress =
65  (lambda unrestricted destination : Nat . (lambda unrestricted left : Nat . (lambda unrestricted offset : Nat .
66    (constructor SM86InstructionBody SM86IntegerMultiplyAddWideConstant (tgR destination) (tgR left) (tgR 2) (byte 0) (tgU offset) tgControl))))
67
68-- Registers: R0 row; R1 block; R2 4; R4:R5 ⌖ R6 target; R8:R9
69-- &output; R10 row x vocabulary + target; R12:R13 &logit; R14 the logit.
70def targetGatherSM86Program =
71  (lambda unrestricted outputOffset : Nat . (lambda unrestricted logitsOffset : Nat . (lambda unrestricted targetsOffset : Nat .
72  (lambda unrestricted vocabulary : Nat . (lambda unrestricted rows : Nat . (lambda unrestricted blockWidth : Nat .
73    (tgNext (constructor SM86InstructionBody SM86SpecialToRegister (tgR 0) (constructor SM86SpecialRegister SM86ThreadIdX) tgSetSB0)
74    (tgNext (constructor SM86InstructionBody SM86SpecialToRegister (tgR 1) (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX) tgSetSB0)
75    (tgNext (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate (tgR 0) (tgR 1) (tgU blockWidth) (tgR 0) tgWaitSB0)
76    (tgNext (constructor SM86InstructionBody SM86PredicateGreaterThanImmediate (constructor SM86Predicate SM86Predicate0) (tgR 0) (tgU (naturalSaturatingSubtract rows 1)) tgControl)
77    (constructor SM86Program SM86ProgramNext (sm86PredicatedInstruction (constructor SM86Predicate SM86Predicate0) (constructor SM86InstructionBody SM86Exit tgControl))
78    (tgNext (constructor SM86InstructionBody SM86MoveImmediate (tgR 2) (tgU 4) tgControl)
79    (tgNext (tgAddress 4 0 targetsOffset)
80    (tgNext (constructor SM86InstructionBody SM86LoadGlobal (tgR 6) (tgR 4) (tgU 0) tgSetSB1)
81    (tgNext (tgAddress 8 0 outputOffset)
82    (tgNext (constructor SM86InstructionBody SM86IntegerMultiplyAddImmediate (tgR 10) (tgR 0) (tgU vocabulary) (tgR 6) tgWaitSB1)
83    (tgNext (tgAddress 12 10 logitsOffset)
84    (tgNext (constructor SM86InstructionBody SM86LoadGlobal (tgR 14) (tgR 12) (tgU 0) tgSetSB1)
85    (tgNext (constructor SM86InstructionBody SM86StoreGlobal (tgR 8) (tgR 14) (tgU 0) tgWaitSB1)
86    (tgNext (constructor SM86InstructionBody SM86Exit tgControl)
87    (constructor SM86Program SM86ProgramEnd)))))))))))))))))))))
88
89
90
91-- Coppelius's instance: `correct` at 0x160, the logits at 0x168, the
92-- targets at 0x170; 12288 tokens; 1024 rows in blocks of 256
93def targetGatherSM86Coppelius : (family SM86Program) =
94  (targetGatherSM86Program 0x160 0x168 0x170 12288 1024 256)
95
96def targetGatherSM86Encode =
97  (lambda unrestricted program : (family SM86Program) .
98    (eliminate SM86ProgramEncodingResult
99      (lambda unrestricted current : (family SM86ProgramEncodingResult) . Bytes)
100      (sm86EncodeProgram program)
101      (branch SM86ProgramEncodingSucceeded bytes telemetry . bytes)
102      (branch SM86ProgramEncodingFailed index failure telemetry . b"")))

The compiler supplied declaration spans and resolved links from this source snapshot. This page does not assert that this file belongs to a checked closure.