Minor formatting

This commit is contained in:
Johannes Fischer
2021-08-05 18:37:54 +02:00
parent 6afb112277
commit 66bfba3986

View File

@@ -94,8 +94,8 @@ class ValueDicePolicy(IntersimPolicy):
transforms (dict): dictionary of transforms to apply to different fields transforms (dict): dictionary of transforms to apply to different fields
""" """
super(ValueDicePolicy, self).__init__(config, transforms) super(ValueDicePolicy, self).__init__(config, transforms)
self._policy = IntersimStateNet(config["policy_net"]) self._policy = IntersimStateNet(config['policy_net'])
self._value = IntersimStateActionNet(config["value_net"]) self._value = IntersimStateActionNet(config['value_net'])
@property @property
def value(self): def value(self):
@@ -204,9 +204,9 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs):
next_state = policy.transform_observation(next_state) next_state = policy.transform_observation(next_state)
# evaluate value network # evaluate value network
value = (policy.value(state)) value = policy.value(state)
value_init = (policy.value(initial_state)) value_init = policy.value(initial_state)
value_next = (policy.value(next_state)) value_next = policy.value(next_state)
# linear loss # linear loss
linear_loss = (1 - discount) * torch.mean(value_init) linear_loss = (1 - discount) * torch.mean(value_init)
@@ -238,8 +238,10 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs):
loss = f_value_dice_loss(batch) loss = f_value_dice_loss(batch)
policy_loss = -loss #+ ORTHOGONAL_REGULARIZER # In original implementation policy is regularized with orthogonal regularization,
value_loss = loss #+ GRADIENT_REGULARIZER # value with L2 regularization on gradients
policy_loss = -loss
value_loss = loss
# compute loss and step optimizer # compute loss and step optimizer
policy_optimizer.zero_grad() policy_optimizer.zero_grad()