From b50816641948937406937278ccc7fc3ef3628639 Mon Sep 17 00:00:00 2001 From: Tomasz Wrona Date: Sun, 26 Apr 2020 03:49:09 +0200 Subject: [PATCH] Copy initial state of an RNN to a CPU before converting it to a NumPy array (#8097) --- rllib/policy/torch_policy.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/rllib/policy/torch_policy.py b/rllib/policy/torch_policy.py index 861654b0a..ccdb08ff2 100644 --- a/rllib/policy/torch_policy.py +++ b/rllib/policy/torch_policy.py @@ -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.