mirror of
https://github.com/wassname/ray.git
synced 2026-09-09 11:32:43 +08:00
[RLlib] Execution Annotation (#13036)
This commit is contained in:
@@ -12,6 +12,7 @@ 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()
|
||||
|
||||
@@ -25,17 +26,17 @@ class TFMultiGPULearner(LearnerThread):
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
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):
|
||||
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):
|
||||
"""Initialize a multi-gpu learner thread.
|
||||
|
||||
Args:
|
||||
@@ -121,7 +122,7 @@ class TFMultiGPULearner(LearnerThread):
|
||||
learner_queue_timeout, num_sgd_iter)
|
||||
|
||||
@override(LearnerThread)
|
||||
def step(self):
|
||||
def step(self) -> None:
|
||||
assert self.loader_thread.is_alive()
|
||||
with self.load_wait_timer:
|
||||
opt, released = self.minibatch_buffer.get()
|
||||
@@ -139,7 +140,7 @@ class TFMultiGPULearner(LearnerThread):
|
||||
|
||||
|
||||
class _LoaderThread(threading.Thread):
|
||||
def __init__(self, learner, share_stats):
|
||||
def __init__(self, learner: LearnerThread, share_stats: bool):
|
||||
threading.Thread.__init__(self)
|
||||
self.learner = learner
|
||||
self.daemon = True
|
||||
@@ -150,11 +151,11 @@ class _LoaderThread(threading.Thread):
|
||||
self.queue_timer = TimerStat()
|
||||
self.load_timer = TimerStat()
|
||||
|
||||
def run(self):
|
||||
def run(self) -> None:
|
||||
while True:
|
||||
self._step()
|
||||
|
||||
def _step(self):
|
||||
def _step(self) -> None:
|
||||
s = self.learner
|
||||
with self.queue_timer:
|
||||
batch = s.inqueue.get()
|
||||
|
||||
Reference in New Issue
Block a user