step method

void step()

Implementation

void step() {
  t++;
  final double bc1 = 1.0 - math.pow(beta1, t);
  final double bc2 = 1.0 - math.pow(beta2, t);

  for (int i = 0; i < parameters.length; i++) {
    final p = parameters[i];

    // Perform global gradient clipping for this parameter tensor
    _clipGradients(p.grad);

    for (int j = 0; j < p.data.length; j++) {
      final g = p.grad[j];

      // 1. Update biased first moment estimate
      m[i][j] = beta1 * m[i][j] + (1.0 - beta1) * g;

      // 2. Update biased second raw moment estimate
      v[i][j] = beta2 * v[i][j] + (1.0 - beta2) * (g * g);

      // 3. Compute bias-corrected estimates
      final mHat = m[i][j] / bc1;
      final vHat = v[i][j] / bc2;

      // 4. Update parameter with Adam rule
      // The lr is modulated by the signal-to-noise ratio (mHat/sqrt(vHat))
      p.data[j] -= lr * mHat / (math.sqrt(vHat) + epsilon);
    }
  }
}