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 361–388

saEpilogue

Full file
---- the epilogue: O / l to the merged plane, m + log2 l ----
361def saEpilogue = (lambda unrestricted seq : Nat . (lambda unrestricted heads : Nat . (lambda unrestricted tail : (family SM86Program) .
362  (let unrestricted rowBytes = (naturalMultiply heads (naturalMultiply saHeadWidth saHalfBytes)) in
363  (saFor 2 (lambda unrestricted h : Nat . (lambda unrestricted rest : (family SM86Program) .
364    (saShfl (saTmp 2) (saL h) 1 saWait3
365    (saFadd (saL h) (saL h) (saTmp 2) saWait5
366    (saShfl (saTmp 2) (saL h) 2 saWaitNone
367    (saFadd (saL h) (saL h) (saTmp 2) saWait5
368    (saMufu (saTmp (naturalAdd 4 h)) (saL h) (constructor SM86MultiFunction SM86Reciprocal) saWaitNone
369    (saMufu (saTmp (naturalAdd 6 h)) (saL h) (constructor SM86MultiFunction SM86LogarithmBase2) saWaitNone
370    (saFor 16 (lambda unrestricted i : Nat .
371      (let unrestricted e = (saElem (saO (naturalDivideUnchecked i 2)) h (naturalModuloUnchecked i 2)) in
372        (saFmul e e (saTmp (naturalAdd 4 h)) (naturalSelect (naturalIsZero i) saWait4 saWaitNone))))
373    (saFadd (saM h) (saM h) (saTmp (naturalAdd 6 h)) saWaitNone
374      rest))))))))))
375  -- the output: this thread's rows r and r + 8, columns 8 nt + 2 (lane % 4)
376  (saWide saPointer saScratch saOne (saArgument 3)
377  (saFor 16 (lambda unrestricted i : Nat .
378    (let unrestricted nt = (naturalDivideUnchecked i 2) in (let unrestricted h = (naturalModuloUnchecked i 2) in
379      (lambda unrestricted rest : (family SM86Program) .
380        (saPack (saTmp 2) (succ (saElem (saO nt) h 0)) (saElem (saO nt) h 0)
381        (saOp (saStore saPointer (saTmp 2) (naturalAdd (naturalMultiply nt 16) (naturalMultiply h (naturalMultiply 8 rowBytes))) saWaitNone)
382          rest))))))
383  -- the log-sum-exp, by the quad's first lane
384  (saWide saPointer saRow saOne (saArgument 4)
385  (saGreater saP0 saQuadLane 0
386  (saUnless saP0 (saStore saPointer (saM 0) 0 saWaitNone)
387  (saUnless saP0 (saStore saPointer (saM 1) 32 saWaitNone)
388    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.