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
24 changes: 12 additions & 12 deletions iris/algorithms/ars_algorithm_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,12 +41,12 @@ def test_ars_gradient(self, orthogonal_suggestions, quasirandom_suggestions):
suggestions = algo.get_param_suggestions()
self.assertLen(suggestions, 6)
eval_results = [
worker_util.EvaluationResult(np.array([10., 11.]), 10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.empty(0), 0), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 11.]), 10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 9.]), -10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 9.]), -10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.empty(0), 0), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10.0, 11.0]), 10),
worker_util.EvaluationResult(np.empty(0), 0),
worker_util.EvaluationResult(np.array([10.0, 11.0]), 10),
worker_util.EvaluationResult(np.array([10.0, 9.0]), -10),
worker_util.EvaluationResult(np.array([10.0, 9.0]), -10),
worker_util.EvaluationResult(np.empty(0), 0),
]
algo.process_evaluations(eval_results)
np.testing.assert_array_equal(algo._opt_params, np.array([10, 11]))
Expand All @@ -63,12 +63,12 @@ def test_ars_gradient_with_schedule(self):
suggestions = algo.get_param_suggestions()
self.assertLen(suggestions, 6)
eval_results = [
worker_util.EvaluationResult(np.array([10., 11.]), 10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.empty(0), 0), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 11.]), 10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 9.]), -10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 9.]), -10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.empty(0), 0), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10.0, 11.0]), 10),
worker_util.EvaluationResult(np.empty(0), 0),
worker_util.EvaluationResult(np.array([10.0, 11.0]), 10),
worker_util.EvaluationResult(np.array([10.0, 9.0]), -10),
worker_util.EvaluationResult(np.array([10.0, 9.0]), -10),
worker_util.EvaluationResult(np.empty(0), 0),
]
algo.process_evaluations(eval_results)
np.testing.assert_array_equal(algo._opt_params, np.array([10, 11]))
Expand Down
2 changes: 1 addition & 1 deletion iris/algorithms/cma_algorithm_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,7 @@ def test_cma_optimization(self):
np.array(suggestion['params_to_eval']),
test_fn(np.array(suggestion['params_to_eval']))))
if i%10 == 0:
eval_results[0] = worker_util.EvaluationResult(np.empty(0), 0) # pytype: disable=wrong-arg-types # numpy-scalars
eval_results[0] = worker_util.EvaluationResult(np.empty(0), 0)
self.algo.process_evaluations(eval_results)
np.testing.assert_almost_equal(self.algo._opt_params, _TRUE_OPTIMAL)

Expand Down
6 changes: 3 additions & 3 deletions iris/algorithms/es_enas_algorithm.py
Original file line number Diff line number Diff line change
Expand Up @@ -90,7 +90,7 @@ def proper_unserialize(metadata: str) -> pg.DNA:
return dna

if self._multithreading:
dna_list = self._pool.map(proper_unserialize, eval_metadatas) # pytype:disable=attribute-error
dna_list = self._pool.map(proper_unserialize, eval_metadatas)
else:
dna_list = map(proper_unserialize, eval_metadatas)
dna_list = list(dna_list)
Expand All @@ -108,7 +108,7 @@ def get_param_suggestions(self,
# Note that for faster serialization, DNASpec is removed from DNA.
dna_list = [self._controller.propose_dna() for _ in vanilla_suggestions]
if self._multithreading:
metadata_list = self._pool.map(pg.to_json_str, dna_list) # pytype:disable=attribute-error
metadata_list = self._pool.map(pg.to_json_str, dna_list)
else:
metadata_list = map(pg.to_json_str, dna_list)
metadata_list = list(metadata_list)
Expand All @@ -130,7 +130,7 @@ def _get_state(self) -> Dict[str, Any]:
return vanilla_state

def _set_state(self, new_state: Dict[str, Any]) -> None:
super()._set_state(new_state) # pytype: disable=attribute-error
super()._set_state(new_state)
self._interval_counter = new_state["interval_counter"]
self._dna_spec = pg.from_json_str(new_state["serialized_dna_spec"])
self._controller = self._controller_fn(dna_spec=self._dna_spec)
Expand Down
3 changes: 1 addition & 2 deletions iris/algorithms/es_enas_algorithm_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,6 @@
# See the License for the specific language governing permissions and
# limitations under the License.

# pytype: disable=attribute-error
from gym import spaces
from iris.algorithms import es_enas_algorithm
from iris.policies import nas_policy
Expand Down Expand Up @@ -44,7 +43,7 @@ def make_evaluation_results(suggestion_list):
value=np.random.uniform(),
metadata=suggestion['metadata'])
eval_results.append(evaluation_result)
eval_results.append(worker_util.EvaluationResult(np.empty(0), 0)) # pytype: disable=wrong-arg-types # numpy-scalars
eval_results.append(worker_util.EvaluationResult(np.empty(0), 0))
return eval_results


