[rllib] Cleanup RNN support and make it work with multi-GPU optimizer (#2394)

Cleanup: TFPolicyGraph now automatically adds loss input entries for state_in_*, so that graph sub-classes don't need to worry about it.

Multi-GPU support:

Allow setting up model tower replicas with existing state input tensors

Truncate the per-device minibatch slices so that they are always a multiple of max_seq_len.
This commit is contained in:
Eric Liang
2018-07-17 06:55:46 +02:00
committed by GitHub
parent 1b645fcc8b
commit 0cecf6b79c
14 changed files with 163 additions and 138 deletions
@@ -41,16 +41,10 @@ class PGPolicyGraph(TFPolicyGraph):
("advantages", advantages),
]
# LSTM support
for i, ph in enumerate(self.model.state_in):
loss_in.append(("state_in_{}".format(i), ph))
is_training = tf.placeholder_with_default(True, ())
TFPolicyGraph.__init__(
self, obs_space, action_space, sess, obs_input=obs,
action_sampler=action_dist.sample(), loss=loss,
loss_inputs=loss_in, is_training=is_training,
state_inputs=self.model.state_in,
loss_inputs=loss_in, state_inputs=self.model.state_in,
state_outputs=self.model.state_out,
seq_lens=self.model.seq_lens,
max_seq_len=config["model"]["max_seq_len"])