mirror of
https://github.com/wassname/ray.git
synced 2026-08-12 12:20:11 +08:00
[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:
@@ -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 "
|
||||
|
||||
Reference in New Issue
Block a user