Skip to content

Commit c91e44b

Browse files
committed
Relax some state constraints
1 parent 4957051 commit c91e44b

6 files changed

Lines changed: 28 additions & 28 deletions

File tree

lerax/policy/actor_critic/mlp.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -41,9 +41,9 @@ class MLPActorCriticPolicy[
4141
action_model: MLP
4242
log_std: Float[Array, " action_size"]
4343

44-
def __init__(
44+
def __init__[StateType: AbstractEnvLikeState](
4545
self,
46-
env: AbstractEnvLike[AbstractEnvLikeState, ActType, ObsType],
46+
env: AbstractEnvLike[StateType, ActType, ObsType],
4747
*,
4848
feature_size: int = 64,
4949
feature_width: int = 64,

lerax/wrapper/base_wrapper.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@ class AbstractWrapper[
2222
WrapperStateType: AbstractEnvLikeState,
2323
WrapperActType,
2424
WrapperObsType,
25-
StateType: AbstractEnvState,
25+
StateType: AbstractEnvLikeState,
2626
ActType,
2727
ObsType,
2828
](AbstractEnvLike[WrapperStateType, WrapperActType, WrapperObsType]):
@@ -31,7 +31,7 @@ class AbstractWrapper[
3131
env: eqx.AbstractVar[AbstractEnvLike[StateType, ActType, ObsType]]
3232

3333
@property
34-
def unwrapped(self) -> AbstractEnv[StateType, ActType, ObsType]:
34+
def unwrapped(self) -> AbstractEnv:
3535
"""Return the unwrapped environment"""
3636
return self.env.unwrapped
3737

lerax/wrapper/misc.py

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
from jax import numpy as jnp
66
from jaxtyping import Array, ArrayLike, Bool, Float, Int, Key
77

8-
from lerax.env import AbstractEnvLike, AbstractEnvLikeState, AbstractEnvState
8+
from lerax.env import AbstractEnvLike, AbstractEnvLikeState
99
from lerax.space import AbstractSpace
1010

1111
from .base_wrapper import (
@@ -14,7 +14,7 @@
1414
)
1515

1616

17-
class Identity[StateType: AbstractEnvState, ActType, ObsType](
17+
class Identity[StateType: AbstractEnvLikeState, ActType, ObsType](
1818
AbstractWrapper[StateType, ActType, ObsType, StateType, ActType, ObsType]
1919
):
2020
env: AbstractEnvLike[StateType, ActType, ObsType]
@@ -33,7 +33,7 @@ def step(
3333
return self.env.step(state, action, key=key)
3434

3535

36-
class EpisodeStatisticsState[StateType: AbstractEnvState](AbstractWrapperState):
36+
class EpisodeStatisticsState[StateType: AbstractEnvLikeState](AbstractWrapperState):
3737
env_state: StateType
3838

3939
episode_length: Int[Array, ""]
@@ -76,7 +76,7 @@ def info(self) -> dict:
7676
}
7777

7878

79-
class EpisodeStatistics[StateType: AbstractEnvState, ActType, ObsType](
79+
class EpisodeStatistics[StateType: AbstractEnvLikeState, ActType, ObsType](
8080
AbstractWrapper[
8181
EpisodeStatisticsState[StateType], ActType, ObsType, StateType, ActType, ObsType
8282
]
@@ -137,7 +137,7 @@ def __init__(self, step_count: Int[ArrayLike, ""], env_state: StateType):
137137
self.env_state = env_state
138138

139139

140-
class TimeLimit[StateType: AbstractEnvState, ActType, ObsType](
140+
class TimeLimit[StateType: AbstractEnvLikeState, ActType, ObsType](
141141
AbstractWrapper[
142142
TimeLimitState[StateType], ActType, ObsType, StateType, ActType, ObsType
143143
]
@@ -193,7 +193,7 @@ def close(self):
193193
self.env.close()
194194

195195

196-
class AutoClose[StateType: AbstractEnvState, ActType, ObsType]():
196+
class AutoClose[StateType: AbstractEnvLikeState, ActType, ObsType]():
197197
"""
198198
Closes the environment automatically when it is deleted.
199199

lerax/wrapper/transform_action.py

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -6,15 +6,15 @@
66
from jax import numpy as jnp
77
from jaxtyping import Array, Float, Key
88

9-
from lerax.env import AbstractEnvLike, AbstractEnvState
9+
from lerax.env import AbstractEnvLike, AbstractEnvLikeState
1010
from lerax.space import AbstractSpace, Box
1111

1212
from .base_wrapper import AbstractWrapper
1313
from .utils import rescale_box
1414

1515

1616
class AbstractPureTransformActionWrapper[
17-
WrapperActType, StateType: AbstractEnvState, ActType, ObsType
17+
WrapperActType, StateType: AbstractEnvLikeState, ActType, ObsType
1818
](
1919
AbstractWrapper[
2020
StateType,
@@ -47,9 +47,9 @@ def close(self):
4747
self.env.close()
4848

4949

50-
class TransformAction[WrapperActType, StateType: AbstractEnvState, ActType, ObsType](
51-
AbstractPureTransformActionWrapper[WrapperActType, StateType, ActType, ObsType]
52-
):
50+
class TransformAction[
51+
WrapperActType, StateType: AbstractEnvLikeState, ActType, ObsType
52+
](AbstractPureTransformActionWrapper[WrapperActType, StateType, ActType, ObsType]):
5353
"""Apply a function to the action before passing it to the environment"""
5454

5555
env: AbstractEnvLike[StateType, ActType, ObsType]
@@ -67,7 +67,7 @@ def __init__(
6767
self.action_space = action_space
6868

6969

70-
class ClipAction[StateType: AbstractEnvState, ObsType](
70+
class ClipAction[StateType: AbstractEnvLikeState, ObsType](
7171
AbstractPureTransformActionWrapper[
7272
Float[Array, " ..."], StateType, Float[Array, " ..."], ObsType
7373
],
@@ -98,7 +98,7 @@ def clip(action: Float[Array, " ..."]) -> Float[Array, " ..."]:
9898
self.action_space = action_space
9999

100100

101-
class RescaleAction[StateType: AbstractEnvState, ObsType](
101+
class RescaleAction[StateType: AbstractEnvLikeState, ObsType](
102102
AbstractPureTransformActionWrapper[
103103
Float[Array, " ..."], StateType, Float[Array, " ..."], ObsType
104104
],

lerax/wrapper/transform_observation.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -7,15 +7,15 @@
77
from jax import numpy as jnp
88
from jaxtyping import Array, Float, Key
99

10-
from lerax.env import AbstractEnvLike, AbstractEnvState
10+
from lerax.env import AbstractEnvLike, AbstractEnvLikeState
1111
from lerax.space import AbstractSpace, Box
1212

1313
from .base_wrapper import AbstractWrapper
1414
from .utils import rescale_box
1515

1616

1717
class AbstractPureObservationWrapper[
18-
WrapperObsType, StateType: AbstractEnvState, ActType, ObsType
18+
WrapperObsType, StateType: AbstractEnvLikeState, ActType, ObsType
1919
](AbstractWrapper[StateType, ActType, WrapperObsType, StateType, ActType, ObsType]):
2020
"""
2121
Apply a pure function to every observation that leaves the environment.
@@ -42,7 +42,7 @@ def close(self):
4242
self.env.close()
4343

4444

45-
class ClipObservation[StateType: AbstractEnvState](
45+
class ClipObservation[StateType: AbstractEnvLikeState](
4646
AbstractPureObservationWrapper[
4747
Float[Array, " ..."], StateType, Float[Array, " ..."], Float[Array, " ..."]
4848
],
@@ -71,7 +71,7 @@ def __init__(self, env: AbstractEnvLike):
7171
self.observation_space = env.observation_space
7272

7373

74-
class RescaleObservation[StateType: AbstractEnvState](
74+
class RescaleObservation[StateType: AbstractEnvLikeState](
7575
AbstractPureObservationWrapper[
7676
Float[Array, " ..."], StateType, Float[Array, " ..."], Float[Array, " ..."]
7777
],
@@ -101,7 +101,7 @@ def __init__(
101101
self.observation_space = new_box
102102

103103

104-
class FlattenObservation[StateType: AbstractEnvState, ObsType](
104+
class FlattenObservation[StateType: AbstractEnvLikeState, ObsType](
105105
AbstractPureObservationWrapper[
106106
Float[Array, " flat"], StateType, Float[Array, " ..."], ObsType
107107
]

lerax/wrapper/transform_reward.py

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -7,14 +7,14 @@
77
from jax import numpy as jnp
88
from jaxtyping import Array, ArrayLike, Float, Key
99

10-
from lerax.env import AbstractEnvLike, AbstractEnvState
10+
from lerax.env import AbstractEnvLike, AbstractEnvLikeState
1111

1212
from .base_wrapper import AbstractWrapper
1313

1414

15-
class AbstractPureTransformRewardWrapper[StateType: AbstractEnvState, ActType, ObsType](
16-
AbstractWrapper[StateType, ActType, ObsType, StateType, ActType, ObsType]
17-
):
15+
class AbstractPureTransformRewardWrapper[
16+
StateType: AbstractEnvLikeState, ActType, ObsType
17+
](AbstractWrapper[StateType, ActType, ObsType, StateType, ActType, ObsType]):
1818
"""
1919
Apply a *pure* (stateless) function to every reward emitted by the wrapped
2020
environment.
@@ -39,11 +39,11 @@ def close(self):
3939
self.env.close()
4040

4141

42-
class ClipReward[StateType: AbstractEnvState, ActType, ObsType](
42+
class ClipReward[StateType: AbstractEnvLikeState, ActType, ObsType](
4343
AbstractPureTransformRewardWrapper[StateType, ActType, ObsType]
4444
):
4545
"""
46-
Element-wise clip of rewards: `reward clamp(min, max)`.
46+
Element-wise clip of rewards: `reward -> clamp(min, max)`.
4747
"""
4848

4949
env: AbstractEnvLike[StateType, ActType, ObsType]

0 commit comments

Comments
 (0)