Copy initial state of an RNN to a CPU before converting it to a NumPy array (#8097)

This commit is contained in:
Tomasz Wrona
2020-04-25 18:49:09 -07:00
committed by GitHub
parent b506f87117
commit b508166419
+3 -1
View File
@@ -333,7 +333,9 @@ class TorchPolicy(Policy):
@override(Policy)
def get_initial_state(self):
return [s.numpy() for s in self.model.get_initial_state()]
return [
s.cpu().detach().numpy() for s in self.model.get_initial_state()
]
def extra_grad_process(self, optimizer, loss):
"""Called after each optimizer.zero_grad() + loss.backward() call.