Expand Down
16 changes: 8 additions & 8 deletions iris/algorithms/multi_agent_ars_algorithm_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,42 +49,42 @@ def test_init(self, agent_keys, expected_agent_keys, expected_num_agents):

def _build_evaluation_results(self) -> list[worker_util.EvaluationResult]:
eval_results = [
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(
params_evaluated=np.array([10.0, 11.0, 12.0, 13.0]),
value=10,
metrics={'reward_arm': 10, 'reward_opp': -5},
),
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(
params_evaluated=np.array([10.0, 11.0, 14.0, 15.0]),
value=10,
metrics={'reward_arm': 10, 'reward_opp': -10},
),
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(
params_evaluated=np.empty(0),
value=0,
metrics={'reward_arm': 0, 'reward_opp': 0},
),
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(
params_evaluated=np.array([1.0, 2.0, 3.0, 4.0]),
value=10,
metrics={'reward_arm': 10, 'reward_opp': -10},
),
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(
params_evaluated=np.array([10.0, 11.0, 12.0, 13.0]),
value=-10,
metrics={'reward_arm': -10, 'reward_opp': 5},
),
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(
params_evaluated=np.array([10.0, 11.0, 14.0, 15.0]),
value=-10,
metrics={'reward_arm': -10, 'reward_opp': 10},
),
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(
params_evaluated=np.array([5.0, 6.0, 7.0, 8.0]),
value=-10,
metrics={'reward_arm': -10, 'reward_opp': 10},
),
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(
params_evaluated=np.empty(0),
value=0,
metrics={'reward_arm': 0, 'reward_opp': 0},
Expand Down
2 changes: 1 addition & 1 deletion iris/algorithms/optimizers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]) # pyrefly: ignore[unsupported-operation]
res_list.append(result[i])
return np.array(res_list)


Expand Down
30 changes: 18 additions & 12 deletions iris/algorithms/pes_algorithm_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,18 +37,24 @@ def test_pes_gradient(self, orthogonal_suggestions, quasirandom_suggestions):
init_state = {'init_params': np.array([10., 10.])}
algo.initialize(init_state)
eval_results = [
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
np.array([10., 11.]), 10, metrics={'current_step': 5}),
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
np.empty(0), 0, metrics={'current_step': 5}),
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
np.array([10., 11.]), 10, metrics={'current_step': 5}),
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
np.array([10., 9.]), -10, metrics={'current_step': 5}),
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
np.array([10., 9.]), -10, metrics={'current_step': 5}),
worker_util.EvaluationResult( # pytype: disable=wrong-arg-types # numpy-scalars
np.empty(0), 0, metrics={'current_step': 5}),
worker_util.EvaluationResult(
np.array([10.0, 11.0]), 10, metrics={'current_step': 5}
),
worker_util.EvaluationResult(
np.empty(0), 0, metrics={'current_step': 5}
),
worker_util.EvaluationResult(
np.array([10.0, 11.0]), 10, metrics={'current_step': 5}
),
worker_util.EvaluationResult(
np.array([10.0, 9.0]), -10, metrics={'current_step': 5}
),
worker_util.EvaluationResult(
np.array([10.0, 9.0]), -10, metrics={'current_step': 5}
),
worker_util.EvaluationResult(
np.empty(0), 0, metrics={'current_step': 5}
),
]
algo.process_evaluations(eval_results)
np.testing.assert_array_equal(algo._opt_params, np.array([10, 11]))
Expand Down
20 changes: 10 additions & 10 deletions iris/algorithms/piars_algorithm.py
Original file line number Diff line number Diff line change
Expand Up @@ -207,8 +207,8 @@ def __init__(
else:
self.policy = policy

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]
obs_spec = gym_wrapper.spec_from_gym_space(self._env.observation_space)
action_spec = gym_wrapper.spec_from_gym_space(self._env.action_space)
time_step_spec = ts.time_step_spec(observation_spec=obs_spec)
policy_step_spec = policy_step.PolicyStep(action=action_spec) # pyrefly: ignore[missing-argument]
collect_data_spec = trajectory.from_transition(
Expand Down Expand Up @@ -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) # pyrefly: ignore[not-callable]
s, _ = self.policy.h_model(obs)
outputs = []
for i in range(self._rollout_length):
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]
p, v = self.policy.f_model(s)
u_next, s_next = self.policy.g_model([s, actions[:, i, ...]])
outputs.append((p, v, u_next, s))
s = s_next
p, v = self.policy.f_model(s) # pyrefly: ignore[not-callable]
p, v = self.policy.f_model(s)
outputs.append((p, v, None, s))
return outputs

