Masking error. With t*valid_mask, we get the error np.inf*0 = np.inf (#9407)

This commit is contained in:
Tanay Wakhare
2020-07-12 22:59:35 +02:00
committed by GitHub
parent 381c242f6b
commit 3536d8e4b3
+3 -2
View File
@@ -70,7 +70,7 @@ class PPOLoss:
num_valid = torch.sum(valid_mask)
def reduce_mean_valid(t):
return torch.sum(t * valid_mask) / num_valid
return torch.sum(t[valid_mask]) / num_valid
else:
@@ -195,7 +195,8 @@ class ValueNetworkMixin:
np.asarray([prev_reward])),
"is_training": False,
}, [convert_to_torch_tensor(np.asarray([s])) for s in state],
convert_to_torch_tensor(np.asarray([1])))
convert_to_torch_tensor(
np.asarray([1])))
return self.model.value_function()[0]
else: