Source code for torchrl.trainers.algorithms.configs.transforms
# Copyright (c) Meta Platforms, Inc. and affiliates.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
from omegaconf import MISSING
from torchrl.envs.transforms import DoneTransform, ExpandAs, LastAction, RewardSum
from torchrl.envs.utils import ExplorationType
from torchrl.trainers.algorithms.configs.common import (
_normalize_hydra_key,
_normalize_hydra_keys,
ConfigBase,
)
[docs]
@dataclass
class NoopResetEnvConfig(TransformConfig):
"""Configuration for NoopResetEnv transform."""
noops: int = 30
random: bool = True
_target_: str = "torchrl.envs.transforms.transforms.NoopResetEnv"
def __post_init__(self) -> None:
"""Post-initialization hook for NoopResetEnv configuration."""
super().__post_init__()
[docs]
@dataclass
class StepCounterConfig(TransformConfig):
"""Configuration for StepCounter transform."""
max_steps: int | None = None
truncated_key: str | None = "truncated"
step_count_key: str | None = "step_count"
update_done: bool = True
_target_: str = "torchrl.envs.transforms.transforms.StepCounter"
def __post_init__(self) -> None:
"""Post-initialization hook for StepCounter configuration."""
super().__post_init__()
[docs]
@dataclass
class ComposeConfig(TransformConfig):
"""Configuration for Compose transform."""
transforms: list[Any] | None = None
_target_: str = "torchrl.envs.transforms.transforms.Compose"
def __post_init__(self) -> None:
"""Post-initialization hook for Compose configuration."""
super().__post_init__()
if self.transforms is None:
self.transforms = []
[docs]
@dataclass
class DoubleToFloatConfig(TransformConfig):
"""Configuration for DoubleToFloat transform."""
in_keys: list[str] | None = None
out_keys: list[str] | None = None
in_keys_inv: list[str] | None = None
out_keys_inv: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.DoubleToFloat"
def __post_init__(self) -> None:
"""Post-initialization hook for DoubleToFloat configuration."""
super().__post_init__()
[docs]
@dataclass
class ToTensorImageConfig(TransformConfig):
"""Configuration for ToTensorImage transform."""
from_int: bool | None = None
unsqueeze: bool = False
dtype: str | None = None
in_keys: list[str] | None = None
out_keys: list[str] | None = None
shape_tolerant: bool = False
_target_: str = "torchrl.envs.transforms.transforms.ToTensorImage"
def __post_init__(self) -> None:
"""Post-initialization hook for ToTensorImage configuration."""
super().__post_init__()
[docs]
@dataclass
class ResizeConfig(TransformConfig):
"""Configuration for Resize transform."""
w: int = 84
h: int = 84
interpolation: str = "bilinear"
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.Resize"
def __post_init__(self) -> None:
"""Post-initialization hook for Resize configuration."""
super().__post_init__()
[docs]
@dataclass
class CenterCropConfig(TransformConfig):
"""Configuration for CenterCrop transform."""
height: int = 84
width: int = 84
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.CenterCrop"
def __post_init__(self) -> None:
"""Post-initialization hook for CenterCrop configuration."""
super().__post_init__()
[docs]
@dataclass
class FlattenObservationConfig(TransformConfig):
"""Configuration for FlattenObservation transform."""
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.FlattenObservation"
def __post_init__(self) -> None:
"""Post-initialization hook for FlattenObservation configuration."""
super().__post_init__()
[docs]
@dataclass
class GrayScaleConfig(TransformConfig):
"""Configuration for GrayScale transform."""
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.GrayScale"
def __post_init__(self) -> None:
"""Post-initialization hook for GrayScale configuration."""
super().__post_init__()
[docs]
@dataclass
class ObservationNormConfig(TransformConfig):
"""Configuration for ObservationNorm transform."""
loc: float = 0.0
scale: float = 1.0
in_keys: list[str] | None = None
out_keys: list[str] | None = None
standard_normal: bool = False
eps: float = 1e-8
_target_: str = "torchrl.envs.transforms.transforms.ObservationNorm"
def __post_init__(self) -> None:
"""Post-initialization hook for ObservationNorm configuration."""
super().__post_init__()
[docs]
@dataclass
class CatFramesConfig(TransformConfig):
"""Configuration for CatFrames transform."""
N: int = 4
dim: int = -3
in_keys: list[str] | None = None
out_keys: list[str] | None = None
padding: str = "same"
padding_value: float = 0.0
as_inverse: bool = False
reset_key: str | None = None
done_key: str | None = None
future: bool = False
mask_key: str | None = None
_target_: str = "torchrl.envs.transforms.transforms.CatFrames"
def __post_init__(self) -> None:
"""Post-initialization hook for CatFrames configuration."""
super().__post_init__()
[docs]
@dataclass
class RewardClippingConfig(TransformConfig):
"""Configuration for RewardClipping transform."""
clamp_min: float | None = None
clamp_max: float | None = None
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.RewardClipping"
def __post_init__(self) -> None:
"""Post-initialization hook for RewardClipping configuration."""
super().__post_init__()
[docs]
@dataclass
class RewardScalingConfig(TransformConfig):
"""Configuration for RewardScaling transform."""
loc: float = 0.0
scale: float = 1.0
in_keys: list[str] | None = None
out_keys: list[str] | None = None
standard_normal: bool = False
eps: float = 1e-8
_target_: str = "torchrl.envs.transforms.transforms.RewardScaling"
def __post_init__(self) -> None:
"""Post-initialization hook for RewardScaling configuration."""
super().__post_init__()
[docs]
@dataclass
class VecNormConfig(TransformConfig):
"""Configuration for VecNorm transform."""
in_keys: list[str] | None = None
out_keys: list[str] | None = None
decay: float = 0.99
eps: float = 1e-8
_target_: str = "torchrl.envs.transforms.transforms.VecNorm"
def __post_init__(self) -> None:
"""Post-initialization hook for VecNorm configuration."""
super().__post_init__()
[docs]
@dataclass
class TargetReturnConfig(TransformConfig):
"""Configuration for TargetReturn transform."""
target_return: float = 10.0
mode: str = "reduce"
in_keys: list[str] | None = None
out_keys: list[str] | None = None
reset_key: str | None = None
_target_: str = "torchrl.envs.transforms.transforms.TargetReturn"
def __post_init__(self) -> None:
"""Post-initialization hook for TargetReturn configuration."""
super().__post_init__()
[docs]
@dataclass
class BinarizeRewardConfig(TransformConfig):
"""Configuration for BinarizeReward transform."""
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.BinarizeReward"
def __post_init__(self) -> None:
"""Post-initialization hook for BinarizeReward configuration."""
super().__post_init__()
[docs]
@dataclass
class ActionDiscretizerConfig(TransformConfig):
"""Configuration for ActionDiscretizer transform."""
num_intervals: int = 10
action_key: str = "action"
out_action_key: str | None = None
sampling: str | None = None
categorical: bool = True
_target_: str = "torchrl.envs.transforms.transforms.ActionDiscretizer"
def __post_init__(self) -> None:
"""Post-initialization hook for ActionDiscretizer configuration."""
super().__post_init__()
[docs]
@dataclass
class CatTensorsConfig(TransformConfig):
"""Configuration for CatTensors transform."""
dim: int = -1
in_keys: list[str] | None = None
out_key: Any = "observation_vector"
_target_: str = "torchrl.envs.transforms.transforms.CatTensors"
def __post_init__(self) -> None:
"""Post-initialization hook for CatTensors configuration."""
super().__post_init__()
[docs]
@dataclass
class StackConfig(TransformConfig):
"""Configuration for Stack transform."""
dim: int = 0
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.Stack"
def __post_init__(self) -> None:
"""Post-initialization hook for Stack configuration."""
super().__post_init__()
[docs]
@dataclass
class DiscreteActionProjectionConfig(TransformConfig):
"""Configuration for DiscreteActionProjection transform."""
num_actions: int = 4
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.DiscreteActionProjection"
def __post_init__(self) -> None:
"""Post-initialization hook for DiscreteActionProjection configuration."""
super().__post_init__()
[docs]
@dataclass
class TensorDictPrimerConfig(TransformConfig):
"""Configuration for TensorDictPrimer transform."""
primer_spec: Any = None
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.TensorDictPrimer"
def __post_init__(self) -> None:
"""Post-initialization hook for TensorDictPrimer configuration."""
super().__post_init__()
[docs]
@dataclass
class RewardSumConfig(TransformConfig):
"""Configuration for RewardSum transform."""
in_keys: list[str] | None = None
out_keys: list[str] | None = None
reset_keys: list[str] | None = None
_target_: str = (
"torchrl.trainers.algorithms.configs.transforms._make_reward_sum_transform"
)
def __post_init__(self) -> None:
"""Post-initialization hook for RewardSum configuration."""
super().__post_init__()
[docs]
@dataclass
class TimeMaxPoolConfig(TransformConfig):
"""Configuration for TimeMaxPool transform."""
dim: int = -1
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.TimeMaxPool"
def __post_init__(self) -> None:
"""Post-initialization hook for TimeMaxPool configuration."""
super().__post_init__()
[docs]
@dataclass
class RandomCropTensorDictConfig(TransformConfig):
"""Configuration for RandomCropTensorDict transform."""
crop_size: list[int] | None = None
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.RandomCropTensorDict"
def __post_init__(self) -> None:
"""Post-initialization hook for RandomCropTensorDict configuration."""
super().__post_init__()
if self.crop_size is None:
self.crop_size = [84, 84]
[docs]
@dataclass
class InitTrackerConfig(TransformConfig):
"""Configuration for InitTracker transform."""
init_key: str = "is_init"
_target_: str = "torchrl.envs.transforms.transforms.InitTracker"
def __post_init__(self) -> None:
"""Post-initialization hook for InitTracker configuration."""
super().__post_init__()
[docs]
@dataclass
class LastActionConfig(TransformConfig):
"""Hydra configuration for :class:`~torchrl.envs.transforms.LastAction`."""
in_keys: list[Any] | None = None
out_keys: list[Any] | None = None
default: Any = "zeros"
reset_key: Any | None = None
_target_: str = (
"torchrl.trainers.algorithms.configs.transforms._make_last_action_transform"
)
def __post_init__(self) -> None:
"""Post-initialization hook for LastAction configuration."""
super().__post_init__()
[docs]
@dataclass
class ActionMaskConfig(TransformConfig):
"""Configuration for ActionMask transform."""
mask_key: str = "action_mask"
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.ActionMask"
def __post_init__(self) -> None:
"""Post-initialization hook for ActionMask configuration."""
super().__post_init__()
[docs]
@dataclass
class RemoveEmptySpecsConfig(TransformConfig):
"""Configuration for RemoveEmptySpecs transform."""
_target_: str = "torchrl.envs.transforms.transforms.RemoveEmptySpecs"
def __post_init__(self) -> None:
"""Post-initialization hook for RemoveEmptySpecs configuration."""
super().__post_init__()
[docs]
@dataclass
class TrajCounterConfig(TransformConfig):
"""Configuration for TrajCounter transform."""
out_key: str = "traj_count"
repeats: int | None = None
_target_: str = "torchrl.envs.transforms.transforms.TrajCounter"
def __post_init__(self) -> None:
"""Post-initialization hook for TrajCounter configuration."""
super().__post_init__()
[docs]
@dataclass
class LineariseRewardsConfig(TransformConfig):
"""Configuration for LineariseRewards transform."""
in_keys: list[str] | None = None
out_keys: list[str] | None = None
weights: list[float] | None = None
_target_: str = "torchrl.envs.transforms.transforms.LineariseRewards"
def __post_init__(self) -> None:
"""Post-initialization hook for LineariseRewards configuration."""
super().__post_init__()
if self.in_keys is None:
self.in_keys = []
[docs]
@dataclass
class ConditionalSkipConfig(TransformConfig):
"""Configuration for ConditionalSkip transform."""
cond: Any = None
_target_: str = "torchrl.envs.transforms.transforms.ConditionalSkip"
def __post_init__(self) -> None:
"""Post-initialization hook for ConditionalSkip configuration."""
super().__post_init__()
[docs]
@dataclass
class MultiActionConfig(TransformConfig):
"""Hydra configuration for :class:`~torchrl.envs.transforms.MultiAction`."""
dim: int = 1
stack_rewards: bool = True
stack_observations: bool = False
action_key: Any | None = None
chunk_key: Any | None = None
reward_aggregation: str | None = None
_target_: str = "torchrl.envs.transforms.transforms.MultiAction"
def __post_init__(self) -> None:
"""Post-initialization hook for MultiAction configuration."""
super().__post_init__()
[docs]
@dataclass
class ClosedLoopMultiActionConfig(TransformConfig):
"""Hydra configuration for :class:`~torchrl.envs.transforms.ClosedLoopMultiAction`.
Install controller primers first, or override the target with
ClosedLoopMultiAction.from_env and supply the environment to instantiate.
"""
controller: Any = MISSING
steps: int = MISSING
decision_spec: Any = None
reward_aggregation: str = "sum"
exploration_type: ExplorationType = ExplorationType.DETERMINISTIC
no_grad: bool = True
dim: int = 1
stack_observations: bool = False
_target_: str = "torchrl.envs.transforms.ClosedLoopMultiAction"
[docs]
@dataclass
class TimerConfig(TransformConfig):
"""Configuration for Timer transform."""
out_keys: list[str] | None = None
time_key: str = "time"
_target_: str = "torchrl.envs.transforms.transforms.Timer"
def __post_init__(self) -> None:
"""Post-initialization hook for Timer configuration."""
super().__post_init__()
[docs]
@dataclass
class ConditionalPolicySwitchConfig(TransformConfig):
"""Configuration for ConditionalPolicySwitch transform."""
policy: Any = None
condition: Any = None
_target_: str = "torchrl.envs.transforms.transforms.ConditionalPolicySwitch"
def __post_init__(self) -> None:
"""Post-initialization hook for ConditionalPolicySwitch configuration."""
super().__post_init__()
[docs]
@dataclass
class VecNormV2Config(TransformConfig):
"""Configuration for VecNormV2 transform."""
in_keys: list[str] | None = None
out_keys: list[str] | None = None
decay: float = 0.99
eps: float = 1e-8
_target_: str = "torchrl.envs.transforms.vecnorm.VecNormV2"
def __post_init__(self) -> None:
"""Post-initialization hook for VecNormV2 configuration."""
super().__post_init__()
[docs]
@dataclass
class FiniteTensorDictCheckConfig(TransformConfig):
"""Configuration for FiniteTensorDictCheck transform."""
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.FiniteTensorDictCheck"
def __post_init__(self) -> None:
"""Post-initialization hook for FiniteTensorDictCheck configuration."""
super().__post_init__()
[docs]
@dataclass
class HashConfig(TransformConfig):
"""Configuration for Hash transform."""
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.Hash"
def __post_init__(self) -> None:
"""Post-initialization hook for Hash configuration."""
super().__post_init__()
[docs]
@dataclass
class TokenizerConfig(TransformConfig):
"""Configuration for Tokenizer transform."""
vocab_size: int = 1000
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.Tokenizer"
def __post_init__(self) -> None:
"""Post-initialization hook for Tokenizer configuration."""
super().__post_init__()
[docs]
@dataclass
class CropConfig(TransformConfig):
"""Configuration for Crop transform."""
top: int = 0
left: int = 0
height: int = 84
width: int = 84
in_keys: list[str] | None = None
out_keys: list[str] | None = None
_target_: str = "torchrl.envs.transforms.transforms.Crop"
def __post_init__(self) -> None:
"""Post-initialization hook for Crop configuration."""
super().__post_init__()
@dataclass
class FlattenTensorDictConfig(TransformConfig):
"""Configuration for flattening TensorDict during inverse pass.
This transform reshapes the tensordict to have a flat batch dimension
during the inverse pass, which is useful for replay buffers that need
to store data with a flat batch structure.
"""
_target_: str = "torchrl.envs.transforms.transforms.FlattenTensorDict"
def __post_init__(self) -> None:
"""Post-initialization hook for FlattenTensorDict configuration."""
super().__post_init__()
@dataclass
class ModuleTransformConfig(TransformConfig):
"""Configuration for ModuleTransform."""
module: Any = None
device: Any = None
no_grad: bool = False
inverse: bool = False
_target_: str = "torchrl.envs.transforms.module.ModuleTransform"
_partial_: bool = False
def __post_init__(self) -> None:
"""Post-initialization hook for ModuleTransform configuration."""
super().__post_init__()
@dataclass
class ExpandAsConfig(TransformConfig):
"""Configuration for ExpandAs transform."""
ref_key: list[str] | None = None
in_key: list[str] | None = None
out_key: list[str] | None = None
_target_: str = (
"torchrl.trainers.algorithms.configs.transforms._make_expand_as_transform"
)
def __post_init__(self) -> None:
super().__post_init__()
def _make_last_action_transform(*args, **kwargs) -> LastAction:
in_keys = _normalize_hydra_keys(kwargs.pop("in_keys", None))
out_keys = _normalize_hydra_keys(kwargs.pop("out_keys", None))
reset_key = _normalize_hydra_key(kwargs.pop("reset_key", None))
return LastAction(
in_keys=in_keys,
out_keys=out_keys,
reset_key=reset_key,
**kwargs,
)
def _make_reward_sum_transform(*args, **kwargs) -> RewardSum:
in_keys = _normalize_hydra_keys(kwargs.pop("in_keys", None))
out_keys = _normalize_hydra_keys(kwargs.pop("out_keys", None))
reset_keys = _normalize_hydra_keys(kwargs.pop("reset_keys", None))
return RewardSum(in_keys=in_keys, out_keys=out_keys, reset_keys=reset_keys)
def _make_expand_as_transform(*args, **kwargs) -> ExpandAs:
ref_key = _normalize_hydra_key(kwargs.pop("ref_key", None))
in_key = _normalize_hydra_key(kwargs.pop("in_key", None))
out_key = _normalize_hydra_key(kwargs.pop("out_key", None))
return ExpandAs(ref_key=ref_key, in_key=in_key, out_key=out_key)
def _make_done_transform(*args, **kwargs) -> DoneTransform:
in_keys = _normalize_hydra_keys(kwargs.pop("in_keys", None))
out_keys = _normalize_hydra_keys(kwargs.pop("out_keys", None))
done_keys = _normalize_hydra_keys(kwargs.pop("done_keys", None))
reward_key = kwargs.pop("reward_key", None)
if reward_key is not None:
reward_key = _normalize_hydra_key(reward_key)
return DoneTransform(
in_keys=in_keys,
out_keys=out_keys,
reward_key=reward_key,
done_keys=done_keys,
)