Expand Down Expand Up @@ -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) # pyrefly: ignore[not-callable]
_, hy_vision = self.policy.h_model(obs_k)
# A trick from https://arxiv.org/abs/2011.10566
hy_vision = tf.stop_gradient(hy_vision)
zx = self.policy.px_model(hx) # pyrefly: ignore[not-callable]
zy = self.policy.py_model(hy_vision) # pyrefly: ignore[not-callable]
zx = self.policy.px_model(hx)
zy = self.policy.py_model(hy_vision)
iyz, _, _ = infonce(zx, zy, temperature=0.1)
loss_pi = -iyz

Expand Down Expand Up @@ -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) # pyrefly: ignore[bad-argument-type]
inner_dims = list(space.shape)
outer_dims = list(x.shape)[: -len(inner_dims)]
x = np.reshape(x, outer_dims + [np.prod(inner_dims)])
return x
Expand Down
8 changes: 4 additions & 4 deletions iris/algorithms/pyglove_algorithm.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,7 +70,7 @@ def proper_unserialize(metadata: str) -> pg.DNA:
return dna

if self._multithreading:
dna_list = self._pool.map(proper_unserialize, eval_metadatas) # pytype:disable=attribute-error
dna_list = self._pool.map(proper_unserialize, eval_metadatas)
else:
dna_list = map(proper_unserialize, eval_metadatas)
dna_list = list(dna_list)
Expand All @@ -89,7 +89,7 @@ def get_param_suggestions(self,
]
# Note that for faster serialization, DNASpec is removed from DNA.
if self._multithreading:
metadata_list = self._pool.map(pg.to_json_str, dna_list) # pytype:disable=attribute-error
metadata_list = self._pool.map(pg.to_json_str, dna_list)
else:
metadata_list = map(pg.to_json_str, dna_list)

