mirror of
https://github.com/wassname/ray.git
synced 2026-09-10 12:38:43 +08:00
[RLlib] Sample batch docs and cleanup. (#8778)
This commit is contained in:
@@ -75,30 +75,47 @@ class MultiAgentSampleBatchBuilder:
|
||||
def __init__(self, policy_map, clip_rewards, callbacks):
|
||||
"""Initialize a MultiAgentSampleBatchBuilder.
|
||||
|
||||
Arguments:
|
||||
policy_map (dict): Maps policy ids to policy instances.
|
||||
clip_rewards (bool): Whether to clip rewards before postprocessing.
|
||||
Args:
|
||||
policy_map (Dict[str,Policy]): Maps policy ids to policy instances.
|
||||
clip_rewards (Union[bool,float]): Whether to clip rewards before
|
||||
postprocessing (at +/-1.0) or the actual value to +/- clip.
|
||||
callbacks (DefaultCallbacks): RLlib callbacks.
|
||||
"""
|
||||
|
||||
self.policy_map = policy_map
|
||||
self.clip_rewards = clip_rewards
|
||||
# Build the Policies' SampleBatchBuilders.
|
||||
self.policy_builders = {
|
||||
k: SampleBatchBuilder()
|
||||
for k in policy_map.keys()
|
||||
}
|
||||
# Whenever we observe a new agent, add a new SampleBatchBuilder for
|
||||
# this agent.
|
||||
self.agent_builders = {}
|
||||
# Internal agent-to-policy map.
|
||||
self.agent_to_policy = {}
|
||||
self.callbacks = callbacks
|
||||
self.count = 0 # increment this manually
|
||||
# Number of "inference" steps taken in the environment.
|
||||
# Regardless of the number of agents involved in each of these steps.
|
||||
self.count = 0
|
||||
|
||||
def total(self):
|
||||
"""Returns summed number of steps across all agent buffers."""
|
||||
"""Returns the total number of steps taken in the env (all agents).
|
||||
|
||||
Returns:
|
||||
int: The number of steps taken in total in the environment over all
|
||||
agents.
|
||||
"""
|
||||
|
||||
return sum(a.count for a in self.agent_builders.values())
|
||||
|
||||
def has_pending_agent_data(self):
|
||||
"""Returns whether there is pending unprocessed data."""
|
||||
"""Returns whether there is pending unprocessed data.
|
||||
|
||||
Returns:
|
||||
bool: True if there is at least one per-agent builder (with data
|
||||
in it).
|
||||
"""
|
||||
|
||||
return len(self.agent_builders) > 0
|
||||
|
||||
@@ -115,32 +132,37 @@ class MultiAgentSampleBatchBuilder:
|
||||
if agent_id not in self.agent_builders:
|
||||
self.agent_builders[agent_id] = SampleBatchBuilder()
|
||||
self.agent_to_policy[agent_id] = policy_id
|
||||
builder = self.agent_builders[agent_id]
|
||||
builder.add_values(**values)
|
||||
self.agent_builders[agent_id].add_values(**values)
|
||||
|
||||
def postprocess_batch_so_far(self, episode):
|
||||
def postprocess_batch_so_far(self, episode=None):
|
||||
"""Apply policy postprocessors to any unprocessed rows.
|
||||
|
||||
This pushes the postprocessed per-agent batches onto the per-policy
|
||||
builders, clearing per-agent state.
|
||||
|
||||
Args:
|
||||
episode (Optional[MultiAgentEpisode]): Current MultiAgentEpisode
|
||||
object.
|
||||
episode (Optional[MultiAgentEpisode]): The Episode object that
|
||||
holds this MultiAgentBatchBuilder object.
|
||||
"""
|
||||
|
||||
# Materialize the batches so far
|
||||
# Materialize the batches so far.
|
||||
pre_batches = {}
|
||||
for agent_id, builder in self.agent_builders.items():
|
||||
pre_batches[agent_id] = (
|
||||
self.policy_map[self.agent_to_policy[agent_id]],
|
||||
builder.build_and_reset())
|
||||
|
||||
# Apply postprocessor
|
||||
# Apply postprocessor.
|
||||
post_batches = {}
|
||||
if self.clip_rewards:
|
||||
if self.clip_rewards is True:
|
||||
for _, (_, pre_batch) in pre_batches.items():
|
||||
pre_batch["rewards"] = np.sign(pre_batch["rewards"])
|
||||
elif self.clip_rewards:
|
||||
for _, (_, pre_batch) in pre_batches.items():
|
||||
pre_batch["rewards"] = np.clip(
|
||||
pre_batch["rewards"],
|
||||
a_min=-self.clip_rewards,
|
||||
a_max=self.clip_rewards)
|
||||
for agent_id, (_, pre_batch) in pre_batches.items():
|
||||
other_batches = pre_batches.copy()
|
||||
del other_batches[agent_id]
|
||||
@@ -193,15 +215,19 @@ class MultiAgentSampleBatchBuilder:
|
||||
"Alternatively, set no_done_at_end=True to allow this.")
|
||||
|
||||
@DeveloperAPI
|
||||
def build_and_reset(self, episode):
|
||||
def build_and_reset(self, episode=None):
|
||||
"""Returns the accumulated sample batches for each policy.
|
||||
|
||||
Any unprocessed rows will be first postprocessed with a policy
|
||||
postprocessor. The internal state of this builder will be reset.
|
||||
|
||||
Args:
|
||||
episode (Optional[MultiAgentEpisode]): Current MultiAgentEpisode
|
||||
object.
|
||||
episode (Optional[MultiAgentEpisode]): The Episode object that
|
||||
holds this MultiAgentBatchBuilder object or None.
|
||||
|
||||
Returns:
|
||||
MultiAgentBatch: Returns the accumulated sample batches for each
|
||||
policy.
|
||||
"""
|
||||
|
||||
self.postprocess_batch_so_far(episode)
|
||||
|
||||
Reference in New Issue
Block a user