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