diff --git a/src/lerobot/rl/buffer.py b/src/lerobot/rl/buffer.py index cec80b723..015868196 100644 --- a/src/lerobot/rl/buffer.py +++ b/src/lerobot/rl/buffer.py @@ -18,7 +18,7 @@ import functools import threading from collections.abc import Callable, Sequence from contextlib import suppress -from typing import TypedDict +from typing import NotRequired, TypedDict import torch import torch.nn.functional as F # noqa: N812 @@ -36,7 +36,7 @@ class BatchTransition(TypedDict): next_state: dict[str, torch.Tensor] done: torch.Tensor truncated: torch.Tensor - complementary_info: dict[str, torch.Tensor | float | int] | None = None + complementary_info: NotRequired[dict[str, torch.Tensor | float | int] | None] def random_crop_vectorized(images: torch.Tensor, output_size: tuple) -> torch.Tensor: diff --git a/src/lerobot/utils/transition.py b/src/lerobot/utils/transition.py index a79b95151..878114f72 100644 --- a/src/lerobot/utils/transition.py +++ b/src/lerobot/utils/transition.py @@ -14,7 +14,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -from typing import TypedDict +from typing import NotRequired, TypedDict import torch @@ -28,7 +28,7 @@ class Transition(TypedDict): next_state: dict[str, torch.Tensor] done: bool truncated: bool - complementary_info: dict[str, torch.Tensor | float | int] | None = None + complementary_info: NotRequired[dict[str, torch.Tensor | float | int] | None] def move_transition_to_device(transition: Transition, device: str = "cpu") -> Transition: