Source/Packages

Realization.Nvidia.SM86.StreamingAttentionSM86

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

1,129 lines205 declarations66.1 KiBSHA-256 aefcbc3b00ca

def · lines 1081–1104

streamingAttentionReduceGroupedGradientUncheckedSM86

Full file
Sum the query-head dK or dV planes of one K/V group into its compact gradient plane. Both backward outputs have 64*T binary32 elements per head, though dV is transposed within that plane. A CTA owns 128 elements of one K/V head; no atomics or cross-CTA ordering are needed. The caller supplies disjoint source/destination regions and launches grid (64*T/128, keyValueHeads, 1) for T divisible by 64.
1081def streamingAttentionReduceGroupedGradientUncheckedSM86 =
1082  (lambda unrestricted seq : Nat .
1083    (lambda unrestricted heads : Nat .
1084      (lambda unrestricted keyValueHeads : Nat .
1085        (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in
1086        (let unrestricted planeWords = (naturalMultiply saHeadWidth seq) in
1087        (let unrestricted planeBytes = (naturalMultiply planeWords 4) in
1088        (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX)
1089        (saS2R 2 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
1090        (saS2R 3 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
1091        (saMovImmAfter 7 0 saWait5
1092        (saImad 4 2 saThreads saTid
1093        (saImad 5 3 planeWords 4
1094        (saImad 6 3 (naturalMultiply groups planeWords) 4
1095        (saMovImm 11 4
1096        (saWide 12 6 11 (saArgument 0)
1097        (saWide 14 5 11 (saArgument 1)
1098        (saMovImm 9 0
1099        (saFor groups (lambda unrestricted group : Nat .
1100          (lambda unrestricted rest : (family SM86Program) .
1101            (saLoad 8 12 (naturalMultiply group planeBytes) saSB0 saWaitNone
1102              (saFadd 9 9 8 saWait0 rest))))
1103        (saOp (saStore 14 9 0 saWaitNone)
1104        (saOp (constructor SM86InstructionBody SM86Exit (saAfter saWaitAll)) sm86ProgramEmpty))))))))))))))))))))

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.