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 = 256The compiler supplied declaration spans and resolved links from this source snapshot. This page does not assert that this file belongs to a checked closure.