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 677–751

streamingAttentionQueryGroupedUncheckedSM86

Full file
677def streamingAttentionQueryGroupedUncheckedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat .
678  (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in
679  (let unrestricted headBytes = (naturalMultiply seq (naturalMultiply saHeadWidth saHalfBytes)) in
680  (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in
681  (let unrestricted body = (sqTileBody seq sm86ProgramEmpty) in
682  (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX)
683  (saS2R saTileIndex (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
684  (saS2R saKeyValueHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
685  (saS2R saHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdZ)
686  (saMovImm saOne 1
687  (saMovImmAfter saScratch 31 saWait5
688  (saImad saHead saKeyValueHead groups saHead
689  (saAnd saLane saTid saScratch
690  (saShr saWarp saTid 5
691  (saShr saQuad saLane 2
692  (saMovImm saScratch 3
693  (saAnd saQuadLane saLane saScratch
694  (saImad saRow saWarp 16 saQuad
695  (saImad saRow saTileIndex saTile saRow
696  -- Q's and dO's A fragments: head h's plane, the row, the quad lane's pair
697  (saMovImm saScratch 0
698  (saImad saScratch saRow (naturalMultiply saHeadWidth saHalfBytes) saScratch
699  (saImad saScratch saQuadLane 4 saScratch
700  (saImad saScratch saHead headBytes saScratch
701  (saWide saPointer saScratch saOne (saArgument 0)
702  (saFor 16 (lambda unrestricted i : Nat .
703    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
704      (saLoad (saQReg kt j) saPointer
705        (naturalAdd (naturalMultiply kt 32)
706          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
707            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
708        saSB1 saWaitNone))))
709  (saWide saPointer saScratch saOne (saArgument 4)
710  (saFor 16 (lambda unrestricted i : Nat .
711    (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
712      (saLoad (naturalAdd (sqDOut kt) j) saPointer
713        (naturalAdd (naturalMultiply kt 32)
714          (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
715            (naturalMultiply (naturalDivideUnchecked j 2) 16)))
716        saSB1 saWaitNone))))
717  -- dQ's offset (binary32 rows), kept in saKeyPointer's partner until the
718  -- epilogue: saScratch becomes O's (merged plane) for the row terms
719  (saMovImm saKeyOffset 0
720  (saImad saKeyOffset saRow (naturalMultiply saHeadWidth 4) saKeyOffset
721  (saImad saKeyOffset saQuadLane 8 saKeyOffset
722  (saImad saKeyOffset saHead (naturalMultiply seq (naturalMultiply saHeadWidth 4)) saKeyOffset
723  (saMovImm saScratch 0
724  (saImad saScratch saRow rowBytes saScratch
725  (saImad saScratch saHead (naturalMultiply saHeadWidth saHalfBytes) saScratch
726  (saImad saScratch saQuadLane 4 saScratch
727  -- the row terms' offset: 4 (h seq + row)
728  (saImad saRow saHead seq saRow
729  (saMovImm saValueOffset 0
730  (saImad saRow saRow 4 saValueOffset
731  (sqRowTerms seq heads
732  -- dQ's offset into saScratch; K's and V's from key 0, K^T's
733  (saAddImm saScratch saKeyOffset 0
734  (saMovImm saKeyOffset 0
735  (saImad saKeyOffset saQuad (naturalMultiply saHeadWidth saHalfBytes) saKeyOffset
736  (saImad saKeyOffset saQuadLane 4 saKeyOffset
737  (saImad saKeyOffset saKeyValueHead headBytes saKeyOffset
738  (saMovImm sqKTOffset 0
739  (saImad sqKTOffset saQuad (naturalMultiply seq saHalfBytes) sqKTOffset
740  (saImad sqKTOffset saQuadLane 4 sqKTOffset
741  (saImad sqKTOffset saKeyValueHead headBytes sqKTOffset
742  (saAddImm saCount saTileIndex 1
743  (saImad saMaskBase saWarp 16 saQuad
744  (saImad saMaskBase saQuadLane 4294967294 saMaskBase
745  (saAddImm saMaskBase saMaskBase (naturalSaturatingSubtract 4294967296 (naturalSaturatingSubtract 4096 64))
746  (saMovConst saScale (saArgument 9)
747  (saFor 32 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 68 i) 0))
748  (sm86ProgramAppend body
749    (sm86ProgramAppend (sqLoopTail (sm86ProgramCount body))
750      (sqEpilogue
751        (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.