trainStep method

double trainStep({
  1. int batchSize = 32,
})

The optimization step (identical logic to previous version)

Implementation

double trainStep({int batchSize = 32}) {
  if (replayBuffer.length < batchSize) return 0.0;

  optimizer.zeroGrad();
  double totalBatchLoss = 0;
  final random = math.Random();

  for (int i = 0; i < batchSize; i++) {
    final sample = replayBuffer[random.nextInt(replayBuffer.length)];
    final state = model.represent(sample.observations);
    final pred = model.predict(state);

    // Policy Loss
    double pLoss = 0;
    final policyData = pred['policy']!.data;
    double maxLogit = policyData.reduce(math.max);
    double sumExp = 0;
    for (var val in policyData)
      sumExp += math.exp((val - maxLogit).clamp(-10, 10));

    for (int j = 0; j < sample.targetPi.length; j++) {
      if (sample.targetPi[j] > 0) {
        double logSoftmax =
            (policyData[j] - maxLogit) - math.log(sumExp + 1e-10);
        pLoss -= sample.targetPi[j] * logSoftmax;
      }
    }

    // Value Loss
    double vPred = pred['value']!.data[0];
    double vErr = vPred - sample.targetValue!;
    double vLoss = vErr * vErr;

    totalBatchLoss += (pLoss + 0.5 * vLoss);
  }

  optimizer.step();
  return totalBatchLoss / batchSize;
}