Revert "[RLLib] Execution-Folder Type Annotations (#12760)" (#12886)

This reverts commit becca1424d.
This commit is contained in:
Edward Oakes
2020-12-15 11:03:02 -08:00
committed by GitHub
parent ba12fb1451
commit cde711aaf1
11 changed files with 98 additions and 126 deletions
+3 -5
View File
@@ -1,7 +1,5 @@
from ray.util.iter import LocalIterator
from ray.rllib.policy.sample_batch import SampleBatch, MultiAgentBatch
from ray.rllib.utils.typing import Dict, SampleBatchType
from ray.util.iter_metrics import MetricsContext
# Counters for training progress (keys for metrics.counters).
STEPS_SAMPLED_COUNTER = "num_steps_sampled"
@@ -25,19 +23,19 @@ LEARNER_INFO = "learner"
# Asserts that an object is a type of SampleBatch.
def _check_sample_batch_type(batch: SampleBatchType) -> None:
def _check_sample_batch_type(batch):
if not isinstance(batch, (SampleBatch, MultiAgentBatch)):
raise ValueError("Expected either SampleBatch or MultiAgentBatch, "
"got {}: {}".format(type(batch), batch))
# Returns pipeline global vars that should be periodically sent to each worker.
def _get_global_vars() -> Dict:
def _get_global_vars():
metrics = LocalIterator.get_metrics()
return {"timestep": metrics.counters[STEPS_SAMPLED_COUNTER]}
def _get_shared_metrics() -> MetricsContext:
def _get_shared_metrics():
"""Return shared metrics for the training workflow.
This only applies if this trainer has an execution plan."""
+6 -9
View File
@@ -1,17 +1,15 @@
from typing import List, Optional, Any
from typing import List
import queue
from ray.util.iter import LocalIterator, _NextValueNotReady
from ray.util.iter_metrics import SharedMetrics
from ray.rllib.utils.typing import SampleBatchType
def Concurrently(ops: List[LocalIterator],
*,
mode: str = "round_robin",
output_indexes: Optional[List[int]] = None,
round_robin_weights: Optional[List[int]] = None
) -> LocalIterator[SampleBatchType]:
mode="round_robin",
output_indexes=None,
round_robin_weights=None):
"""Operator that runs the given parent iterators concurrently.
Args:
@@ -93,7 +91,7 @@ class Enqueue:
type(output_queue)))
self.queue = output_queue
def __call__(self, x: Any) -> Any:
def __call__(self, x):
try:
self.queue.put_nowait(x)
except queue.Full:
@@ -101,8 +99,7 @@ class Enqueue:
return x
def Dequeue(input_queue: queue.Queue,
check=lambda: True) -> LocalIterator[SampleBatchType]:
def Dequeue(input_queue: queue.Queue, check=lambda: True):
"""Dequeue data items from a queue.Queue instance.
The dequeue is non-blocking, so Dequeue operations can executed with
+5 -8
View File
@@ -1,4 +1,3 @@
from typing import Dict
import threading
import copy
@@ -9,7 +8,6 @@ from ray.rllib.execution.minibatch_buffer import MinibatchBuffer
from ray.rllib.utils.framework import try_import_tf
from ray.rllib.utils.timer import TimerStat
from ray.rllib.utils.window_stat import WindowStat
from ray.rllib.evaluation.rollout_worker import RolloutWorker
tf1, tf, tfv = try_import_tf()
@@ -23,9 +21,8 @@ class LearnerThread(threading.Thread):
improves overall throughput.
"""
def __init__(self, local_worker: RolloutWorker, minibatch_buffer_size: int,
num_sgd_iter: int, learner_queue_size: int,
learner_queue_timeout: int):
def __init__(self, local_worker, minibatch_buffer_size, num_sgd_iter,
learner_queue_size, learner_queue_timeout):
"""Initialize the learner thread.
Args:
@@ -60,14 +57,14 @@ class LearnerThread(threading.Thread):
self.stopped = False
self.num_steps = 0
def run(self) -> None:
def run(self):
# Switch on eager mode if configured.
if self.local_worker.policy_config.get("framework") in ["tf2", "tfe"]:
tf1.enable_eager_execution()
while not self.stopped:
self.step()
def step(self) -> None:
def step(self):
with self.queue_timer:
batch, _ = self.minibatch_buffer.get()
@@ -80,7 +77,7 @@ class LearnerThread(threading.Thread):
self.outqueue.put((batch.count, self.stats))
self.learner_queue_size.push(self.inqueue.qsize())
def add_learner_metrics(self, result: Dict) -> Dict:
def add_learner_metrics(self, result):
"""Add internal metrics to a trainer result dict."""
def timer_to_ms(timer):
+9 -9
View File
@@ -1,4 +1,4 @@
from typing import Any, List, Dict
from typing import Any, List
import time
from ray.util.iter import LocalIterator
@@ -59,9 +59,9 @@ class CollectMetrics:
"""
def __init__(self,
workers: WorkerSet,
min_history: int = 100,
timeout_seconds: int = 180,
workers,
min_history=100,
timeout_seconds=180,
selected_workers: List["ActorHandle"] = None):
self.workers = workers
self.episode_history = []
@@ -70,7 +70,7 @@ class CollectMetrics:
self.timeout_seconds = timeout_seconds
self.selected_workers = selected_workers
def __call__(self, _: Any) -> Dict:
def __call__(self, _):
# Collect worker metrics.
episodes, self.to_be_collected = collect_episodes(
self.workers.local_worker(),
@@ -124,11 +124,11 @@ class OncePerTimeInterval:
5.00001 # will be greater than 5 seconds
"""
def __init__(self, delay: int):
def __init__(self, delay):
self.delay = delay
self.last_called = 0
def __call__(self, item: Any) -> bool:
def __call__(self, item):
if self.delay <= 0.0:
return True
now = time.time()
@@ -151,11 +151,11 @@ class OncePerTimestepsElapsed:
# will only return after 1000 steps have elapsed
"""
def __init__(self, delay_steps: int):
def __init__(self, delay_steps):
self.delay_steps = delay_steps
self.last_called = 0
def __call__(self, item: Any) -> bool:
def __call__(self, item):
if self.delay_steps <= 0:
return True
metrics = _get_shared_metrics()
+2 -11
View File
@@ -1,19 +1,10 @@
from typing import Any, Tuple
import queue
class MinibatchBuffer:
"""Ring buffer of recent data batches for minibatch SGD.
This is for use with AsyncSamplesOptimizer.
"""
def __init__(self,
inqueue: queue.Queue,
size: int,
timeout: float,
num_passes: int,
init_num_passes: int = 1):
def __init__(self, inqueue, size, timeout, num_passes, init_num_passes=1):
"""Initialize a minibatch buffer.
Args:
@@ -32,7 +23,7 @@ class MinibatchBuffer:
self.ttl = [0] * size
self.idx = 0
def get(self) -> Tuple[Any, bool]:
def get(self):
"""Get a new batch from the internal ring buffer.
Returns:
+15 -16
View File
@@ -12,7 +12,6 @@ from ray.rllib.execution.multi_gpu_impl import LocalSyncParallelOptimizer
from ray.rllib.utils.annotations import override
from ray.rllib.utils.framework import try_import_tf
from ray.rllib.utils.timer import TimerStat
from ray.rllib.evaluation.rollout_worker import RolloutWorker
tf1, tf, tfv = try_import_tf()
@@ -26,17 +25,17 @@ class TFMultiGPULearner(LearnerThread):
"""
def __init__(self,
local_worker: RolloutWorker,
num_gpus: int = 1,
lr: float = 0.0005,
train_batch_size: int = 500,
num_data_loader_buffers: int = 1,
minibatch_buffer_size: int = 1,
num_sgd_iter: int = 1,
learner_queue_size: int = 16,
learner_queue_timeout: int = 300,
num_data_load_threads: int = 16,
_fake_gpus: bool = False):
local_worker,
num_gpus=1,
lr=0.0005,
train_batch_size=500,
num_data_loader_buffers=1,
minibatch_buffer_size=1,
num_sgd_iter=1,
learner_queue_size=16,
learner_queue_timeout=300,
num_data_load_threads=16,
_fake_gpus=False):
"""Initialize a multi-gpu learner thread.
Args:
@@ -122,7 +121,7 @@ class TFMultiGPULearner(LearnerThread):
learner_queue_timeout, num_sgd_iter)
@override(LearnerThread)
def step(self) -> None:
def step(self):
assert self.loader_thread.is_alive()
with self.load_wait_timer:
opt, released = self.minibatch_buffer.get()
@@ -140,7 +139,7 @@ class TFMultiGPULearner(LearnerThread):
class _LoaderThread(threading.Thread):
def __init__(self, learner: LearnerThread, share_stats: bool):
def __init__(self, learner, share_stats):
threading.Thread.__init__(self)
self.learner = learner
self.daemon = True
@@ -151,11 +150,11 @@ class _LoaderThread(threading.Thread):
self.queue_timer = TimerStat()
self.load_timer = TimerStat()
def run(self) -> None:
def run(self):
while True:
self._step()
def _step(self) -> None:
def _step(self):
s = self.learner
with self.queue_timer:
batch = s.inqueue.get()
+22 -23
View File
@@ -3,7 +3,7 @@ import logging
import numpy as np
import platform
import random
from typing import List, Dict
from typing import List
# Import ray before psutil will make sure we use psutil's bundled version
import ray # noqa F401
@@ -64,11 +64,11 @@ class ReplayBuffer:
self._evicted_hit_stats = WindowStat("evicted_hit", 1000)
self._est_size_bytes = 0
def __len__(self) -> int:
def __len__(self):
return len(self._storage)
@DeveloperAPI
def add(self, item: SampleBatchType, weight: float) -> None:
def add(self, item: SampleBatchType, weight: float):
warn_replay_buffer_size(
item=item, num_items=self._maxsize / item.count)
assert item.count > 0, item
@@ -116,7 +116,7 @@ class ReplayBuffer:
return self._encode_sample(idxes)
@DeveloperAPI
def stats(self, debug=False) -> dict:
def stats(self, debug=False):
data = {
"added_count": self._num_timesteps_added,
"sampled_count": self._num_timesteps_sampled,
@@ -156,7 +156,7 @@ class PrioritizedReplayBuffer(ReplayBuffer):
self._prio_change_stats = WindowStat("reprio", 1000)
@DeveloperAPI
def add(self, item: SampleBatchType, weight: float) -> None:
def add(self, item: SampleBatchType, weight: float):
idx = self._next_idx
super(PrioritizedReplayBuffer, self).add(item, weight)
if weight is None:
@@ -164,7 +164,7 @@ class PrioritizedReplayBuffer(ReplayBuffer):
self._it_sum[idx] = weight**self._alpha
self._it_min[idx] = weight**self._alpha
def _sample_proportional(self, num_items: int) -> List[int]:
def _sample_proportional(self, num_items: int):
res = []
for _ in range(num_items):
# TODO(szymon): should we ensure no repeats?
@@ -215,8 +215,7 @@ class PrioritizedReplayBuffer(ReplayBuffer):
return batch
@DeveloperAPI
def update_priorities(self, idxes: List[int],
priorities: List[float]) -> None:
def update_priorities(self, idxes, priorities):
"""Update priorities of sampled transitions.
sets priority of transition at index idxes[i] in buffer
@@ -243,7 +242,7 @@ class PrioritizedReplayBuffer(ReplayBuffer):
self._max_priority = max(self._max_priority, priority)
@DeveloperAPI
def stats(self, debug: bool = False) -> Dict:
def stats(self, debug=False):
parent = ReplayBuffer.stats(self, debug)
if debug:
parent.update(self._prio_change_stats.stats())
@@ -261,15 +260,15 @@ class LocalReplayBuffer(ParallelIteratorWorker):
may be created to increase parallelism."""
def __init__(self,
num_shards: int = 1,
learning_starts: int = 1000,
buffer_size: int = 10000,
replay_batch_size: int = 1,
prioritized_replay_alpha: float = 0.6,
prioritized_replay_beta: float = 0.4,
prioritized_replay_eps: float = 1e-6,
replay_mode: str = "independent",
replay_sequence_length: int = 1):
num_shards=1,
learning_starts=1000,
buffer_size=10000,
replay_batch_size=1,
prioritized_replay_alpha=0.6,
prioritized_replay_beta=0.4,
prioritized_replay_eps=1e-6,
replay_mode="independent",
replay_sequence_length=1):
self.replay_starts = learning_starts // num_shards
self.buffer_size = buffer_size // num_shards
self.replay_batch_size = replay_batch_size
@@ -319,10 +318,10 @@ class LocalReplayBuffer(ParallelIteratorWorker):
global _local_replay_buffer
return _local_replay_buffer
def get_host(self) -> str:
def get_host(self):
return platform.node()
def add_batch(self, batch: SampleBatchType) -> None:
def add_batch(self, batch):
# Make a copy so the replay buffer doesn't pin plasma memory.
batch = batch.copy()
# Handle everything as if multiagent
@@ -343,7 +342,7 @@ class LocalReplayBuffer(ParallelIteratorWorker):
self.replay_buffers[policy_id].add(s, weight=weight)
self.num_added += batch.count
def replay(self) -> SampleBatchType:
def replay(self):
if self._fake_batch:
fake_batch = SampleBatch(self._fake_batch)
return MultiAgentBatch({
@@ -365,7 +364,7 @@ class LocalReplayBuffer(ParallelIteratorWorker):
beta=self.prioritized_replay_beta)
return MultiAgentBatch(samples, self.replay_batch_size)
def update_priorities(self, prio_dict: Dict) -> None:
def update_priorities(self, prio_dict):
with self.update_priorities_timer:
for policy_id, (batch_indexes, td_errors) in prio_dict.items():
new_priorities = (
@@ -373,7 +372,7 @@ class LocalReplayBuffer(ParallelIteratorWorker):
self.replay_buffers[policy_id].update_priorities(
batch_indexes, new_priorities)
def stats(self, debug: bool = False) -> Dict:
def stats(self, debug=False):
stat = {
"add_batch_time_ms": round(1000 * self.add_batch_timer.mean, 3),
"replay_time_ms": round(1000 * self.replay_timer.mean, 3),
+9 -11
View File
@@ -1,4 +1,4 @@
from typing import List, Any, Optional
from typing import List
import random
from ray.util.iter import from_actors, LocalIterator, _NextValueNotReady
@@ -55,7 +55,7 @@ class StoreToReplayBuffer:
def Replay(*,
local_buffer: LocalReplayBuffer = None,
actors: List["ActorHandle"] = None,
num_async: int = 4) -> LocalIterator[SampleBatchType]:
num_async=4):
"""Replay experiences from the given buffer or actors.
This should be combined with the StoreToReplayActors operation using the
@@ -99,10 +99,10 @@ def Replay(*,
class WaitUntilTimestepsElapsed:
"""Callable that returns True once a given number of timesteps are hit."""
def __init__(self, target_num_timesteps: int):
def __init__(self, target_num_timesteps):
self.target_num_timesteps = target_num_timesteps
def __call__(self, item: Any) -> bool:
def __call__(self, item):
metrics = _get_shared_metrics()
ts = metrics.counters[STEPS_SAMPLED_COUNTER]
return ts > self.target_num_timesteps
@@ -112,9 +112,7 @@ class WaitUntilTimestepsElapsed:
class SimpleReplayBuffer:
"""Simple replay buffer that operates over batches."""
def __init__(self,
num_slots: int,
replay_proportion: Optional[float] = None):
def __init__(self, num_slots, replay_proportion: float = None):
"""Initialize SimpleReplayBuffer.
Args:
@@ -124,7 +122,7 @@ class SimpleReplayBuffer:
self.replay_batches = []
self.replay_index = 0
def add_batch(self, sample_batch: SampleBatchType) -> None:
def add_batch(self, sample_batch):
warn_replay_buffer_size(item=sample_batch, num_items=self.num_slots)
if self.num_slots > 0:
if len(self.replay_batches) < self.num_slots:
@@ -134,7 +132,7 @@ class SimpleReplayBuffer:
self.replay_index += 1
self.replay_index %= self.num_slots
def replay(self) -> SampleBatchType:
def replay(self):
return random.choice(self.replay_batches)
@@ -147,7 +145,7 @@ class MixInReplay:
number of replay slots.
"""
def __init__(self, num_slots: int, replay_proportion: float):
def __init__(self, num_slots, replay_proportion: float):
"""Initialize MixInReplay.
Args:
@@ -173,7 +171,7 @@ class MixInReplay:
self.replay_buffer = SimpleReplayBuffer(num_slots)
self.replay_proportion = replay_proportion
def __call__(self, sample_batch: SampleBatchType) -> List[SampleBatchType]:
def __call__(self, sample_batch):
# Put in replay buffer if enabled.
self.replay_buffer.add_batch(sample_batch)
+9 -13
View File
@@ -1,5 +1,4 @@
import operator
from typing import Any, Optional
class SegmentTree:
@@ -29,10 +28,7 @@ class SegmentTree:
`tree[0]` accesses `internal_array[4]` in the above example.
"""
def __init__(self,
capacity: int,
operation: Any,
neutral_element: Optional[Any] = None):
def __init__(self, capacity, operation, neutral_element=None):
"""Initializes a Segment Tree object.
Args:
@@ -56,7 +52,7 @@ class SegmentTree:
self.value = [self.neutral_element for _ in range(2 * capacity)]
self.operation = operation
def reduce(self, start: int = 0, end: Optional[int] = None) -> Any:
def reduce(self, start=0, end=None):
"""Applies `self.operation` to subsequence of our values.
Subsequence is contiguous, includes `start` and excludes `end`.
@@ -126,7 +122,7 @@ class SegmentTree:
return result
def __setitem__(self, idx: int, val: float) -> None:
def __setitem__(self, idx, val):
"""
Inserts/overwrites a value in/into the tree.
@@ -151,7 +147,7 @@ class SegmentTree:
self.value[update_idx + 1])
idx = idx >> 1 # Divide by 2 (faster than division).
def __getitem__(self, idx: int) -> Any:
def __getitem__(self, idx):
assert 0 <= idx < self.capacity
return self.value[idx + self.capacity]
@@ -159,15 +155,15 @@ class SegmentTree:
class SumSegmentTree(SegmentTree):
"""A SegmentTree with the reduction `operation`=operator.add."""
def __init__(self, capacity: int):
def __init__(self, capacity):
super(SumSegmentTree, self).__init__(
capacity=capacity, operation=operator.add)
def sum(self, start: int = 0, end: Optional[Any] = None) -> Any:
def sum(self, start=0, end=None):
"""Returns the sum over a sub-segment of the tree."""
return self.reduce(start, end)
def find_prefixsum_idx(self, prefixsum: float) -> int:
def find_prefixsum_idx(self, prefixsum):
"""Finds highest i, for which: sum(arr[0]+..+arr[i - i]) <= prefixsum.
Args:
@@ -192,9 +188,9 @@ class SumSegmentTree(SegmentTree):
class MinSegmentTree(SegmentTree):
def __init__(self, capacity: int):
def __init__(self, capacity):
super(MinSegmentTree, self).__init__(capacity=capacity, operation=min)
def min(self, start: int = 0, end: Optional[Any] = None) -> Any:
def min(self, start=0, end=None):
"""Returns min(arr[start], ..., arr[end])"""
return self.reduce(start, end)
+11 -12
View File
@@ -2,7 +2,7 @@ from collections import defaultdict
import logging
import numpy as np
import math
from typing import List, Tuple, Any
from typing import List
import ray
from ray.rllib.evaluation.metrics import get_learner_stats, LEARNER_STATS_KEY
@@ -18,7 +18,7 @@ from ray.rllib.policy.sample_batch import SampleBatch, DEFAULT_POLICY_ID, \
MultiAgentBatch
from ray.rllib.utils.framework import try_import_tf
from ray.rllib.utils.sgd import do_minibatch_sgd, averaged
from ray.rllib.utils.typing import PolicyID, SampleBatchType, ModelGradients
from ray.rllib.utils.typing import PolicyID, SampleBatchType
tf1, tf, tfv = try_import_tf()
@@ -242,10 +242,10 @@ class ComputeGradients:
Updates the LEARNER_INFO info field in the local iterator context.
"""
def __init__(self, workers: WorkerSet):
def __init__(self, workers):
self.workers = workers
def __call__(self, samples: SampleBatchType) -> Tuple[ModelGradients, int]:
def __call__(self, samples: SampleBatchType):
_check_sample_batch_type(samples)
metrics = _get_shared_metrics()
with metrics.timers[COMPUTE_GRADS_TIMER]:
@@ -283,7 +283,7 @@ class ApplyGradients:
self.policies = policies or workers.local_worker().policies_to_train
self.update_all = update_all
def __call__(self, item: Tuple[ModelGradients, int]) -> None:
def __call__(self, item):
if not isinstance(item, tuple) or len(item) != 2:
raise ValueError(
"Input must be a tuple of (grad_dict, count), got {}".format(
@@ -333,8 +333,7 @@ class AverageGradients:
{"var_0": ..., ...}, 1600 # averaged grads, summed batch count
"""
def __call__(self, gradients: List[Tuple[ModelGradients, int]]
) -> Tuple[ModelGradients, int]:
def __call__(self, gradients):
acc = None
sum_count = 0
for grad, count in gradients:
@@ -367,10 +366,10 @@ class UpdateTargetNetwork:
"""
def __init__(self,
workers: WorkerSet,
target_update_freq: int,
by_steps_trained: bool = False,
policies: List[PolicyID] = frozenset([])):
workers,
target_update_freq,
by_steps_trained=False,
policies=frozenset([])):
self.workers = workers
self.target_update_freq = target_update_freq
self.policies = (policies or workers.local_worker().policies_to_train)
@@ -379,7 +378,7 @@ class UpdateTargetNetwork:
else:
self.metric = STEPS_SAMPLED_COUNTER
def __call__(self, _: Any) -> None:
def __call__(self, _):
metrics = _get_shared_metrics()
cur_ts = metrics.counters[self.metric]
last_update = metrics.counters[LAST_TARGET_UPDATE_TS]
+7 -9
View File
@@ -1,6 +1,6 @@
import logging
import platform
from typing import List, Dict, Any
from typing import List
import ray
from ray.rllib.execution.common import STEPS_SAMPLED_COUNTER, \
@@ -9,9 +9,8 @@ from ray.rllib.execution.replay_ops import MixInReplay
from ray.rllib.execution.rollout_ops import ParallelRollouts, ConcatBatches
from ray.rllib.utils.actors import create_colocated
from ray.util.iter import ParallelIterator, ParallelIteratorWorker, \
from_actors, LocalIterator
from ray.rllib.utils.typing import SampleBatchType, ModelWeights
from ray.rllib.evaluation.worker_set import WorkerSet
from_actors
from ray.rllib.utils.typing import SampleBatchType
logger = logging.getLogger(__name__)
@@ -26,7 +25,7 @@ class Aggregator(ParallelIteratorWorker):
work to be offloaded to these actors instead of run in the learner.
"""
def __init__(self, config: Dict,
def __init__(self, config: dict,
rollout_group: "ParallelIterator[SampleBatchType]"):
self.weights = None
self.global_vars = None
@@ -61,16 +60,15 @@ class Aggregator(ParallelIteratorWorker):
super().__init__(generator, repeat=False)
def get_host(self) -> str:
def get_host(self):
return platform.node()
def set_weights(self, weights: ModelWeights, global_vars: Dict) -> None:
def set_weights(self, weights, global_vars):
self.weights = weights
self.global_vars = global_vars
def gather_experiences_tree_aggregation(workers: WorkerSet,
config: Dict) -> "LocalIterator[Any]":
def gather_experiences_tree_aggregation(workers, config):
"""Tree aggregation version of gather_experiences_directly()."""
rollouts = ParallelRollouts(workers, mode="raw")