Documentation

Hesper.Training.FFNBackward

FFN Backward GPU Kernels #

Backward pass for the FFN (Feed-Forward Network) sub-layer:

Forward: gate = W_gate @ normed2 up = W_up @ normed2 hidden = ReLU²(gate) × up ffnNormed = RMSNorm(hidden) output = residual + W_down @ ffnNormed

Backward (reverse):

  1. dFFNNormed = W_down^T @ dOutput
  2. dHidden = RMSNorm_bwd(hidden, gamma, dFFNNormed)
  3. dGate, dUp = ReLU²Mul_bwd(gate, up, dHidden)
  4. dNormed2 = W_gate^T @ dGate + W_up^T @ dUp
  5. dResidual += RMSNorm_bwd(residual, gamma, dNormed2)

ReLU²×Mul Backward #

Forward: h = max(0, gate)² × up Backward: dGate = dH × up × 2 × ReLU(gate) dUp = dH × max(0, gate)²

ReLU²×Mul backward kernel. Forward: hidden[i] = max(0, gate[i])² × up[i] Backward: dGate[i] = dHidden[i] × up[i] × 2 × max(0, gate[i]) dUp[i] = dHidden[i] × max(0, gate[i])²

Reads: gate, up, dHidden Writes: dGate, dUp

Equations
  • One or more equations did not get rendered due to their size.
Instances For
    def Hesper.Training.FFNBackward.executeReluSqrMulBackward (device : WebGPU.Device) (gateBuf upBuf dHiddenBuf dGateBuf dUpBuf : WebGPU.Buffer) (numElements : Nat) :
    Equations
    • One or more equations did not get rendered due to their size.
    Instances For
      def Hesper.Training.FFNBackward.executeFFNBackward (device : WebGPU.Device) (wDown wGate wUp : Layers.BitLinear.BitLinear WebGPU.Buffer WGSL.Execute.PreparedDispatch WGSL.Execute.CompiledKernel) (ffnSubNormScale ffnNormScale dOutputBuf savedHidden savedResidual1 savedGate savedUp dFFNNormed dFFNHidden dGate dUp dNormed2 dHiddenBuf : WebGPU.Buffer) (dim ffnDim : Nat) :

      Execute full FFN backward for one layer. Requires saved forward activations: gate, up, hidden, residual1.

      @param device GPU device @param block Transformer block (for weight access) @param dOutputBuf Gradient from next layer [dim] @param dHiddenBuf Scratch buffer [dim] — will contain dResidual contribution @param savedGate Saved gate buffer [ffnDim] from forward @param savedUp Saved up buffer [ffnDim] from forward @param savedHidden Saved hidden buffer [ffnDim] from forward (pre sub-norm) @param savedResidual1 Saved residual1 buffer [dim] from forward (pre ffn-norm) @param dFFNNormed Scratch [ffnDim] @param dFFNHidden Scratch [ffnDim] @param dGate Scratch [ffnDim] @param dUp Scratch [ffnDim] @param dNormed2 Scratch [dim]

      Equations
      • One or more equations did not get rendered due to their size.
      Instances For