mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-07-29 20:14:17 +00:00
Downgrade sphinx-autodoc-typehints (#1291)
* Update setup.py * black * hotfix pytype
This commit is contained in:
parent
92f7a6f23b
commit
69fdf155e1
3 changed files with 10 additions and 10 deletions
2
setup.py
2
setup.py
|
|
@ -117,7 +117,7 @@ setup(
|
|||
# For spelling
|
||||
"sphinxcontrib.spelling",
|
||||
# Type hints support
|
||||
"sphinx-autodoc-typehints",
|
||||
"sphinx-autodoc-typehints==1.21.1", # TODO: remove version constraint, see #1290
|
||||
# Copy button for code snippets
|
||||
"sphinx_copybutton",
|
||||
],
|
||||
|
|
|
|||
|
|
@ -474,7 +474,7 @@ class RolloutBuffer(BaseBuffer):
|
|||
yield self._get_samples(indices[start_idx : start_idx + batch_size])
|
||||
start_idx += batch_size
|
||||
|
||||
def _get_samples(self, batch_inds: np.ndarray, env: Optional[VecNormalize] = None) -> RolloutBufferSamples:
|
||||
def _get_samples(self, batch_inds: np.ndarray, env: Optional[VecNormalize] = None) -> RolloutBufferSamples: # type: ignore[signature-mismatch] #FIXME
|
||||
data = (
|
||||
self.observations[batch_inds],
|
||||
self.actions[batch_inds],
|
||||
|
|
@ -603,7 +603,7 @@ class DictReplayBuffer(ReplayBuffer):
|
|||
self.full = True
|
||||
self.pos = 0
|
||||
|
||||
def sample(self, batch_size: int, env: Optional[VecNormalize] = None) -> DictReplayBufferSamples:
|
||||
def sample(self, batch_size: int, env: Optional[VecNormalize] = None) -> DictReplayBufferSamples: # type: ignore[signature-mismatch] #FIXME:
|
||||
"""
|
||||
Sample elements from the replay buffer.
|
||||
|
||||
|
|
@ -614,7 +614,7 @@ class DictReplayBuffer(ReplayBuffer):
|
|||
"""
|
||||
return super(ReplayBuffer, self).sample(batch_size=batch_size, env=env)
|
||||
|
||||
def _get_samples(self, batch_inds: np.ndarray, env: Optional[VecNormalize] = None) -> DictReplayBufferSamples:
|
||||
def _get_samples(self, batch_inds: np.ndarray, env: Optional[VecNormalize] = None) -> DictReplayBufferSamples: # type: ignore[signature-mismatch] #FIXME:
|
||||
# Sample randomly the env idx
|
||||
env_indices = np.random.randint(0, high=self.n_envs, size=(len(batch_inds),))
|
||||
|
||||
|
|
@ -743,7 +743,7 @@ class DictRolloutBuffer(RolloutBuffer):
|
|||
if self.pos == self.buffer_size:
|
||||
self.full = True
|
||||
|
||||
def get(self, batch_size: Optional[int] = None) -> Generator[DictRolloutBufferSamples, None, None]:
|
||||
def get(self, batch_size: Optional[int] = None) -> Generator[DictRolloutBufferSamples, None, None]: # type: ignore[signature-mismatch] #FIXME
|
||||
assert self.full, ""
|
||||
indices = np.random.permutation(self.buffer_size * self.n_envs)
|
||||
# Prepare the data
|
||||
|
|
@ -767,7 +767,7 @@ class DictRolloutBuffer(RolloutBuffer):
|
|||
yield self._get_samples(indices[start_idx : start_idx + batch_size])
|
||||
start_idx += batch_size
|
||||
|
||||
def _get_samples(self, batch_inds: np.ndarray, env: Optional[VecNormalize] = None) -> DictRolloutBufferSamples:
|
||||
def _get_samples(self, batch_inds: np.ndarray, env: Optional[VecNormalize] = None) -> DictRolloutBufferSamples: # type: ignore[signature-mismatch] #FIXME
|
||||
|
||||
return DictRolloutBufferSamples(
|
||||
observations={key: self.to_torch(obs[batch_inds]) for (key, obs) in self.observations.items()},
|
||||
|
|
|
|||
|
|
@ -71,7 +71,7 @@ class IdentityEnvBox(IdentityEnv[np.ndarray]):
|
|||
super().__init__(ep_length=ep_length, space=space)
|
||||
self.eps = eps
|
||||
|
||||
def step(self, action: np.ndarray) -> GymStepReturn:
|
||||
def step(self, action: np.ndarray) -> Tuple[np.ndarray, float, bool, Dict[str, Any]]:
|
||||
reward = self._get_reward(action)
|
||||
self._choose_next_state()
|
||||
self.current_step += 1
|
||||
|
|
@ -83,7 +83,7 @@ class IdentityEnvBox(IdentityEnv[np.ndarray]):
|
|||
|
||||
|
||||
class IdentityEnvMultiDiscrete(IdentityEnv[np.ndarray]):
|
||||
def __init__(self, dim: int = 1, ep_length: int = 100):
|
||||
def __init__(self, dim: int = 1, ep_length: int = 100) -> None:
|
||||
"""
|
||||
Identity environment for testing purposes
|
||||
|
||||
|
|
@ -95,7 +95,7 @@ class IdentityEnvMultiDiscrete(IdentityEnv[np.ndarray]):
|
|||
|
||||
|
||||
class IdentityEnvMultiBinary(IdentityEnv[np.ndarray]):
|
||||
def __init__(self, dim: int = 1, ep_length: int = 100):
|
||||
def __init__(self, dim: int = 1, ep_length: int = 100) -> None:
|
||||
"""
|
||||
Identity environment for testing purposes
|
||||
|
||||
|
|
@ -126,7 +126,7 @@ class FakeImageEnv(gym.Env):
|
|||
n_channels: int = 1,
|
||||
discrete: bool = True,
|
||||
channel_first: bool = False,
|
||||
):
|
||||
) -> None:
|
||||
self.observation_shape = (screen_height, screen_width, n_channels)
|
||||
if channel_first:
|
||||
self.observation_shape = (n_channels, screen_height, screen_width)
|
||||
|
|
|
|||
Loading…
Reference in a new issue