Each query head writes a separate dK/dV plane. A later reduction sums
those planes into the compact K/V gradient; concurrent CTAs never race on
one compact gradient address.
973def streamingAttentionKeyGroupedUncheckedSM86 = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted keyValueHeads : Nat .
974 (let unrestricted groups = (naturalDivideUnchecked heads keyValueHeads) in
975 (let unrestricted headBytes = (naturalMultiply seq (naturalMultiply saHeadWidth saHalfBytes)) in
976 (let unrestricted subTiles = (naturalDivideUnchecked seq skSubTile) in
977 (let unrestricted body = (skTileBody seq sm86ProgramEmpty) in
978 (saS2R saTid (constructor SM86SpecialRegister SM86ThreadIdX)
979 (saS2R saTileIndex (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdX)
980 (saS2R skKeyValueHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdY)
981 (saS2R skHead (constructor SM86SpecialRegister SM86CooperativeThreadArrayIdZ)
982 (saMovImm saOne 1
983 (saMovImmAfter saScratch 31 saWait5
984 (saImad skHead skKeyValueHead groups skHead
985 (saAnd saLane saTid saScratch
986 (saShr saWarp saTid 5
987 (saShr saQuad saLane 2
988 (saMovImm saScratch 3
989 (saAnd saQuadLane saLane saScratch
990 -- the thread's key row: 32 tile + 16 warp + lane / 4
991 (saImad skRow saWarp 16 saQuad
992 (saImad skRow saTileIndex skSubTile skRow
993 -- K's and V's A fragments
994 (saMovImm saScratch 0
995 (saImad saScratch skRow (naturalMultiply saHeadWidth saHalfBytes) saScratch
996 (saImad saScratch saQuadLane 4 saScratch
997 (saImad saScratch skKeyValueHead headBytes saScratch
998 (saWide saPointer saScratch saOne (saArgument 1)
999 (saFor 16 (lambda unrestricted i : Nat .
1000 (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
1001 (saLoad (naturalAdd (skK kt) j) saPointer
1002 (naturalAdd (naturalMultiply kt 32)
1003 (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
1004 (naturalMultiply (naturalDivideUnchecked j 2) 16)))
1005 saSB1 saWaitNone))))
1006 (saWide saPointer saScratch saOne (saArgument 2)
1007 (saFor 16 (lambda unrestricted i : Nat .
1008 (let unrestricted kt = (naturalDivideUnchecked i 4) in (let unrestricted j = (naturalModuloUnchecked i 4) in
1009 (saLoad (naturalAdd (skV kt) j) saPointer
1010 (naturalAdd (naturalMultiply kt 32)
1011 (naturalAdd (naturalMultiply (naturalModuloUnchecked j 2) (naturalMultiply 8 (naturalMultiply saHeadWidth saHalfBytes)))
1012 (naturalMultiply (naturalDivideUnchecked j 2) 16)))
1013 saSB1 saWaitNone))))
1014 -- dK's offset (binary32 rows) in saScratch
1015 (saMovImm saScratch 0
1016 (saImad saScratch skRow (naturalMultiply saHeadWidth 4) saScratch
1017 (saImad saScratch saQuadLane 8 saScratch
1018 (saImad saScratch skHead (naturalMultiply seq (naturalMultiply saHeadWidth 4)) saScratch
1019 -- the queries start at the last sub-tile: Q's and dO's rows (row lane / 4,
1020 -- the quad lane's pair), Q^T's and dO^T's columns (row lane / 4), L's and D's
1021 (saMovImm saKeyOffset 0
1022 (saImad saKeyOffset saQuad (naturalMultiply saHeadWidth saHalfBytes) saKeyOffset
1023 (saImad saKeyOffset saQuadLane 4 saKeyOffset
1024 (saImad saKeyOffset skHead headBytes saKeyOffset
1025 (saAddImm saKeyOffset saKeyOffset (naturalMultiply (naturalSaturatingSubtract subTiles 1) (naturalMultiply skSubTile (naturalMultiply saHeadWidth saHalfBytes)))
1026 (saMovImm saValueOffset 0
1027 (saImad saValueOffset saQuad (naturalMultiply seq saHalfBytes) saValueOffset
1028 (saImad saValueOffset saQuadLane 4 saValueOffset
1029 (saImad saValueOffset skHead headBytes saValueOffset
1030 (saAddImm saValueOffset saValueOffset (naturalMultiply (naturalSaturatingSubtract subTiles 1) (naturalMultiply skSubTile saHalfBytes))
1031 (saMovImm skLogOffset 0
1032 (saImad skLogOffset saQuadLane 8 skLogOffset
1033 (saImad skLogOffset skHead (naturalMultiply seq 4) skLogOffset
1034 (saAddImm skLogOffset skLogOffset (naturalMultiply (naturalSaturatingSubtract subTiles 1) (naturalMultiply skSubTile 4))
1035 -- the sub-tiles from the last down to the diagonal; the mask's base,
1036 -- 2 (lane % 4) - 16 warp - lane / 4 + 128 - 4096
1037 (saMovImm saCount subTiles
1038 (saImad saCount saTileIndex 4294967295 saCount
1039 (saMovImm saMaskBase 0
1040 (saImad saMaskBase saQuadLane 2 saMaskBase
1041 (saImad saMaskBase saWarp 4294967280 saMaskBase
1042 (saImad saMaskBase saQuad 4294967295 saMaskBase
1043 (saAddImm saMaskBase saMaskBase (naturalSaturatingSubtract 4294967296 (naturalSaturatingSubtract 4096 128))
1044 (saMovConst saScale (saArgument 10)
1045 (saFor 64 (lambda unrestricted i : Nat . (saMovImm (naturalAdd 84 i) 0))
1046 (sm86ProgramAppend body
1047 (sm86ProgramAppend (skLoopTail (sm86ProgramCount body))
1048 (skEpilogue seq
1049 (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.