From 19954c1c31db50f99db92f62546d38a823511b80 Mon Sep 17 00:00:00 2001 From: Amer Elsheikh Date: Mon, 28 Sep 2026 15:13:56 -0700 Subject: [PATCH] Automated Code Change PiperOrigin-RevId: 989871237 --- iris/normalizer.py | 8 ++++---- iris/policies/implicit_policy.py | 2 +- iris/policies/linear_policy.py | 2 +- iris/policies/nas_policy.py | 2 +- iris/policies/nn_policy.py | 2 +- 5 files changed, 8 insertions(+), 8 deletions(-) diff --git a/iris/normalizer.py b/iris/normalizer.py index 94b4fa7..9067c34 100644 --- a/iris/normalizer.py +++ b/iris/normalizer.py @@ -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 @@ -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 @@ -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 diff --git a/iris/policies/implicit_policy.py b/iris/policies/implicit_policy.py index a1bf47e..a2777e8 100644 --- a/iris/policies/implicit_policy.py +++ b/iris/policies/implicit_policy.py @@ -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 diff --git a/iris/policies/linear_policy.py b/iris/policies/linear_policy.py index 4719d65..becfb4c 100644 --- a/iris/policies/linear_policy.py +++ b/iris/policies/linear_policy.py @@ -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 diff --git a/iris/policies/nas_policy.py b/iris/policies/nas_policy.py index 0307947..7054a06 100644 --- a/iris/policies/nas_policy.py +++ b/iris/policies/nas_policy.py @@ -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]) diff --git a/iris/policies/nn_policy.py b/iris/policies/nn_policy.py index 9e8d1ca..9ba0623 100644 --- a/iris/policies/nn_policy.py +++ b/iris/policies/nn_policy.py @@ -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)