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 247–256

saMaskTile

Full file
the diagonal tile's mask: element (row r + 8h, column 8 nt + 2 (lane % 4) + b) is masked when its column exceeds its row: when 8 nt + b - 8 h exceeds 16 warp + lane / 4 - 2 (lane % 4). saMask holds that plus 64, plus 4096 for every tile after this one, so no element of an earlier tile is masked
247def saMaskTile = (lambda unrestricted tail : (family SM86Program) .
248  (saFor 32 (lambda unrestricted i : Nat .
249    (let unrestricted nt = (naturalDivideUnchecked i 4) in (let unrestricted h = (naturalDivideUnchecked (naturalModuloUnchecked i 4) 2) in
250      (let unrestricted b = (naturalModuloUnchecked i 2) in
251      (let unrestricted threshold = (naturalSaturatingSubtract (naturalAdd (naturalAdd (naturalMultiply 8 nt) b) 63) (naturalMultiply 8 h)) in
252      (lambda unrestricted rest : (family SM86Program) .
253        (saGreater saP0 saMask threshold
254          (saUnless saP0 (constructor SM86InstructionBody SM86MoveImmediate (saR (saElem (saS nt) h b)) (saU saMinusInfinity) saPlain)
255            rest))))))))
256    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.