Expand All @@ -104,8 +104,8 @@ def get_param_suggestions(self,

def _get_state(self) -> Dict[str, Any]:
vanilla_state = {}
vanilla_state["serialized_dna_spec"] = pg.to_json_str(self._dna_spec) # pytype:disable=attribute-error
vanilla_state["controller_alg_state"] = self._controller.get_state() # pytype:disable=attribute-error
vanilla_state["serialized_dna_spec"] = pg.to_json_str(self._dna_spec)
vanilla_state["controller_alg_state"] = self._controller.get_state()
return vanilla_state

def _set_state(self, new_state: Dict[str, Any]) -> None:
Expand Down
4 changes: 2 additions & 2 deletions iris/algorithms/pyribs_algorithm_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -208,7 +208,7 @@ def test_process_evaluations(self):
worker_util.EvaluationResult(
params_evaluated=np.ones((13,)),
value=1,
obs_norm_buffer_data={ # pyrefly: ignore[bad-argument-type]
obs_norm_buffer_data={
buffer.N: 1, # pyrefly: ignore[bad-assignment]
buffer.STD: np.ones((8,)),
buffer.MEAN: np.ones((8,)),
Expand All @@ -219,7 +219,7 @@ def test_process_evaluations(self):
worker_util.EvaluationResult(
params_evaluated=np.ones((13,) * 2),
value=2,
obs_norm_buffer_data={ # pyrefly: ignore[bad-argument-type]
obs_norm_buffer_data={
buffer.N: 2, # pyrefly: ignore[bad-assignment]
buffer.STD: np.ones((8,)) * 2,
buffer.MEAN: np.ones((8,)) * 2,
Expand Down
32 changes: 16 additions & 16 deletions iris/algorithms/rbo_algorithm_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,14 +32,14 @@ def test_rbo_gradient(self):
init_state = {'init_params': np.array([10., 10.])}
algo.initialize(init_state)
eval_results = [
worker_util.EvaluationResult(np.array([10., 11.]), 10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 9.]), -10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.empty(0), 0), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 11.]), 10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 11.]), 10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 9.]), -10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 9.]), -10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.empty(0), 0), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10.0, 11.0]), 10),
worker_util.EvaluationResult(np.array([10.0, 9.0]), -10),
worker_util.EvaluationResult(np.empty(0), 0),
worker_util.EvaluationResult(np.array([10.0, 11.0]), 10),
worker_util.EvaluationResult(np.array([10.0, 11.0]), 10),
worker_util.EvaluationResult(np.array([10.0, 9.0]), -10),
worker_util.EvaluationResult(np.array([10.0, 9.0]), -10),
worker_util.EvaluationResult(np.empty(0), 0),
]
algo.process_evaluations(eval_results)
np.testing.assert_array_almost_equal(
Expand Down Expand Up @@ -67,14 +67,14 @@ def test_rbo_gradient_2(self, regression_method, orthogonal_suggestions,
init_state = {'init_params': np.array([10., 10.])}
algo.initialize(init_state)
eval_results = [
worker_util.EvaluationResult(np.array([10., 11.]), 10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 9.]), -10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.empty(0), 0), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 11.]), 10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 11.]), 10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 9.]), -10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10., 9.]), -10), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.empty(0), 0), # pytype: disable=wrong-arg-types # numpy-scalars
worker_util.EvaluationResult(np.array([10.0, 11.0]), 10),
worker_util.EvaluationResult(np.array([10.0, 9.0]), -10),
worker_util.EvaluationResult(np.empty(0), 0),
worker_util.EvaluationResult(np.array([10.0, 11.0]), 10),
worker_util.EvaluationResult(np.array([10.0, 11.0]), 10),
worker_util.EvaluationResult(np.array([10.0, 9.0]), -10),
worker_util.EvaluationResult(np.array([10.0, 9.0]), -10),
worker_util.EvaluationResult(np.empty(0), 0),
]
algo.process_evaluations(eval_results)
np.testing.assert_equal(len(algo._opt_params), 2)
Expand Down
4 changes: 2 additions & 2 deletions iris/policies/gym_space_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,8 +38,8 @@ def filter_space(space: gym.Space,
def extend_space(space: gym.Space, key: str, value: gym.Space):
"""Adds new keys or dimensions to the space."""
if isinstance(space, gym.spaces.Box):
low = np.concatenate((space.low, value.low)) # pyrefly: ignore[missing-attribute]
high = np.concatenate((space.high, value.high)) # pyrefly: ignore[missing-attribute]
low = np.concatenate((space.low, value.low))
high = np.concatenate((space.high, value.high))
return gym.spaces.Box(low=low, high=high)
elif isinstance(space, gym.spaces.Dict):
extended_space = dict(space.spaces)
Expand Down
8 changes: 5 additions & 3 deletions iris/policies/implicit_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -438,8 +438,9 @@ def act(self, state: np.ndarray) -> np.ndarray:
else:
phi_state = self._energy.linearized_energy_state(state)
if self._bootstrapped_samples == self._num_samples:
return self._actions[np.argmax(
np.dot(self._lat_reps_for_actions, phi_state))] # pyrefly: ignore[bad-argument-type, no-matching-overload]
return self._actions[
np.argmax(np.dot(self._lat_reps_for_actions, phi_state)) # pyrefly: ignore[no-matching-overload]
]
else:
random_indices = np.random.choice(np.arange(len(self._actions)))
return self._actions[random_indices[np.argmax(
Expand Down Expand Up @@ -522,7 +523,8 @@ def act(self, state: np.ndarray) -> np.ndarray:
base_prefix_sum = self._prefix_sum_table[seg_start_index - 1]
prob = np.dot( # pyrefly: ignore[no-matching-overload]
self._prefix_sum_table[seg_end_index - 1] - base_prefix_sum,
phi_state) # pyrefly: ignore[bad-argument-type]
phi_state,
)
probs.append(prob)
start_end_indices.append([seg_start_index, seg_end_index])
seg_start_index = seg_end_index
Expand Down
Loading
Loading