diff --git a/iris/algorithms/ars_algorithm.py b/iris/algorithms/ars_algorithm.py index 1ecab04..3336476 100644 --- a/iris/algorithms/ars_algorithm.py +++ b/iris/algorithms/ars_algorithm.py @@ -189,7 +189,7 @@ def get_param_suggestions( param_suggestions = self._np_random_state.normal( 0, 1, (self._num_suggestions, dimensions) ) - self._last_std_used = self._std + self._last_std_used = self._std # pyrefly: ignore[bad-assignment] if callable(self._std): self._last_std_used = self._std(self._iteration) param_suggestions = np.vstack([ @@ -213,7 +213,7 @@ def state(self) -> Dict[str, Any]: def _get_state(self) -> Dict[str, Any]: state = {"params_to_eval": self._opt_params} if self._obs_norm_data_buffer is not None: - state["obs_norm_state"] = self._obs_norm_data_buffer.state + state["obs_norm_state"] = self._obs_norm_data_buffer.state # pyrefly: ignore[bad-assignment] return state @state.setter diff --git a/iris/algorithms/ars_algorithm_test.py b/iris/algorithms/ars_algorithm_test.py index 69075cc..8b55d69 100644 --- a/iris/algorithms/ars_algorithm_test.py +++ b/iris/algorithms/ars_algorithm_test.py @@ -92,7 +92,7 @@ def test_restore_state_from_checkpoint(self, expected_obs_norm_state): ) init_state = {'init_params': np.array([10.0, 10.0])} if expected_obs_norm_state: - init_state['obs_norm_buffer_data'] = { + init_state['obs_norm_buffer_data'] = { # pyrefly: ignore[bad-assignment] 'mean': np.asarray([0.0, 0.0]), 'std': np.asarray([1.0, 1.0]), 'n': 0, diff --git a/iris/algorithms/cma_algorithm.py b/iris/algorithms/cma_algorithm.py index 4d35908..c9baf4d 100644 --- a/iris/algorithms/cma_algorithm.py +++ b/iris/algorithms/cma_algorithm.py @@ -103,7 +103,7 @@ def process_evaluations(self, # Update the observation buffer if self._obs_norm_data_buffer is not None: for r in filtered_eval_results: - self._obs_norm_data_buffer.merge(r.obs_norm_buffer_data) + self._obs_norm_data_buffer.merge(r.obs_norm_buffer_data) # pyrefly: ignore[bad-argument-type] def get_param_suggestions(self, evaluate: bool = False @@ -141,7 +141,7 @@ def state(self) -> Dict[str, Any]: def _get_state(self) -> Dict[str, Any]: state = {"params_to_eval": self._opt_params} if self._obs_norm_data_buffer is not None: - state["obs_norm_state"] = self._obs_norm_data_buffer.state + state["obs_norm_state"] = self._obs_norm_data_buffer.state # pyrefly: ignore[bad-assignment] return state @state.setter diff --git a/iris/algorithms/learnable_ars_algorithm.py b/iris/algorithms/learnable_ars_algorithm.py index 6d5cf97..a839d2e 100644 --- a/iris/algorithms/learnable_ars_algorithm.py +++ b/iris/algorithms/learnable_ars_algorithm.py @@ -169,7 +169,7 @@ def get_param_suggestions( param_suggestions = self._np_random_state.normal( 0, 1, (self._num_suggestions, dimensions) ) - self._last_std_used = self._std + self._last_std_used = self._std # pyrefly: ignore[bad-assignment] param_suggestions = np.vstack([ self._opt_params, self._opt_params + self._last_std_used * param_suggestions, @@ -194,7 +194,7 @@ def process_evaluations( model_input = np.concatenate([[self._iteration], rewards]) if self._tree_weights is None: - self._model_state = self._restore_state_from_checkpoint(self._model_path) + self._model_state = self._restore_state_from_checkpoint(self._model_path) # pyrefly: ignore[bad-argument-type] self._tree_weights = self._model.init( jax.random.PRNGKey(seed=self._seed), model_input, self._model_state ) diff --git a/iris/algorithms/multi_agent_ars_algorithm.py b/iris/algorithms/multi_agent_ars_algorithm.py index 7edb4ed..4b53a25 100644 --- a/iris/algorithms/multi_agent_ars_algorithm.py +++ b/iris/algorithms/multi_agent_ars_algorithm.py @@ -110,7 +110,7 @@ def restore_state_from_checkpoint(self, new_state: Dict[str, Any]) -> None: ) } if self._obs_norm_data_buffer is not None: - duplicated_state["obs_norm_state"] = {} + duplicated_state["obs_norm_state"] = {} # pyrefly: ignore[bad-assignment] duplicated_state["obs_norm_state"]["mean"] = np.tile( new_state["obs_norm_state"]["mean"], self._num_agents ) @@ -193,10 +193,10 @@ def _get_top_evaluation_results( neg_eval_results: Sequence[worker_util.EvaluationResult], ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: pos_evals = np.array( - [r.metrics[f"reward_{agent_key}"] for r in pos_eval_results] + [r.metrics[f"reward_{agent_key}"] for r in pos_eval_results] # pyrefly: ignore[unsupported-operation] ) neg_evals = np.array( - [r.metrics[f"reward_{agent_key}"] for r in neg_eval_results] + [r.metrics[f"reward_{agent_key}"] for r in neg_eval_results] # pyrefly: ignore[unsupported-operation] ) if self._top_sort_type == "max": max_evals = np.max(np.vstack([pos_evals, neg_evals]), axis=0) diff --git a/iris/algorithms/multi_agent_ars_algorithm_test.py b/iris/algorithms/multi_agent_ars_algorithm_test.py index 2f398e4..1cc662d 100644 --- a/iris/algorithms/multi_agent_ars_algorithm_test.py +++ b/iris/algorithms/multi_agent_ars_algorithm_test.py @@ -224,7 +224,7 @@ def test_restore_state_from_checkpoint( self.assertEqual(algo._num_agents, num_agents) init_state = {'init_params': np.array([10.0, 10.0])} if state['obs_norm_state'] is not None: - init_state['obs_norm_buffer_data'] = { + init_state['obs_norm_buffer_data'] = { # pyrefly: ignore[bad-assignment] 'mean': np.asarray([0.0, 0.0]), 'std': np.asarray([1.0, 1.0]), 'n': 0, diff --git a/iris/algorithms/optimizers.py b/iris/algorithms/optimizers.py index cfbdddc..5b27c99 100644 --- a/iris/algorithms/optimizers.py +++ b/iris/algorithms/optimizers.py @@ -87,7 +87,7 @@ def vector_decoding_function(A, b, optimization_parameters, loss_function): result = x.value res_list = [] for i in range(n): - res_list.append(result[i]) + res_list.append(result[i]) # pyrefly: ignore[unsupported-operation] return np.array(res_list) @@ -129,7 +129,7 @@ def general_jacobian_decoder(atranspose, yprime, optimization_parameters, list_res = [] for j in range(n): list_res.append(res[j]) - final_solutions.append(np.float32(list_res)) + final_solutions.append(np.float32(list_res)) # pyrefly: ignore[bad-argument-type] return np.array(final_solutions) diff --git a/iris/algorithms/pes_algorithm.py b/iris/algorithms/pes_algorithm.py index af94289..eb5eafd 100644 --- a/iris/algorithms/pes_algorithm.py +++ b/iris/algorithms/pes_algorithm.py @@ -111,7 +111,7 @@ def process_evaluations( pos_directions.append((params - self._opt_params) / self._std) pos_directions[-1] = self._positive_cumulative_perturbations[ i] + pos_directions[-1] - if pos_eval_results[i].metrics["current_step"] == 0: + if pos_eval_results[i].metrics["current_step"] == 0: # pyrefly: ignore[unsupported-operation] self._positive_cumulative_perturbations[i] = 0 else: self._positive_cumulative_perturbations[i] = pos_directions[-1] @@ -120,7 +120,7 @@ def process_evaluations( neg_directions.append((params - self._opt_params) / self._std) neg_directions[-1] = self._negative_cumulative_perturbations[ i] + neg_directions[-1] - if neg_eval_results[i].metrics["current_step"] == 0: + if neg_eval_results[i].metrics["current_step"] == 0: # pyrefly: ignore[unsupported-operation] self._negative_cumulative_perturbations[i] = 0 else: self._negative_cumulative_perturbations[i] = neg_directions[-1] @@ -136,7 +136,7 @@ def process_evaluations( max_evals = np.max(np.vstack([pos_evals, neg_evals]), axis=0) elif self._top_sort_type == "diff": max_evals = np.abs(pos_evals - neg_evals) - idx = (-max_evals).argsort()[:self._num_top] + idx = (-max_evals).argsort()[:self._num_top] # pyrefly: ignore[unbound-name] pos_evals = pos_evals[idx] neg_evals = neg_evals[idx] all_top_evals = np.hstack([pos_evals, neg_evals]) @@ -212,7 +212,7 @@ def state(self) -> Dict[str, Any]: def _get_state(self) -> Dict[str, Any]: state = {"params_to_eval": self._opt_params} if self._obs_norm_data_buffer is not None: - state["obs_norm_state"] = self._obs_norm_data_buffer.state + state["obs_norm_state"] = self._obs_norm_data_buffer.state # pyrefly: ignore[bad-assignment] return state @state.setter diff --git a/iris/algorithms/piars_algorithm.py b/iris/algorithms/piars_algorithm.py index 93245a2..7c4b9f9 100644 --- a/iris/algorithms/piars_algorithm.py +++ b/iris/algorithms/piars_algorithm.py @@ -207,10 +207,10 @@ def __init__( else: self.policy = policy - obs_spec = gym_wrapper.spec_from_gym_space(self._env.observation_space) - action_spec = gym_wrapper.spec_from_gym_space(self._env.action_space) + obs_spec = gym_wrapper.spec_from_gym_space(self._env.observation_space) # pyrefly: ignore[bad-argument-type] + action_spec = gym_wrapper.spec_from_gym_space(self._env.action_space) # pyrefly: ignore[bad-argument-type] time_step_spec = ts.time_step_spec(observation_spec=obs_spec) - policy_step_spec = policy_step.PolicyStep(action=action_spec) + policy_step_spec = policy_step.PolicyStep(action=action_spec) # pyrefly: ignore[missing-argument] collect_data_spec = trajectory.from_transition( time_step_spec, policy_step_spec, time_step_spec ) @@ -293,7 +293,7 @@ def train(self, obs_norm_state=None): if self.global_step % self.reverb_checkpoint_period == 0: logging.info("Start checkpointing reverb data.") self.reverb_rb.py_client.checkpoint() - print("train/loss: {}".format(np.mean(loss.numpy()))) + print("train/loss: {}".format(np.mean(loss.numpy()))) # pyrefly: ignore[unbound-name] @tf.function def train_single_step(self, obs, reward, action, discount): @@ -325,14 +325,14 @@ def train_single_step(self, obs, reward, action, discount): @tf.function def rollout(self, obs, actions): """Latent rollout.""" - s, _ = self.policy.h_model(obs) + s, _ = self.policy.h_model(obs) # pyrefly: ignore[not-callable] outputs = [] for i in range(self._rollout_length): - p, v = self.policy.f_model(s) - u_next, s_next = self.policy.g_model([s, actions[:, i, ...]]) + p, v = self.policy.f_model(s) # pyrefly: ignore[not-callable] + u_next, s_next = self.policy.g_model([s, actions[:, i, ...]]) # pyrefly: ignore[not-callable] outputs.append((p, v, u_next, s)) s = s_next - p, v = self.policy.f_model(s) + p, v = self.policy.f_model(s) # pyrefly: ignore[not-callable] outputs.append((p, v, None, s)) return outputs @@ -368,11 +368,11 @@ def infonce(hidden_x, hidden_y, temperature=0.1): # Latent state (from visual + other observations) for the first time step hx = latent_traj[0][-1] # Latent state (from visual observations) for the last time step - _, hy_vision = self.policy.h_model(obs_k) + _, hy_vision = self.policy.h_model(obs_k) # pyrefly: ignore[not-callable] # A trick from https://arxiv.org/abs/2011.10566 hy_vision = tf.stop_gradient(hy_vision) - zx = self.policy.px_model(hx) - zy = self.policy.py_model(hy_vision) + zx = self.policy.px_model(hx) # pyrefly: ignore[not-callable] + zy = self.policy.py_model(hy_vision) # pyrefly: ignore[not-callable] iyz, _, _ = infonce(zx, zy, temperature=0.1) loss_pi = -iyz @@ -406,11 +406,11 @@ def infonce(hidden_x, hidden_y, temperature=0.1): loss_v += self.distributional_value_loss( value_logits=z, value_supports=self.supports, - target_value_logits=last_value_distribution, - target_value_supports=target_value_supports[i], + target_value_logits=last_value_distribution, # pyrefly: ignore[unbound-name] + target_value_supports=target_value_supports[i], # pyrefly: ignore[unbound-name] ) vd = tf.nn.softmax(z) - pred_value_sum += tf.reduce_sum(vd * self.supports[None, ...], axis=-1) + pred_value_sum += tf.reduce_sum(vd * self.supports[None, ...], axis=-1) # pyrefly: ignore[unbound-name] # reward loss loss_r += tf.reduce_sum( tf.math.square(u_next - tf.stop_gradient(rewards[:, i : i + 1])), -1 @@ -424,9 +424,9 @@ def infonce(hidden_x, hidden_y, temperature=0.1): "loss_pi": tf.reduce_mean(loss_pi), } if self.use_value_loss: - metrics["value"] = tf.reduce_mean(pred_value_sum) / self._rollout_length + metrics["value"] = tf.reduce_mean(pred_value_sum) / self._rollout_length # pyrefly: ignore[unbound-name] if self.use_pi_loss: - metrics["iyz"] = tf.reduce_mean(iyz) + metrics["iyz"] = tf.reduce_mean(iyz) # pyrefly: ignore[unbound-name] return loss, metrics def distributional_value_loss( @@ -455,7 +455,7 @@ def flatten_nested(space, x): """Flatten nested.""" if isinstance(space, spaces.Box): x = np.asarray(x, dtype=np.float32) - inner_dims = list(space.shape) + inner_dims = list(space.shape) # pyrefly: ignore[bad-argument-type] outer_dims = list(x.shape)[: -len(inner_dims)] x = np.reshape(x, outer_dims + [np.prod(inner_dims)]) return x diff --git a/iris/algorithms/pyglove_algorithm.py b/iris/algorithms/pyglove_algorithm.py index 12904a7..f379c7c 100644 --- a/iris/algorithms/pyglove_algorithm.py +++ b/iris/algorithms/pyglove_algorithm.py @@ -97,7 +97,7 @@ def get_param_suggestions(self, for metadata in metadata_list: suggestion = {"params_to_eval": np.empty((), dtype=np.float64)} - suggestion["metadata"] = metadata + suggestion["metadata"] = metadata # pyrefly: ignore[bad-assignment] vanilla_suggestions.append(suggestion) return vanilla_suggestions diff --git a/iris/algorithms/pyribs_algorithm.py b/iris/algorithms/pyribs_algorithm.py index f637ff8..31874b8 100644 --- a/iris/algorithms/pyribs_algorithm.py +++ b/iris/algorithms/pyribs_algorithm.py @@ -181,7 +181,7 @@ def get_param_suggestions( buffer_lib.STD: elite[_OBS_NORM_STD], } else: - param_suggestions = self._scheduler.ask() + param_suggestions = self._scheduler.ask() # pyrefly: ignore[missing-attribute] buffer = self._obs_norm_data_buffer.state return [ @@ -202,12 +202,12 @@ def process_evaluations( obs_norm_std = [] obs_norm_mean = [] for result in eval_results: - self._obs_norm_data_buffer.merge(result.obs_norm_buffer_data) + self._obs_norm_data_buffer.merge(result.obs_norm_buffer_data) # pyrefly: ignore[bad-argument-type] objective.append(result.value) - measures.append([result.metrics[name] for name in self._measure_names]) - obs_norm_n.append(result.obs_norm_buffer_data[buffer_lib.N]) - obs_norm_std.append(result.obs_norm_buffer_data[buffer_lib.STD]) - obs_norm_mean.append(result.obs_norm_buffer_data[buffer_lib.MEAN]) + measures.append([result.metrics[name] for name in self._measure_names]) # pyrefly: ignore[unsupported-operation] + obs_norm_n.append(result.obs_norm_buffer_data[buffer_lib.N]) # pyrefly: ignore[unsupported-operation] + obs_norm_std.append(result.obs_norm_buffer_data[buffer_lib.STD]) # pyrefly: ignore[unsupported-operation] + obs_norm_mean.append(result.obs_norm_buffer_data[buffer_lib.MEAN]) # pyrefly: ignore[unsupported-operation] # Store the state of the obs_norm_buffer for each solution so that it can be # reproduced later when evaluating the policy, similar to other algorithms @@ -218,7 +218,7 @@ def process_evaluations( _OBS_NORM_N: obs_norm_n, } - self._scheduler.tell( + self._scheduler.tell( # pyrefly: ignore[missing-attribute] objective=objective, measures=measures, **extra_fields, diff --git a/iris/algorithms/pyribs_algorithm_test.py b/iris/algorithms/pyribs_algorithm_test.py index b1e1ac1..1f44193 100644 --- a/iris/algorithms/pyribs_algorithm_test.py +++ b/iris/algorithms/pyribs_algorithm_test.py @@ -102,7 +102,7 @@ def test_get_param_suggestions_for_eval(self): # Give the first evaluation a high score so it is the elite. evaluations[0].value = 1000 if evaluations[0].obs_norm_buffer_data is not None: - evaluations[0].obs_norm_buffer_data[buffer.N] = 1000 + evaluations[0].obs_norm_buffer_data[buffer.N] = 1000 # pyrefly: ignore[unsupported-operation] self.test_algorithm.process_evaluations(evaluations) eval_suggestions = self.test_algorithm.get_param_suggestions(evaluate=True) @@ -115,7 +115,7 @@ def test_get_param_suggestions_for_eval(self): ) np.testing.assert_equal( eval_suggestion[algorithm.OBS_NORM_BUFFER_STATE][buffer.N], - evaluations[0].obs_norm_buffer_data[buffer.N], + evaluations[0].obs_norm_buffer_data[buffer.N], # pyrefly: ignore[unsupported-operation] ) self.assertFalse(eval_suggestion[algorithm.UPDATE_OBS_NORM_BUFFER]) @@ -208,7 +208,7 @@ def test_process_evaluations(self): worker_util.EvaluationResult( params_evaluated=np.ones((13,)), value=1, - obs_norm_buffer_data={ + obs_norm_buffer_data={ # pyrefly: ignore[bad-argument-type] buffer.N: 1, buffer.STD: np.ones((8,)), buffer.MEAN: np.ones((8,)), @@ -219,7 +219,7 @@ def test_process_evaluations(self): worker_util.EvaluationResult( params_evaluated=np.ones((13,) * 2), value=2, - obs_norm_buffer_data={ + obs_norm_buffer_data={ # pyrefly: ignore[bad-argument-type] buffer.N: 2, buffer.STD: np.ones((8,)) * 2, buffer.MEAN: np.ones((8,)) * 2, diff --git a/iris/buffer.py b/iris/buffer.py index e1443af..ca143a8 100644 --- a/iris/buffer.py +++ b/iris/buffer.py @@ -197,7 +197,7 @@ def state(self, new_state: Dict[str, Any]) -> None: @property def _var(self) -> np.ndarray: return ( - self._data[UNNORM_VAR] / (self._data[N] - 1) + self._data[UNNORM_VAR] / (self._data[N] - 1) # pyrefly: ignore[bad-return] if self._data[N] > 1 else np.ones_like(self._data[MEAN]) ) diff --git a/iris/checkpoint_evaluator.py b/iris/checkpoint_evaluator.py index 741e23c..dc2cfe3 100644 --- a/iris/checkpoint_evaluator.py +++ b/iris/checkpoint_evaluator.py @@ -49,7 +49,7 @@ def main(argv): worker_config.worker_args.write_to_replay = False worker = worker_config["worker_class"]( worker_id=0, **worker_config["worker_args"]) - state = checkpoint_util.load_checkpoint_state(_CHECKPOINT_FILE.value) + state = checkpoint_util.load_checkpoint_state(_CHECKPOINT_FILE.value) # pyrefly: ignore[bad-argument-type] returns = [] times = [] metric_dict = collections.defaultdict(list) @@ -61,7 +61,7 @@ def main(argv): gfile.MakeDirs( _VIDEO_PATH.value, mode=gfile.LEGACY_GROUP_WRITABLE_WORLD_READABLE ) - video_path = os.path.join(_VIDEO_PATH.value, "video_" + str(i) + ".mp4") + video_path = os.path.join(_VIDEO_PATH.value, "video_" + str(i) + ".mp4") # pyrefly: ignore[no-matching-overload] result = worker.work( **state, enable_logging=True, diff --git a/iris/coordinator.py b/iris/coordinator.py index 5f6c758..1ede87c 100644 --- a/iris/coordinator.py +++ b/iris/coordinator.py @@ -476,7 +476,7 @@ def _restore_checkpoint( int(checkpoint_paths_sorted[0].split("_")[-1]) + 1 ) if latest_checkpoint_num_iterations > max_allowed_iteration_for_restart: - raise checkpoint_load_error + raise checkpoint_load_error # pyrefly: ignore[bad-raise] return None, 0 def evaluate( diff --git a/iris/coordinator_rl_test.py b/iris/coordinator_rl_test.py index 5c79c3e..4ab6d3e 100644 --- a/iris/coordinator_rl_test.py +++ b/iris/coordinator_rl_test.py @@ -94,7 +94,7 @@ def make_bb_program( workers.append(worker_handle) if warmstartdir: - warmstartdir = pathlib.Path(warmstartdir) + warmstartdir = pathlib.Path(warmstartdir) # pyrefly: ignore[bad-assignment] algo = algo_config["algorithm_class"](**algo_config["algorithm_args"]) # Launches eval worker instances if there is at least one num_eval_workers. diff --git a/iris/coordinator_test.py b/iris/coordinator_test.py index 72b3639..62912fe 100644 --- a/iris/coordinator_test.py +++ b/iris/coordinator_test.py @@ -69,7 +69,7 @@ def make_bb_program( workers.append(worker_handle) if warmstartdir: - warmstartdir = pathlib.Path(warmstartdir) + warmstartdir = pathlib.Path(warmstartdir) # pyrefly: ignore[bad-assignment] algo = algo_config["algorithm_class"](**algo_config["algorithm_args"]) # Launches eval worker instances if there is at least one num_eval_workers. diff --git a/iris/normalizer.py b/iris/normalizer.py index 8d68e57..94b4fa7 100644 --- a/iris/normalizer.py +++ b/iris/normalizer.py @@ -144,7 +144,7 @@ def __call__( """ del update_buffer # No buffer to update action = action.copy() - ignored_action = self._filter_ignored_input(action) + ignored_action = self._filter_ignored_input(action) # pyrefly: ignore[bad-argument-type] action = utils.flatten(self._space, action) action = (action * self._state["half_range"]) + self._state["mid"] action = utils.unflatten(self._space, action) @@ -184,7 +184,7 @@ def __call__( """ del update_buffer # No buffer to update observation = observation.copy() - ignored_observation = self._filter_ignored_input(observation) + ignored_observation = self._filter_ignored_input(observation) # pyrefly: ignore[bad-argument-type] observation = utils.flatten(self._space, observation) observation = (observation - self._state["mid"]) / self._state["half_range"] @@ -215,7 +215,7 @@ def __call__( update_buffer: bool = True, ) -> Union[np.ndarray, Dict[str, np.ndarray]]: observation = observation.copy() - ignored_observation = self._filter_ignored_input(observation) + 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) @@ -245,7 +245,7 @@ def __call__( ) -> Union[np.ndarray, Dict[str, np.ndarray]]: observation = observation.copy() - opp_observation = self._filter_ignored_input(observation) + opp_observation = self._filter_ignored_input(observation) # pyrefly: ignore[bad-argument-type] opp_observation = utils.flatten(self._space_ignored, opp_observation) arm_observation = utils.flatten(self._space, observation) diff --git a/iris/policies/layers/keras_image_encoder_layer_test.py b/iris/policies/layers/keras_image_encoder_layer_test.py index 92efbd5..79b1dba 100644 --- a/iris/policies/layers/keras_image_encoder_layer_test.py +++ b/iris/policies/layers/keras_image_encoder_layer_test.py @@ -24,7 +24,7 @@ def test_layer_output(self): """Tests the output of ImageEncoder layer.""" input_layer = tf.keras.layers.Input( batch_input_shape=(2, 5, 6, 2), dtype="float", name="input") - output_layer = keras_image_encoder_layer.ImageEncoder( + output_layer = keras_image_encoder_layer.ImageEncoder( # pyrefly: ignore[not-callable] patch_height=2, patch_width=2, stride_height=1, diff --git a/iris/policies/layers/keras_masking_attention_layer_test.py b/iris/policies/layers/keras_masking_attention_layer_test.py index f19407b..72ae799 100644 --- a/iris/policies/layers/keras_masking_attention_layer_test.py +++ b/iris/policies/layers/keras_masking_attention_layer_test.py @@ -29,7 +29,7 @@ def test_layer_output(self): batch_input_shape=(2, 3, 4), dtype="float", name="keys") value_layer = tf.keras.layers.Input( batch_input_shape=(2, 3, 4), dtype="float", name="values") - output_layer = keras_masking_attention_layer.FavorMaskingAttention( + output_layer = keras_masking_attention_layer.FavorMaskingAttention( # pyrefly: ignore[not-callable] kernel_transformation=favor.relu_kernel_transformation, top_k=2)(query_layer, key_layer, value_layer) model = tf.keras.models.Model( diff --git a/iris/policies/layers/keras_positional_encoding_layer.py b/iris/policies/layers/keras_positional_encoding_layer.py index 865e88c..85a203c 100644 --- a/iris/policies/layers/keras_positional_encoding_layer.py +++ b/iris/policies/layers/keras_positional_encoding_layer.py @@ -29,7 +29,7 @@ def call(self, indices = tf.expand_dims(tf.range(seq_len), 0) indices = tf.tile(indices, [num_freq, 1]) freq_fn = lambda k: 1.0/(10000 ** (2*k/encoding_dimension)) - freq = tf.keras.layers.Lambda(freq_fn)(tf.range(num_freq)) + freq = tf.keras.layers.Lambda(freq_fn)(tf.range(num_freq)) # pyrefly: ignore[not-callable] freq = tf.expand_dims(freq, 1) freq = tf.tile(freq, [1, seq_len]) args = tf.multiply(freq, tf.cast(indices, dtype=tf.float64)) diff --git a/iris/policies/layers/keras_positional_encoding_layer_test.py b/iris/policies/layers/keras_positional_encoding_layer_test.py index 806c83e..7faaa82 100644 --- a/iris/policies/layers/keras_positional_encoding_layer_test.py +++ b/iris/policies/layers/keras_positional_encoding_layer_test.py @@ -20,7 +20,7 @@ class PositionalEncodingTest(absltest.TestCase): def test_layer_output(self): """Tests the output of PositionalEncoding layer.""" - encoding = keras_positional_encoding_layer.PositionalEncoding()(7, 4) + encoding = keras_positional_encoding_layer.PositionalEncoding()(7, 4) # pyrefly: ignore[not-callable] self.assertEqual(encoding.shape, (1, 7, 4)) if __name__ == "__main__": diff --git a/iris/policies/layers/keras_ranking_attention_layer_test.py b/iris/policies/layers/keras_ranking_attention_layer_test.py index eb70412..9c1fcfe 100644 --- a/iris/policies/layers/keras_ranking_attention_layer_test.py +++ b/iris/policies/layers/keras_ranking_attention_layer_test.py @@ -29,7 +29,7 @@ def test_layer_output(self): batch_input_shape=(2, 3, 4), dtype="float", name="keys") value_layer = tf.keras.layers.Input( batch_input_shape=(2, 3, 4), dtype="float", name="values") - output_layer = keras_ranking_attention_layer.FavorRankingAttention( + output_layer = keras_ranking_attention_layer.FavorRankingAttention( # pyrefly: ignore[not-callable] kernel_transformation=favor.relu_kernel_transformation, top_k=2)(query_layer, key_layer, value_layer) model = tf.keras.models.Model( diff --git a/iris/policies/layers/keras_trans_attention_layer_test.py b/iris/policies/layers/keras_trans_attention_layer_test.py index ced9ec0..b3d3a90 100644 --- a/iris/policies/layers/keras_trans_attention_layer_test.py +++ b/iris/policies/layers/keras_trans_attention_layer_test.py @@ -29,7 +29,7 @@ def test_layer_output(self): batch_input_shape=(2, 3, 4), dtype="float", name="keys") value_layer = tf.keras.layers.Input( batch_input_shape=(2, 3, 4), dtype="float", name="values") - output_layer = keras_trans_attention_layer.FavorTransAttention( + output_layer = keras_trans_attention_layer.FavorTransAttention( # pyrefly: ignore[not-callable] kernel_transformation=favor.relu_kernel_transformation)( query_layer, key_layer, value_layer) model = tf.keras.models.Model( diff --git a/iris/workers/maml_worker.py b/iris/workers/maml_worker.py index 9a2ce88..9f38017 100644 --- a/iris/workers/maml_worker.py +++ b/iris/workers/maml_worker.py @@ -33,7 +33,7 @@ def _multiple_eval( ) -> Tuple[float, Sequence[worker_util.EvaluationResult]]: """Evaluates parameters multiple times and averages results.""" results = [work_fn(params_to_eval, **work_kwargs) for _ in range(num_evals)] - return np.mean([r.value for r in results]), results + return np.mean([r.value for r in results]), results # pyrefly: ignore[no-matching-overload] # TODO: Potentially make this a subclass of BlackboxAlgorithm. @@ -300,7 +300,7 @@ def __init__( self._adaptation_optimizer = adaptation_constructor() self._init_state = self._worker._init_state - def work( + def work( # pyrefly: ignore[bad-override] self, params_to_eval: Any, **work_kwargs # pytype: disable=signature-mismatch # overriding-parameter-count-checks ) -> worker_util.EvaluationResult: """Uses another Worker's work() function for adaptation. diff --git a/iris/workers/maml_worker_test.py b/iris/workers/maml_worker_test.py index b469be8..0d5c75f 100644 --- a/iris/workers/maml_worker_test.py +++ b/iris/workers/maml_worker_test.py @@ -45,7 +45,7 @@ def test_multiple_eval(self): work_fn=self.worker_obj.work, ) self.assertLen(results, 5) - self.assertEqual(mean_val, np.mean([result.value for result in results])) + self.assertEqual(mean_val, np.mean([result.value for result in results])) # pyrefly: ignore[no-matching-overload] def test_gradient_adaptation(self): num_iterations = 2 @@ -66,7 +66,7 @@ def test_gradient_adaptation(self): ) self.assertEqual( - val, np.mean([result.value for result in results[-num_adapted_evals:]]) + val, np.mean([result.value for result in results[-num_adapted_evals:]]) # pyrefly: ignore[no-matching-overload] ) meta_value = self.worker_obj.work(self.init_params).value @@ -98,7 +98,7 @@ def test_hillclimb_adaptation(self, parallel_alg: str): + num_adapted_evals, ) self.assertEqual( - val, np.mean([result.value for result in results[-num_adapted_evals:]]) + val, np.mean([result.value for result in results[-num_adapted_evals:]]) # pyrefly: ignore[no-matching-overload] ) self.assertGreaterEqual(val, meta_value) @@ -109,7 +109,7 @@ def test_hillclimb_adaptation(self, parallel_alg: str): ) self.assertEqual( val, - np.mean([result.value for result in new_results[-num_adapted_evals:]]), + np.mean([result.value for result in new_results[-num_adapted_evals:]]), # pyrefly: ignore[no-matching-overload] ) self.assertLen( new_results, diff --git a/iris/workers/multi_agent_rl_worker.py b/iris/workers/multi_agent_rl_worker.py index 272c33d..d5818bf 100644 --- a/iris/workers/multi_agent_rl_worker.py +++ b/iris/workers/multi_agent_rl_worker.py @@ -68,11 +68,11 @@ def work( # pytype: disable=signature-mismatch # overriding-default-value-chec self._observation_normalizer.buffer.reset() if obs_norm_state is not None: - self._observation_normalizer.state = obs_norm_state + self._observation_normalizer.state = obs_norm_state # pyrefly: ignore[bad-argument-type] video = None if record_video: - video = video_recorder.VideoRecorder(video_path, video_framerate) + video = video_recorder.VideoRecorder(video_path, video_framerate) # pyrefly: ignore[bad-argument-type] reward_dict = collections.defaultdict(float) agent_1 = None @@ -115,9 +115,9 @@ def work( # pytype: disable=signature-mismatch # overriding-default-value-chec "done": done, rl_worker.INFO: info, } - mdict = self._metrics_fn(self._env, step_output) + mdict = self._metrics_fn(self._env, step_output) # pyrefly: ignore[not-callable] for metric_name, metric_value in mdict.items(): - metrics[metric_name] += metric_value + metrics[metric_name] += metric_value # pyrefly: ignore[unsupported-operation] for rkey, rval in reward_dict.items(): metrics[f"reward_{rkey}"] = rval diff --git a/iris/workers/pyglove_rl_worker.py b/iris/workers/pyglove_rl_worker.py index 27bd303..fed62b2 100644 --- a/iris/workers/pyglove_rl_worker.py +++ b/iris/workers/pyglove_rl_worker.py @@ -40,7 +40,7 @@ def __init__( self._policy.dna_spec # pytype: disable=attribute-error ) - def work( + def work( # pyrefly: ignore[bad-override] self, metadata: Optional[str] = None, **kwargs ) -> worker_util.EvaluationResult: if metadata: diff --git a/iris/workers/rl_representation_worker.py b/iris/workers/rl_representation_worker.py index 33252a8..c0d7ebe 100644 --- a/iris/workers/rl_representation_worker.py +++ b/iris/workers/rl_representation_worker.py @@ -56,10 +56,10 @@ def __init__( if reverb_client is not None: self._init_state["reverb_server_addr"] = reverb_client.server_address - obs_spec = gym_wrapper.spec_from_gym_space(self._env.observation_space) - action_spec = gym_wrapper.spec_from_gym_space(self._env.action_space) + obs_spec = gym_wrapper.spec_from_gym_space(self._env.observation_space) # pyrefly: ignore[bad-argument-type] + action_spec = gym_wrapper.spec_from_gym_space(self._env.action_space) # pyrefly: ignore[bad-argument-type] time_step_spec = ts.time_step_spec(observation_spec=obs_spec) - policy_step_spec = policy_step.PolicyStep(action=action_spec) + policy_step_spec = policy_step.PolicyStep(action=action_spec) # pyrefly: ignore[missing-argument] collect_data_spec = trajectory.from_transition( time_step_spec, policy_step_spec, time_step_spec ) @@ -111,14 +111,14 @@ def work( # pytype: disable=signature-mismatch # overriding-default-value-chec self._observation_normalizer.buffer.reset() if obs_norm_state is not None: - self._observation_normalizer.state = obs_norm_state + self._observation_normalizer.state = obs_norm_state # pyrefly: ignore[bad-argument-type] if env_seed is not None: self._env.seed(env_seed) video = None if record_video: - video = video_recorder.VideoRecorder(video_path, video_framerate) + video = video_recorder.VideoRecorder(video_path, video_framerate) # pyrefly: ignore[bad-argument-type] reward = 0.0 metrics = collections.defaultdict(float) @@ -134,7 +134,7 @@ def work( # pytype: disable=signature-mismatch # overriding-default-value-chec for st in range(self._rollout_length): normalized_obs = self._observation_normalizer(obs, update_obs_norm_buffer) action = self._policy.act(normalized_obs) - action_step = policy_step.PolicyStep(action) + action_step = policy_step.PolicyStep(action) # pyrefly: ignore[missing-argument] action = self._action_denormalizer(action) next_obs, r, done, info = self._env.step(action) reward += r @@ -158,9 +158,9 @@ def work( # pytype: disable=signature-mismatch # overriding-default-value-chec "done": done, rl_worker.INFO: info, } - mdict = self._metrics_fn(self._env, step_output) + mdict = self._metrics_fn(self._env, step_output) # pyrefly: ignore[not-callable] for metric_name, metric_value in mdict.items(): - metrics[metric_name] += metric_value + metrics[metric_name] += metric_value # pyrefly: ignore[unsupported-operation] obs = next_obs time_step = next_time_step diff --git a/iris/workers/rl_representation_worker_test.py b/iris/workers/rl_representation_worker_test.py index 5075b77..c96e349 100644 --- a/iris/workers/rl_representation_worker_test.py +++ b/iris/workers/rl_representation_worker_test.py @@ -52,7 +52,7 @@ def test_rl_representation_worker(self): self.assertIn('INFO:absl:Total Reward:', logs.output[-1]) self.assertLen(logs.output, 401) self.assertLessEqual(result1.value, 0) - self.assertLessEqual(result1.metrics['extra_metric'], 500.0) + self.assertLessEqual(result1.metrics['extra_metric'], 500.0) # pyrefly: ignore[unsupported-operation] result2 = worker_obj.work( params_to_eval=np.ones(3), diff --git a/iris/workers/rl_worker.py b/iris/workers/rl_worker.py index 70f068d..a219903 100644 --- a/iris/workers/rl_worker.py +++ b/iris/workers/rl_worker.py @@ -124,13 +124,13 @@ def __init__( if not isinstance(observation_normalizer, normalizer.Normalizer): self._observation_normalizer = observation_normalizer( - self._env.observation_space + self._env.observation_space # pyrefly: ignore[bad-argument-type] ) else: self._observation_normalizer = observation_normalizer if not isinstance(action_denormalizer, normalizer.Normalizer): - self._action_denormalizer = action_denormalizer(self._env.action_space) + self._action_denormalizer = action_denormalizer(self._env.action_space) # pyrefly: ignore[bad-argument-type] else: self._action_denormalizer = action_denormalizer @@ -246,11 +246,11 @@ def _run_rollout( self._observation_normalizer.buffer.reset() if obs_norm_state is not None: - self._observation_normalizer.state = obs_norm_state + self._observation_normalizer.state = obs_norm_state # pyrefly: ignore[bad-argument-type] video = None if record_video: - video = video_recorder.VideoRecorder(video_path, video_framerate) + video = video_recorder.VideoRecorder(video_path, video_framerate) # pyrefly: ignore[bad-argument-type] rewards = [] metrics = collections.defaultdict(list) @@ -286,7 +286,7 @@ def _run_rollout( "done": done, INFO: info, } - mdict = self._metrics_fn(self._env, step_output) + mdict = self._metrics_fn(self._env, step_output) # pyrefly: ignore[not-callable] for metric_name, metric_value in mdict.items(): metrics[metric_name].append(metric_value) @@ -306,12 +306,12 @@ def _run_rollout( self._step = 0 break - aggregate_reward, reward_stats = self._stats_fn( + aggregate_reward, reward_stats = self._stats_fn( # pyrefly: ignore[not-callable] "reward", rewards, **self._stats_fn_args ) aggregate_metrics_with_stats = {} for metric_name, metric_values in metrics.items(): - metric, metric_stats = self._stats_fn( + metric, metric_stats = self._stats_fn( # pyrefly: ignore[not-callable] metric_name, metric_values, **self._stats_fn_args ) aggregate_metrics_with_stats[metric_name] = metric diff --git a/iris/workers/worker_util.py b/iris/workers/worker_util.py index 268705e..7ad10d9 100644 --- a/iris/workers/worker_util.py +++ b/iris/workers/worker_util.py @@ -42,21 +42,21 @@ def merge_eval_results(results: Sequence[EvaluationResult]) -> EvaluationResult: if len(results) == 1: return results[0] - merged_value = np.mean([r.value for r in results]) + merged_value = np.mean([r.value for r in results]) # pyrefly: ignore[no-matching-overload] merged_obs_norm_buffer_data = None if results[0].obs_norm_buffer_data: merged_buffer = buffer.MeanStdBuffer() merged_buffer.data = results[0].obs_norm_buffer_data for result in itertools.islice(results, 1, None): - merged_buffer.merge(result.obs_norm_buffer_data) + merged_buffer.merge(result.obs_norm_buffer_data) # pyrefly: ignore[bad-argument-type] merged_obs_norm_buffer_data = merged_buffer.data if results[0].metrics is not None: merged_metrics = {} for metric_name in results[0].metrics: - merged_metrics[metric_name] = np.mean( - [result.metrics[metric_name] for result in results]) + merged_metrics[metric_name] = np.mean( # pyrefly: ignore[no-matching-overload] + [result.metrics[metric_name] for result in results]) # pyrefly: ignore[unsupported-operation] else: merged_metrics = {} diff --git a/iris/workers/worker_util_test.py b/iris/workers/worker_util_test.py index 8349c76..5fd657e 100644 --- a/iris/workers/worker_util_test.py +++ b/iris/workers/worker_util_test.py @@ -24,7 +24,7 @@ def test_merge(self): result1 = worker_util.EvaluationResult( params_evaluated=np.zeros(6), value=np.float64(5.0), - obs_norm_buffer_data={ + obs_norm_buffer_data={ # pyrefly: ignore[bad-argument-type] 'n': 5, 'mean': np.zeros(7), 'unnorm_var': np.ones(7), @@ -35,7 +35,7 @@ def test_merge(self): result2 = worker_util.EvaluationResult( params_evaluated=np.zeros(6), value=np.float64(10.0), - obs_norm_buffer_data={ + obs_norm_buffer_data={ # pyrefly: ignore[bad-argument-type] 'n': 10, 'mean': np.ones(7), 'unnorm_var': 2 * np.ones(7), @@ -44,7 +44,7 @@ def test_merge(self): metrics={'extra_metric': np.float64(3.0)}, ) merged_result = worker_util.merge_eval_results([result1, result2]) - mean_value = np.mean([result1.value, result2.value]) + mean_value = np.mean([result1.value, result2.value]) # pyrefly: ignore[no-matching-overload] buffer_data_mean = 10 * np.ones(7) / 15.0 self.assertEqual(merged_result.value, mean_value) self.assertIsNotNone(merged_result.obs_norm_buffer_data) @@ -52,7 +52,7 @@ def test_merge(self): np.testing.assert_array_equal( merged_result.obs_norm_buffer_data['mean'], buffer_data_mean ) - self.assertEqual(merged_result.metrics['extra_metric'], 2.0) + self.assertEqual(merged_result.metrics['extra_metric'], 2.0) # pyrefly: ignore[unsupported-operation] def test_merge_empty(self): with self.assertRaisesRegex(ValueError, '(?=.*empty)(?=.*merge)'):