ValueEstimatorHook#
- class torchrl.trainers.ValueEstimatorHook(value_estimator: Callable[[TensorDictBase], TensorDictBase])[source]#
A hook that computes value estimates over a collected batch.
Wraps a value estimator module such as
GAEand applies it to the whole collected batch at thepre_epochstage, so that advantage and value-target entries are available when the loss module consumes sub-batches during optimization.In async-collection mode the training loop has no batch to hand to the
pre_epochstage (batchisNone); the hook then passes the batch through untouched, and value estimates are expected to be computed elsewhere (e.g. by a replay-buffer transform).- Parameters:
value_estimator (Callable[[TensorDictBase], TensorDictBase]) – the value estimator to apply to the collected batch, e.g. an instance of
GAE.
Examples
>>> gae = GAE(gamma=0.99, lmbda=0.95, value_network=critic, average_gae=True) >>> value_estimator_hook = ValueEstimatorHook(gae) >>> value_estimator_hook.register(trainer)
- register(trainer: Trainer, name: str = 'value_estimator')[source]#
Registers the hook in the trainer at a default location.
- Parameters:
trainer (Trainer) – the trainer where the hook must be registered.
name (str) – the name of the hook.
Note
To register the hook at another location than the default, use
register_op().