Source/Packages

Realization.Nvidia.SM86.GatedGELUSM86

packages/realizations/cooperative/nvidia-sm86/src/Realization/Nvidia/SM86/GatedGELUSM86.alpha

170 lines18 declarations7.4 KiBSHA-256 fe6a7ec05912

def · lines 36–36

ggThreads

Full file
A gated feed-forward's activation in one pass each way, the same arithmetic as the separate launches it replaces (Realization.Nvidia.SM86.GELUSM86's forward and backward, the elementwise product, the binary32 -> half cast), instruction for instruction, so the results are bit-identical: forward gelu = GELU(gate) (binary32, kept for the backward) gated = half(up gelu) (the down product's operand) backward upGradient = half(dGated gelu) gateGradient = half(GELU'(gate) (dGated up)) with GELU(x) = x ((tanh(s (x^3 c + x)) + 1) h) and GELU'(x) g = g ((1 + t) h + x (1 - t^2) (x^2 3c + 1) s h), t = tanh(s (x^3 c + x)), the scalars s = sqrt(2/pi), c = 0.044715, h = 1/2, 1 and 3c as GELUSM86 reads them. The half packing is the cast's: F2FP with element 2i + 1 high, which the realization's half format makes binary16 or bfloat16. Two elements a thread, 256 threads a block, the element count a multiple of 512. Parameters (constant bank 0 from the SM86 parameter base): the pointers (forward: gelu, gated, up, gate; backward: upGradient, gateGradient, dGated, gelu, up, gate), then after six pointer slots the scalars s, c, h, 1, 3c.
36def ggThreads : Nat = 256

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.