trainStep method
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;
}