mirror of
https://github.com/wassname/ray.git
synced 2026-09-09 11:32:43 +08:00
MADDPG implementation in RLlib (#5348)
This commit is contained in:
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user