From 6afb1122773815a1e620608ed95f8e0de0b791f0 Mon Sep 17 00:00:00 2001 From: Johannes Fischer Date: Thu, 5 Aug 2021 18:35:14 +0200 Subject: [PATCH] Bugfix in value dice FIRST backward() has to be called on both, policy and value, before step() is called for either of them --- src/value_dice/value_dice.py | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/src/value_dice/value_dice.py b/src/value_dice/value_dice.py index 6f9ee16..050b893 100644 --- a/src/value_dice/value_dice.py +++ b/src/value_dice/value_dice.py @@ -1,6 +1,7 @@ import torch import torch.nn as nn from torch.utils.data import DataLoader +from torch.nn.utils import clip_grad_norm_ import pickle import itertools from torch.utils.tensorboard import SummaryWriter @@ -164,6 +165,7 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs): train_epochs = config['train_epochs'] train_batch_size = config['train_batch_size'] discount = config['discount'] + clip_grad_norm = config['clip_grad_norm'] cv_every = 1 print_epoch_every = 1000 @@ -236,16 +238,19 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs): loss = f_value_dice_loss(batch) - # TODO: Regularization is done in original source code policy_loss = -loss #+ ORTHOGONAL_REGULARIZER value_loss = loss #+ GRADIENT_REGULARIZER # compute loss and step optimizer policy_optimizer.zero_grad() - policy_loss.backward(retain_graph=True) - policy_optimizer.step() value_optimizer.zero_grad() + policy_loss.backward(retain_graph=True) value_loss.backward() + + clip_grad_norm_(policy.policy.parameters(), clip_grad_norm) + clip_grad_norm_(policy.value.parameters(), clip_grad_norm) + + policy_optimizer.step() value_optimizer.step() epoch_loss += loss.item() / len(train_dataset)