Source/Packages

Realization.Nvidia.SM86.AdamWHalfSM86

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

205 lines38 declarations10.6 KiBSHA-256 557f99a5eb69

def · lines 48–48

ahThreads

Full file
AdamW over the whole flat bank with the parameters' half copies written as it goes: the half cast that followed the update, fused into it. Each element's update is Realization.Nvidia.SM86.AdamWSM86's (adamWF32MathProgram), instruction for instruction, and each half pair is the cast's (F2FP, element 2i + 1 high, in the realization's half format), so the parameters, the moments and the halves are bit-identical to the update followed by the cast: g = g one; m = m beta1 + g (1 - beta1); v = v beta2 + g^2 (1 - beta2) p = p decay + -(m stepSize (1 / (sqrt v + epsilon))) half[i] = half(p[i]) The half bank is parallel to the parameter bank: element i's half at byte 2 i of it. A thread takes the quad 4T .. 4T + 3 (T = block 256 + thread): each bank's four elements in one 128-bit load and store, the four halves in one 64-bit store -- the bank is memory-bound, and wide accesses keep more of it in flight a thread -- then T + stride, ... `iterations` times; quads at or past `quads` are untouched. The gradient is read, never written (every gradient the step reads was written whole by the backward before it). Parameters (constant bank 0 from the SM86 parameter base): the pointers P, G, M, V (arguments 0-3, 16-byte aligned); the scalars one, beta1, 1 - beta1, beta2, 1 - beta2, decay, stepSize, epsilon (words from argument 4, where AdamWSM86's ABI puts them, so the host writes the step's scalars at the same offsets); the half bank's pointer (argument 9, 8-byte aligned). Registers: R0 thread, R1 block, R2 T, R3 sixteen, R4 eight; R6:R7, R8:R9, R10:R11, R12:R13 the quad's P, G, M, V; R18..R25 the scalars; R29 the iterations left; R30:R31 the halves' address; R32..R35, R36..R39, R40..R43, R44..R47 the quad's p, g, m, v; element e's temporaries R48 + 3e .. R50 + 3e; R60, R61 the half pairs. Loads under SB0..SB3 (P, G, M, V), the square roots and reciprocals under SB4; the stores' reads under SB5, which the next iteration's first load waits for.
48def ahThreads : 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.