mirror of
https://github.com/wassname/ray.git
synced 2026-09-09 11:32:43 +08:00
[RLlib] Trajectory view API - 03 Fast LSTM + prev actions/rewards (#9950)
This commit is contained in:
+19
-5
@@ -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
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user