[RLlib] Trajectory view API - 03 Fast LSTM + prev actions/rewards (#9950)

This commit is contained in:
Sven Mika
2020-08-21 12:35:16 +02:00
committed by GitHub
parent 92664249e8
commit e968b52cb7
25 changed files with 1230 additions and 413 deletions
+19 -5
View File
@@ -63,7 +63,8 @@ def minibatches(samples, sgd_minibatch_size):
raise NotImplementedError(
"Minibatching not implemented for multi-agent in simple mode")
if "state_in_0" in samples.data:
# Replace with `if samples.seq_lens` check.
if "state_in_0" in samples.data or "state_out_0" in samples.data:
if log_once("not_shuffling_rnn_data_in_simple_mode"):
logger.warning("Not shuffling RNN data for SGD in simple mode")
else:
@@ -71,9 +72,22 @@ def minibatches(samples, sgd_minibatch_size):
i = 0
slices = []
while i < samples.count:
slices.append((i, i + sgd_minibatch_size))
i += sgd_minibatch_size
if samples.seq_lens:
seq_no = 0
while i < samples.count:
seq_no_end = seq_no
actual_count = 0
while actual_count < sgd_minibatch_size and len(
samples.seq_lens) > seq_no_end:
actual_count += samples.seq_lens[seq_no_end]
seq_no_end += 1
slices.append((seq_no, seq_no_end))
i += actual_count
seq_no = seq_no_end
else:
while i < samples.count:
slices.append((i, i + sgd_minibatch_size))
i += sgd_minibatch_size
random.shuffle(slices)
for i, j in slices:
@@ -100,7 +114,7 @@ def do_minibatch_sgd(samples, policies, local_worker, num_sgd_iter,
samples = MultiAgentBatch({DEFAULT_POLICY_ID: samples}, samples.count)
fetches = {}
for policy_id, policy in policies.items():
for policy_id in policies.keys():
if policy_id not in samples.policy_batches:
continue
+3
View File
@@ -43,6 +43,9 @@ EnvID = int
# Represents an episode id.
EpisodeID = int
# Represents an "unroll" (maybe across different sub-envs in a vector env).
UnrollID = int
# A dict keyed by agent ids, e.g. {"agent-1": value}.
MultiAgentDict = Dict[AgentID, Any]