From 88b4466e5704148e06cccc7d31f7c5a016c69856 Mon Sep 17 00:00:00 2001 From: Johannes Fischer Date: Fri, 6 Aug 2021 18:52:17 +0200 Subject: [PATCH] MInor change in value dice loss, activate print statements, only do EITHER value OR policy update for each batch --- src/value_dice/value_dice.py | 41 ++++++++++++++++++++++++------------ 1 file changed, 27 insertions(+), 14 deletions(-) diff --git a/src/value_dice/value_dice.py b/src/value_dice/value_dice.py index 977d008..d9e29d6 100644 --- a/src/value_dice/value_dice.py +++ b/src/value_dice/value_dice.py @@ -1,3 +1,5 @@ +import numpy as np + import torch import torch.nn as nn from torch.utils.data import DataLoader @@ -213,7 +215,7 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs): # nonlinear loss value_diff = value - discount * value_next - nonlinear_loss = torch.logsumexp(value_diff, dim=0) + nonlinear_loss = torch.logsumexp(value_diff, dim=0) - np.log(len(value_diff)) loss = nonlinear_loss - linear_loss return loss @@ -243,22 +245,33 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs): policy_loss = -loss value_loss = loss - # compute loss and step optimizer - policy_optimizer.zero_grad() - value_optimizer.zero_grad() - policy_loss.backward(retain_graph=True) - value_loss.backward() + # # compute loss and step optimizer + # policy_optimizer.zero_grad() + # 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) + # clip_grad_norm_(policy.policy.parameters(), clip_grad_norm) + # clip_grad_norm_(policy.value.parameters(), clip_grad_norm) - policy_optimizer.step() - value_optimizer.step() + # policy_optimizer.step() + # value_optimizer.step() + + if batch_idx % 2 == 0: + policy_optimizer.zero_grad() + policy_loss.backward() + clip_grad_norm_(policy.policy.parameters(), clip_grad_norm) + policy_optimizer.step() + else: + value_optimizer.zero_grad() + value_loss.backward() + clip_grad_norm_(policy.value.parameters(), clip_grad_norm) + value_optimizer.step() epoch_loss += loss.item() / len(train_dataset) - # if i % print_epoch_every == 0: - # print('Epoch: {}, Training Loss: {}'.format(i, epoch_loss)) + if i % print_epoch_every == 0: + print('Epoch: {}, Training Loss: {}'.format(i, epoch_loss)) # measure cv loss if i % cv_every == 0: @@ -280,8 +293,8 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs): if i % cv_every == 0: writer.add_scalar('cv loss', cv_loss, i) - # if i % print_cv_every == 0: - # print('Epoch: {}, CV Loss: {}'.format(i, cv_loss)) + if i % print_cv_every == 0: + print('Epoch: {}, CV Loss: {}'.format(i, cv_loss)) policy.save_model(filestr)