Bugfix in value dice

FIRST backward() has to be called on both, policy and value, before step() is called for either of them
This commit is contained in:
Johannes Fischer
2021-08-05 18:35:14 +02:00
parent bf4c19a4d0
commit 6afb112277

View File

@@ -1,6 +1,7 @@
import torch import torch
import torch.nn as nn import torch.nn as nn
from torch.utils.data import DataLoader from torch.utils.data import DataLoader
from torch.nn.utils import clip_grad_norm_
import pickle import pickle
import itertools import itertools
from torch.utils.tensorboard import SummaryWriter 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_epochs = config['train_epochs']
train_batch_size = config['train_batch_size'] train_batch_size = config['train_batch_size']
discount = config['discount'] discount = config['discount']
clip_grad_norm = config['clip_grad_norm']
cv_every = 1 cv_every = 1
print_epoch_every = 1000 print_epoch_every = 1000
@@ -236,16 +238,19 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs):
loss = f_value_dice_loss(batch) loss = f_value_dice_loss(batch)
# TODO: Regularization is done in original source code
policy_loss = -loss #+ ORTHOGONAL_REGULARIZER policy_loss = -loss #+ ORTHOGONAL_REGULARIZER
value_loss = loss #+ GRADIENT_REGULARIZER value_loss = loss #+ GRADIENT_REGULARIZER
# compute loss and step optimizer # compute loss and step optimizer
policy_optimizer.zero_grad() policy_optimizer.zero_grad()
policy_loss.backward(retain_graph=True)
policy_optimizer.step()
value_optimizer.zero_grad() value_optimizer.zero_grad()
policy_loss.backward(retain_graph=True)
value_loss.backward() 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() value_optimizer.step()
epoch_loss += loss.item() / len(train_dataset) epoch_loss += loss.item() / len(train_dataset)