Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions iris/normalizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -148,7 +148,7 @@ def __call__(
action = utils.flatten(self._space, action)
action = (action * self._state["half_range"]) + self._state["mid"]
action = utils.unflatten(self._space, action)
action = self._add_ignored_input(action, ignored_action)
action = self._add_ignored_input(action, ignored_action) # pyrefly: ignore[bad-argument-type]
return action


Expand Down Expand Up @@ -189,7 +189,7 @@ def __call__(
observation = utils.flatten(self._space, observation)
observation = (observation - self._state["mid"]) / self._state["half_range"]
observation = utils.unflatten(self._space, observation)
observation = self._add_ignored_input(observation, ignored_observation)
observation = self._add_ignored_input(observation, ignored_observation) # pyrefly: ignore[bad-argument-type]
return observation


Expand Down Expand Up @@ -218,11 +218,11 @@ def __call__(
ignored_observation = self._filter_ignored_input(observation) # pyrefly: ignore[bad-argument-type]
observation = utils.flatten(self._space, observation)
if update_buffer:
self._buffer.push(observation)
self._buffer.push(observation) # pyrefly: ignore[bad-argument-type]
observation -= self._state[buffer.MEAN]
observation /= self._state[buffer.STD] + _EPSILON
observation = utils.unflatten(self._space, observation)
observation = self._add_ignored_input(observation, ignored_observation)
observation = self._add_ignored_input(observation, ignored_observation) # pyrefly: ignore[bad-argument-type]
return observation


Expand Down
2 changes: 1 addition & 1 deletion iris/policies/implicit_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -562,6 +562,6 @@ def act(
The actions in reinforcement learning.
"""
ob = utils.flatten(self._ob_space, ob)
action = self._action_calculator.act(ob)
action = self._action_calculator.act(ob) # pyrefly: ignore[bad-argument-type]
action = utils.unflatten(self._ac_space, action)
return action
2 changes: 1 addition & 1 deletion iris/policies/linear_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,6 @@ def act(self, ob: Union[np.ndarray, Dict[str, np.ndarray]]
"""
ob = utils.flatten(self._ob_space, ob)
matrix_weights = np.reshape(self._weights, (self._ac_dim, self._ob_dim))
actions = self._activation(np.dot(matrix_weights, ob)) # pyrefly: ignore[not-callable]
actions = self._activation(np.dot(matrix_weights, ob)) # pyrefly: ignore[no-matching-overload, not-callable]
actions = utils.unflatten(self._ac_space, actions)
return actions
2 changes: 1 addition & 1 deletion iris/policies/nas_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,7 +82,7 @@ def act(
ob = utils.flatten(self._ob_space, ob)
values = [0.0] * self._total_nb_nodes
for i in range(self._ob_dim):
values[i] = ob[i]
values[i] = ob[i] # pyrefly: ignore[bad-index]
for i in range(self._total_nb_nodes):
if (i > self._ob_dim) and (i < self._total_nb_nodes - self._ac_dim):
values[i] = np.tanh(values[i] + self._biases[i])
Expand Down
2 changes: 1 addition & 1 deletion iris/policies/nn_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,7 +70,7 @@ def act(self, ob: Union[np.ndarray, Dict[str, np.ndarray]]
mat_weight = np.reshape(
self._weights[start:end],
(self._layer_sizes[ith_layer + 1], self._layer_sizes[ith_layer]))
ith_layer_result = np.dot(mat_weight, ith_layer_result)
ith_layer_result = np.dot(mat_weight, ith_layer_result) # pyrefly: ignore[no-matching-overload]
ith_layer_result = self._activation(ith_layer_result) # pyrefly: ignore[not-callable]
actions = ith_layer_result
actions = utils.unflatten(self._ac_space, actions)
Expand Down
Loading