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 373–400

saEpilogue

Full file
---- the epilogue: O / l to the merged plane, m + log2 l ----
373def saEpilogue = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted tail : (family SM86Program) .
374  (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in
375  (saFor 2 (lambda unrestricted h : Nat . (lambda unrestricted rest : (family SM86Program) .
376    (saShfl (saTmp 2) (saL h) 1 saWait3
377    (saFadd (saL h) (saL h) (saTmp 2) saWait5
378    (saShfl (saTmp 2) (saL h) 2 saWaitNone
379    (saFadd (saL h) (saL h) (saTmp 2) saWait5
380    (saMufu (saTmp (naturalAdd 4 h)) (saL h) (constructor SM86MultiFunction SM86Reciprocal) saWaitNone
381    (saMufu (saTmp (naturalAdd 6 h)) (saL h) (constructor SM86MultiFunction SM86LogarithmBase2) saWaitNone
382    (saFor 16 (lambda unrestricted i : Nat .
383      (let unrestricted e = (saElem (saO (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in
384        (saFmul e e (saTmp (naturalAdd 4 h)) (naturalSelect (naturalIsZero i) saWait4 saWaitNone))))
385    (saFadd (saM h) (saM h) (saTmp (naturalAdd 6 h)) saWaitNone
386      rest))))))))))
387  -- the output: this thread's rows r and r + 8, columns 8 nt + 2 (lane % 4)
388  (saWide saPointer saScratch saOne (saArgument 3)
389  (saFor 16 (lambda unrestricted i : Nat .
390    (let unrestricted nt = (naturalDivideUnchecked i 2) in (let unrestricted h = (naturalModuloUnchecked i 2) in
391      (lambda unrestricted rest : (family SM86Program) .
392        (saPack (saForwardOutputPackRegister i) (succ (saElem (saO nt) h 0)) (saElem (saO nt) h 0)
393        (saOp (saStore saPointer (saForwardOutputPackRegister i) (naturalAdd (naturalMultiply nt 16) (naturalMultiply h (naturalMultiply 8 rowBytes))) saWaitNone)
394          rest))))))
395  -- the log-sum-exp, by the quad's first lane
396  (saWide saForwardLogSumPointer saRow saOne (saArgument 4)
397  (saGreater saP0 saQuadLane 0
398  (saUnless saP0 (saStore saForwardLogSumPointer (saM 0) 0 saWaitNone)
399  (saUnless saP0 (saStore saForwardLogSumPointer (saM 1) 32 saWaitNone)
400    tail)))))))))))

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.