Source/Packages

Realization.Nvidia.SM86.StreamingAttentionSM86

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

1,203 lines221 declarations70.9 KiBSHA-256 23d4e5e2aa2a

def · lines 1155–1178

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.
1155def streamingAttentionReduceGroupedGradientUncheckedSM86 =
1156  (lambda unrestricted seq : Nat .
1157    (lambda unrestricted heads : Nat .
1158      (lambda unrestricted keyValueHeads : Nat .
1159        (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in
1160        (let unrestricted planeWords = (naturalMultiply saHeadWidth seq) in
1161        (let unrestricted planeBytes = (naturalMultiply planeWords 4) in
1162        (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX)
1163        (saS2R 2 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
1164        (saS2R 3 (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
1165        (saMovImmAfter 7 0 saWait5
1166        (saImad 4 2 saThreads saTid
1167        (saImad 5 3 planeWords 4
1168        (saImad 6 3 (naturalMultiply groups planeWords) 4
1169        (saMovImm 11 4
1170        (saWide 12 6 11 (saArgument 0)
1171        (saWide 14 5 11 (saArgument 1)
1172        (saMovImm 9 0
1173        (saFor groups (lambda unrestricted group : Nat .
1174          (lambda unrestricted rest : (family SM86Program) .
1175            (saLoad 8 12 (naturalMultiply group planeBytes) saSB0 saWaitNone
1176              (saFadd 9 9 8 saWait0 rest))))
1177        (saOp (saStore 14 9 0 saWaitNone)
1178        (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.