FQLLoss#
- class torchrl.objectives.FQLLoss(*args, **kwargs)[source]#
Flow Q-learning for normalized continuous actions.
Implements https://arxiv.org/abs/2502.02538. The flow learns the behavior distribution; the one-step actor balances distillation and mean critic value. Only the critics have target parameters, updated with
SoftUpdate.- Parameters:
flow_policy (FlowMatchingPolicy or TensorDictSequential) – behavior flow policy, optionally preceded by TensorDict observation encoders.
actor_network (OneStepPolicy or TensorDictSequential) – noise-conditioned one-step policy, optionally preceded by TensorDict observation encoders. Each sequence must end in the corresponding policy. Observation keys are read from these modules’
in_keys.qvalue_network (TensorDictModule or list) – critic reading observations and actions and writing
state_action_value. A single critic is duplicated; a list supplies independently initialized critics.
- Keyword Arguments:
num_qvalue_nets (int, optional) – number of critics. Defaults to 2.
alpha (float, optional) – distillation weight. Defaults to 10.
q_aggregation (str, optional) – target critic aggregation,
meanormin. The actor always uses the mean. Defaults tomean.normalize_q_loss (bool, optional) – divide actor Q loss by the detached mean absolute Q value. Defaults to False.
reduction (str, optional) –
none,meanorsum. Defaults tomean. Action and critic coordinates are always averaged first.
- default_keys#
alias of
_AcceptedKeys
- forward(tensordict: TensorDictBase = None) TensorDict[source]#
It is designed to read an input TensorDict and return another tensordict with loss keys named “loss*”.
Splitting the loss in its component can then be used by the trainer to log the various loss values throughout training. Other scalars present in the output tensordict will be logged too.
- Parameters:
tensordict – an input tensordict with the values required to compute the loss.
- Returns:
A new tensordict with no batch dimension containing various loss scalars which will be named “loss*”. It is essential that the losses are returned with this name as they will be read by the trainer before backpropagation.
- make_value_estimator(value_type=None, **hyperparams)[source]#
Value-function constructor.
If the non-default value function is wanted, it must be built using this method.
- Parameters:
value_type (ValueEstimators, ValueEstimatorBase, or type) –
The value estimator to use. This can be one of the following:
A
ValueEstimatorsenum type indicating which value function to use. If none is provided, the default stored in thedefault_value_estimatorattribute will be used.A
ValueEstimatorBaseinstance, which will be used directly as the value estimator.A
ValueEstimatorBasesubclass, which will be instantiated with the providedhyperparams.
The resulting value estimator class will be registered in
self.value_type, allowing future refinements.**hyperparams – hyperparameters to use for the value function. If not provided, the value indicated by
default_value_kwargs()will be used. When passing aValueEstimatorBasesubclass, these hyperparameters are passed directly to the class constructor.
- Returns:
Returns the loss module for method chaining.
- Return type:
self
Examples
>>> from torchrl.objectives import DQNLoss >>> # initialize the DQN loss >>> actor = torch.nn.Linear(3, 4) >>> dqn_loss = DQNLoss(actor, action_space="one-hot") >>> # updating the parameters of the default value estimator >>> dqn_loss.make_value_estimator(gamma=0.9) >>> dqn_loss.make_value_estimator( ... ValueEstimators.TD1, ... gamma=0.9) >>> # if we want to change the gamma value >>> dqn_loss.make_value_estimator(dqn_loss.value_type, gamma=0.9)
Using a
ValueEstimatorBasesubclass:>>> from torchrl.objectives.value import TD0Estimator >>> dqn_loss.make_value_estimator(TD0Estimator, gamma=0.99, value_network=value_net)
Using a
ValueEstimatorBaseinstance:>>> from torchrl.objectives.value import GAE >>> gae = GAE(gamma=0.99, lmbda=0.95, value_network=value_net) >>> ppo_loss.make_value_estimator(gae)
- policy_operator(policy: TensorDictModuleBase, *, method: str | None = None, method_kwargs: dict | None = None) TensorDictModuleBase[source]#
Bind a policy operation while retaining its observation encoders.