Source code for synora.envs.wrappers

from __future__ import annotations

import datetime
from collections import deque
import uuid
from typing import Any

from synora.envs._actions import clip_box_action
from synora.envs._contract import finalize_step_info
from synora.utils.gym_compat import gym
import numpy as np
from PIL import Image


[docs] class TimeLimit: """Terminate episodes after a fixed number of wrapper steps. If the wrapped environment does not provide a discount flag at timeout, the wrapper injects a default discount of `1.0` for downstream learners. """ def __init__(self, env: Any, duration: int) -> None: self._env = env self._duration = duration self._step: int | None = None def __getattr__(self, name: str) -> Any: return getattr(self._env, name)
[docs] def step(self, action: Any) -> tuple[Any, Any, bool, dict[str, Any]]: assert self._step is not None, "Must reset environment." obs, reward, done, info = self._env.step(action) self._step += 1 timeout = self._step >= self._duration upstream_terminated = bool(info.get("terminated", False)) if info else False if timeout: done = True self._step = None info = finalize_step_info( info, done=done, terminated=upstream_terminated, truncated=timeout and not upstream_terminated, discount=np.array( 1.0 if timeout and not upstream_terminated else ((info or {}).get("discount", 0.0 if done else 1.0)), dtype=np.float32, ), ) return obs, reward, done, info
[docs] def reset(self, *, seed: int | None = None) -> Any: self._step = 0 if seed is None: return self._env.reset() return self._env.reset(seed=seed)
[docs] class FrameStack: """Stack the most recent image observations along the channel axis. The wrapped environment must emit dict observations with an ``"image"`` key in ``(C, H, W)`` layout. Non-image keys such as ``"state"`` are passed through unchanged from the latest observation. """ def __init__(self, env: Any, num_frames: int) -> None: if int(num_frames) < 1: raise ValueError("num_frames must be >= 1") self._env = env self._num_frames = int(num_frames) self._frames: deque[np.ndarray] = deque(maxlen=self._num_frames) image_space = env.observation_space["image"] if len(image_space.shape) != 3: raise ValueError( "FrameStack expects image observations with shape (C, H, W)." ) channels, height, width = image_space.shape spaces = dict(env.observation_space.spaces) image_space_cls = image_space.__class__ spaces["image"] = image_space_cls( low=0, high=255, shape=(channels * self._num_frames, height, width), dtype=image_space.dtype, ) dict_space_cls = env.observation_space.__class__ self._observation_space = dict_space_cls(spaces) def __getattr__(self, name: str) -> Any: return getattr(self._env, name) @property def observation_space(self) -> gym.spaces.Dict: return self._observation_space @property def action_space(self) -> Any: return self._env.action_space def _stack_observation(self, obs: dict[str, Any]) -> dict[str, Any]: image = np.asarray(obs["image"], dtype=np.uint8) if image.ndim != 3: raise ValueError( "FrameStack expects image observations with shape (C, H, W)." ) stacked = np.concatenate(list(self._frames), axis=0) out = dict(obs) out["image"] = stacked.copy() return out
[docs] def reset(self, *, seed: int | None = None) -> dict[str, Any]: obs = self._env.reset() if seed is None else self._env.reset(seed=seed) image = np.asarray(obs["image"], dtype=np.uint8) self._frames.clear() for _ in range(self._num_frames): self._frames.append(image.copy()) return self._stack_observation(obs)
[docs] def step(self, action: Any) -> tuple[dict[str, Any], Any, bool, dict[str, Any]]: obs, reward, done, info = self._env.step(action) image = np.asarray(obs["image"], dtype=np.uint8) self._frames.append(image.copy()) return self._stack_observation(obs), reward, done, info
[docs] class ActionRepeat: """Repeat each action for a fixed number of environment steps. Rewards are accumulated and the loop stops early if the environment terminates, mirroring common action-repeat behavior in world model papers. """ def __init__(self, env: Any, amount: int) -> None: self._env = env self._amount = amount def __getattr__(self, name: str) -> Any: return getattr(self._env, name)
[docs] def step(self, action: Any) -> tuple[Any, float, bool, dict[str, Any]]: done = False total_reward = 0 current_step = 0 info: dict[str, Any] = {} while current_step < self._amount and not done: obs, reward, done, info = self._env.step(action) total_reward += reward current_step += 1 info = dict(info or {}) info["action_repeat"] = current_step return obs, total_reward, done, info
[docs] class NormalizeActions: """Expose a normalized `[-1, 1]` action space for bounded continuous controls. Incoming normalized actions are mapped back to the wrapped environment action bounds before stepping the environment. """ def __init__(self, env: Any) -> None: self._env = env self._mask = np.logical_and( np.isfinite(env.action_space.low), np.isfinite(env.action_space.high) ) self._low = np.where(self._mask, env.action_space.low, -1) self._high = np.where(self._mask, env.action_space.high, 1) def __getattr__(self, name: str) -> Any: return getattr(self._env, name) @property def action_space(self) -> gym.spaces.Box: low = np.where(self._mask, -np.ones_like(self._low), self._low) high = np.where(self._mask, np.ones_like(self._low), self._high) return gym.spaces.Box(low, high, dtype=np.float32)
[docs] def step(self, action: np.ndarray) -> tuple[Any, Any, bool, dict[str, Any]]: normalized = clip_box_action( action, -np.ones_like(self._low, dtype=np.float32), np.ones_like(self._high, dtype=np.float32), ) original = (normalized + 1.0) / 2.0 * (self._high - self._low) + self._low original = np.where(self._mask, original, normalized).astype( np.float32, copy=False ) obs, reward, done, info = self._env.step(original) info = dict(info or {}) if "action" in info and "executed_action" not in info: existing = info["action"] info["executed_action"] = ( int(np.asarray(existing).item()) if np.isscalar(existing) else np.asarray(existing).copy() ) info["action"] = normalized.copy() return obs, reward, done, info
[docs] class ObsDict: """Convert scalar/array observations into a dictionary observation format. This harmonizes outputs for code paths that expect keyed observations (for example `{"image": ...}` style world model inputs). """ def __init__(self, env: Any, key: str = "obs") -> None: self._env = env self._key = key def __getattr__(self, name: str) -> Any: return getattr(self._env, name) @property def observation_space(self) -> gym.spaces.Dict: spaces = {self._key: self._env.observation_space} return gym.spaces.Dict(spaces) @property def action_space(self) -> Any: return self._env.action_space
[docs] def step(self, action: Any) -> tuple[dict[str, Any], Any, bool, dict[str, Any]]: obs, reward, done, info = self._env.step(action) obs = {self._key: np.array(obs)} return obs, reward, done, info
[docs] def reset(self, *, seed: int | None = None) -> dict[str, Any]: result = self._env.reset() if seed is None else self._env.reset(seed=seed) obs = result[0] if isinstance(result, tuple) else result obs = {self._key: np.array(obs)} return obs
[docs] class OneHotAction: """Wrap discrete-action environments to accept one-hot action vectors. The wrapper validates one-hot inputs and converts them to integer action indices before forwarding to the underlying environment. """ def __init__(self, env: Any) -> None: assert isinstance(env.action_space, gym.spaces.Discrete) self._env = env self._random = np.random.default_rng() shape = (self._env.action_space.n,) self._action_space = gym.spaces.Box( low=0, high=1, shape=shape, dtype=np.float32 ) self._action_space.sample = self._sample_action def __getattr__(self, name: str) -> Any: return getattr(self._env, name) @property def action_space(self) -> gym.spaces.Box: return self._action_space
[docs] def step(self, action: np.ndarray) -> tuple[Any, Any, bool, dict[str, Any]]: index = np.argmax(action).astype(int) reference = np.zeros_like(action) reference[index] = 1 if not np.allclose(reference, action): raise ValueError(f"Invalid one-hot action:\n{action}") return self._env.step(index)
[docs] def reset(self, *, seed: int | None = None) -> Any: if seed is not None: self._random = np.random.default_rng(seed) if hasattr(self._action_space, "seed"): try: self._action_space.seed(seed) except Exception: pass return self._env.reset(seed=seed) return self._env.reset()
def _sample_action(self) -> np.ndarray: actions = self._env.action_space.n index = int(self._random.integers(0, actions)) reference = np.zeros(actions, dtype=np.float32) reference[index] = 1.0 return reference
[docs] class RewardObs: """Augment observations with the latest scalar reward under `obs["reward"]`. Useful for agents that consume reward as part of the observation stream during model learning or recurrent policy inference. """ def __init__(self, env: Any) -> None: self._env = env def __getattr__(self, name: str) -> Any: return getattr(self._env, name) @property def observation_space(self) -> gym.spaces.Dict: spaces = dict(self._env.observation_space.spaces) assert "reward" not in spaces spaces["reward"] = gym.spaces.Box(-np.inf, np.inf, dtype=np.float32) return gym.spaces.Dict(spaces)
[docs] def step(self, action: Any) -> tuple[dict[str, Any], Any, bool, dict[str, Any]]: obs, reward, done, info = self._env.step(action) obs["reward"] = reward return obs, reward, done, info
[docs] def reset(self, *, seed: int | None = None) -> dict[str, Any]: result = self._env.reset() if seed is None else self._env.reset(seed=seed) obs = result[0] if isinstance(result, tuple) else result obs["reward"] = 0.0 return obs
[docs] class ResizeImage: """Resize image-like observation entries to a target spatial size. The wrapper discovers image keys from `env.obs_space`, applies nearest neighbor resizing, and updates the advertised observation space shapes. """ def __init__(self, env: Any, size: tuple[int, int] = (64, 64)) -> None: self._env = env self._size = size self._keys = [ k for k, v in env.obs_space.items() if len(v.shape) > 1 and v.shape[:2] != size ] print(f"Resizing keys {','.join(self._keys)} to {self._size}.") if self._keys: self._Image = Image def __getattr__(self, name: str) -> Any: if name.startswith("__"): raise AttributeError(name) try: return getattr(self._env, name) except AttributeError: raise AttributeError(name) @property def obs_space(self) -> dict[str, Any]: spaces = self._env.obs_space for key in self._keys: shape = self._size + spaces[key].shape[2:] spaces[key] = gym.spaces.Box(0, 255, shape, np.uint8) return spaces
[docs] def step(self, action: Any) -> Any: obs = self._env.step(action) for key in self._keys: obs[key] = self._resize(obs[key]) return obs
[docs] def reset(self, *, seed: int | None = None) -> Any: obs = self._env.reset() if seed is None else self._env.reset(seed=seed) for key in self._keys: obs[key] = self._resize(obs[key]) return obs
def _resize(self, image: np.ndarray) -> np.ndarray: img = self._Image.fromarray(image) img = img.resize(self._size, self._Image.Resampling.NEAREST) return np.array(img)
[docs] class RenderImage: """Inject RGB renders from `env.render("rgb_array")` into observations. This is useful when the base environment returns non-image observations but a rendered camera view is needed for world-model training. """ def __init__(self, env: Any, key: str = "image") -> None: self._env = env self._key = key self._shape = self._env.render().shape def __getattr__(self, name: str) -> Any: if name.startswith("__"): raise AttributeError(name) try: return getattr(self._env, name) except AttributeError: raise AttributeError(name) @property def obs_space(self) -> dict[str, Any]: spaces = self._env.obs_space spaces[self._key] = gym.spaces.Box(0, 255, self._shape, np.uint8) return spaces
[docs] def step(self, action: Any) -> Any: obs = self._env.step(action) obs[self._key] = self._env.render("rgb_array") return obs
[docs] def reset(self, *, seed: int | None = None) -> Any: obs = self._env.reset() if seed is None else self._env.reset(seed=seed) obs[self._key] = self._env.render("rgb_array") return obs
[docs] class UUID(gym.Wrapper): """Gym wrapper that tracks a unique run identifier per environment reset. The ID combines timestamp and UUID and can be used to tag episodes or artifacts generated during data collection. """ def __init__(self, env: Any) -> None: super().__init__(env) timestamp = datetime.datetime.now().strftime("%Y%m%dT%H%M%S") self.id = f"{timestamp}-{str(uuid.uuid4().hex)}"
[docs] def reset(self, **kwargs: Any) -> Any: timestamp = datetime.datetime.now().strftime("%Y%m%dT%H%M%S") self.id = f"{timestamp}-{str(uuid.uuid4().hex)}" return self.env.reset(**kwargs)
[docs] class SelectAction(gym.Wrapper): """Gym wrapper for dictionary actions that forwards a selected key only. This enables integration with policies that emit action dicts while the environment expects a single tensor/array action payload. """ def __init__(self, env: Any, key: str) -> None: super().__init__(env) self._key = key
[docs] def step(self, action: dict[str, Any]) -> Any: return self.env.step(action[self._key])