step method

  1. @override
void step()
override

Apply one update step using each parameter's current .grad. Parameters with grad == null are skipped.

Implementation

@override
void step() {
  _step += 1;
  final biasCorr1 = 1.0 - _pow(beta1, _step);
  final biasCorr2 = 1.0 - _pow(beta2, _step);

  for (int i = 0; i < parameters.length; i++) {
    final p = parameters[i];
    final g = p.grad;
    if (g == null) continue;

    // Lazily allocate moment buffers to match parameter shape/device.
    var m = _m[i];
    var v = _v[i];
    if (m == null) {
      m = Tensor.fill(p.shape, 0.0, device: p.device);
      _m[i] = m;
    }
    if (v == null) {
      v = Tensor.fill(p.shape, 0.0, device: p.device);
      _v[i] = v;
    }

    // Update biased moments — dispose every intermediate to avoid
    // leaking one GPU handle per sub-op per step (previously this
    // leaked ~4 handles per param per step → 6 GB VRAM in ~10 steps
    // for a mid-size GPT, silently corrupting later allocations).
    final mBeta = m * beta1;
    final gOneMinusBeta1 = g * (1.0 - beta1);
    final mSum = mBeta + gOneMinusBeta1;
    mBeta.dispose();
    gOneMinusBeta1.dispose();
    m.assign(mSum);

    final vBeta = v * beta2;
    final gSq = g * g;
    final gSqScaled = gSq * (1.0 - beta2);
    gSq.dispose();
    final vSum = vBeta + gSqScaled;
    vBeta.dispose();
    gSqScaled.dispose();
    v.assign(vSum);

    final mHat = m / biasCorr1;
    final vHat = v / biasCorr2;
    final sqrtV = vHat.pow(0.5);
    vHat.dispose();
    final denom = sqrtV + eps;
    sqrtV.dispose();
    final ratio = mHat / denom;
    mHat.dispose();
    denom.dispose();
    var update = ratio * lr;
    ratio.dispose();

    if (weightDecay != 0.0) {
      final pDet = p.detach();
      final wd = pDet * (lr * weightDecay);
      pDet.dispose();
      final newUpdate = update + wd;
      update.dispose();
      wd.dispose();
      update = newUpdate;
    }

    final pDet = p.detach();
    final newP = pDet - update;
    pDet.dispose();
    update.dispose();
    p.assign(newP);
  }
}