[rllib] Rename sample_batch_size => rollout_fragment_length (#7503)

* bulk rename

* deprecation warn

* update doc

* update fig

* line length

* rename

* make pytest comptaible

* fix test

* fi sys

* rename

* wip

* fix more

* lint

* update svg

* comments

* lint

* fix use of batch steps
This commit is contained in:
Eric Liang
2020-03-14 12:05:04 -07:00
committed by GitHub
parent 53549314c5
commit dd70720578
84 changed files with 323 additions and 271 deletions
+8 -6
View File
@@ -41,7 +41,7 @@ class LocalMultiGPUOptimizer(PolicyOptimizer):
workers,
sgd_batch_size=128,
num_sgd_iter=10,
sample_batch_size=200,
rollout_fragment_length=200,
num_envs_per_worker=1,
train_batch_size=1024,
num_gpus=0,
@@ -53,7 +53,8 @@ class LocalMultiGPUOptimizer(PolicyOptimizer):
workers (WorkerSet): all workers
sgd_batch_size (int): SGD minibatch size within train batch size
num_sgd_iter (int): number of passes to learn on per train batch
sample_batch_size (int): size of batches to sample from workers
rollout_fragment_length (int): size of batches to sample from
workers.
num_envs_per_worker (int): num envs in each rollout worker
train_batch_size (int): size of batches to learn on
num_gpus (int): number of GPUs to use for data-parallel SGD
@@ -67,7 +68,7 @@ class LocalMultiGPUOptimizer(PolicyOptimizer):
self.batch_size = sgd_batch_size
self.num_sgd_iter = num_sgd_iter
self.num_envs_per_worker = num_envs_per_worker
self.sample_batch_size = sample_batch_size
self.rollout_fragment_length = rollout_fragment_length
self.train_batch_size = train_batch_size
self.shuffle_sequences = shuffle_sequences
if not num_gpus:
@@ -132,9 +133,10 @@ class LocalMultiGPUOptimizer(PolicyOptimizer):
with self.sample_timer:
if self.workers.remote_workers():
samples = collect_samples(
self.workers.remote_workers(), self.sample_batch_size,
self.num_envs_per_worker, self.train_batch_size)
samples = collect_samples(self.workers.remote_workers(),
self.rollout_fragment_length,
self.num_envs_per_worker,
self.train_batch_size)
if samples.count > self.train_batch_size * 2:
logger.info(
"Collected more training samples than expected "