MADDPG implementation in RLlib (#5348)

This commit is contained in:
Wonseok Jeon
2019-08-06 16:22:06 -07:00
committed by Eric Liang
parent 094ec7adbc
commit 281829e712
13 changed files with 736 additions and 23 deletions
+33
View File
@@ -68,6 +68,18 @@ class ReplayBuffer(object):
return (np.array(obses_t), np.array(actions), np.array(rewards),
np.array(obses_tp1), np.array(dones))
@DeveloperAPI
def sample_idxes(self, batch_size):
return [
random.randint(0,
len(self._storage) - 1) for _ in range(batch_size)
]
@DeveloperAPI
def sample_with_idxes(self, idxes):
self._num_sampled += len(idxes)
return self._encode_sample(idxes)
@DeveloperAPI
def sample(self, batch_size):
"""Sample a batch of experiences.
@@ -164,6 +176,27 @@ class PrioritizedReplayBuffer(ReplayBuffer):
res.append(idx)
return res
@DeveloperAPI
def sample_idxes(self, batch_size):
return self._sample_proportional(batch_size)
@DeveloperAPI
def sample_with_idxes(self, idxes, beta):
assert beta > 0
self._num_sampled += len(idxes)
weights = []
p_min = self._it_min.min() / self._it_sum.sum()
max_weight = (p_min * len(self._storage))**(-beta)
for idx in idxes:
p_sample = self._it_sum[idx] / self._it_sum.sum()
weight = (p_sample * len(self._storage))**(-beta)
weights.append(weight / max_weight)
weights = np.array(weights)
encoded_sample = self._encode_sample(idxes)
return tuple(list(encoded_sample) + [weights, idxes])
@DeveloperAPI
def sample(self, batch_size, beta):
"""Sample a batch of experiences.