diff --git a/vizier/_src/algorithms/designers/bocs.py b/vizier/_src/algorithms/designers/bocs.py index e343b1071..91de2b3cb 100644 --- a/vizier/_src/algorithms/designers/bocs.py +++ b/vizier/_src/algorithms/designers/bocs.py @@ -318,7 +318,7 @@ def surrogate_model(self, x: np.ndarray) -> FloatType: if self._X_inf.shape[0] != 0 and np.equal(x, self._X_inf).all(axis=1).any(): barrier = np.inf - return np.dot(x_all, self._alpha) + barrier # pyrefly: ignore[bad-argument-type, no-matching-overload] + return np.dot(x_all, self._alpha) + barrier # pyrefly: ignore[no-matching-overload] def _order_effects(self, X: np.ndarray) -> np.ndarray: """Function computes data matrix for all coupling.""" @@ -498,9 +498,9 @@ def argmin(self) -> np.ndarray: # Extract vectors and compute Cholesky. try: - L = np.linalg.cholesky(X.value) # pyrefly: ignore[no-matching-overload] + L = np.linalg.cholesky(X.value) except np.linalg.LinAlgError: - XpI = X.value + 1e-15 * np.eye(self._num_vars + 1) # pyrefly: ignore[unsupported-operation] + XpI = X.value + 1e-15 * np.eye(self._num_vars + 1) L = np.linalg.cholesky(XpI) suggest_vect = np.zeros((self._num_vars, self._num_repeats)) diff --git a/vizier/_src/algorithms/designers/gp_bandit.py b/vizier/_src/algorithms/designers/gp_bandit.py index b250c7115..504b13dbd 100644 --- a/vizier/_src/algorithms/designers/gp_bandit.py +++ b/vizier/_src/algorithms/designers/gp_bandit.py @@ -638,4 +638,4 @@ def from_problem( if seed is None else jax.random.PRNGKey(seed) ) - return cls(problem, rng=rng, **kwargs) # pyrefly: ignore[unexpected-keyword] + return cls(problem, rng=rng, **kwargs) diff --git a/vizier/_src/algorithms/designers/gp_bandit_test.py b/vizier/_src/algorithms/designers/gp_bandit_test.py index 46f938070..46d414fb1 100644 --- a/vizier/_src/algorithms/designers/gp_bandit_test.py +++ b/vizier/_src/algorithms/designers/gp_bandit_test.py @@ -82,7 +82,7 @@ def _setup_lambda_search( problem = vz.ProblemStatement( search_space=search_space, metric_information=vz.MetricsConfig( - metrics=[ # pyrefly: ignore[unexpected-keyword] + metrics=[ vz.MetricInformation('obj', goal=vz.ObjectiveMetricGoal.MAXIMIZE), ] ), @@ -99,7 +99,7 @@ def _setup_lambda_search( trial.complete(vz.Measurement(metrics={'obj': f(x)})) # pyrefly: ignore[bad-argument-type] obs_trials.append(trial) - gp_designer = gp_bandit.VizierGPBandit(problem, ard_optimizer=ard_optimizer) # pyrefly: ignore[unexpected-keyword] + gp_designer = gp_bandit.VizierGPBandit(problem, ard_optimizer=ard_optimizer) return gp_designer, obs_trials, problem @@ -138,8 +138,8 @@ class GoogleGpBanditTest(parameterized.TestCase): batch_size=5, num_seed_trials=5, padding_schedule=padding.PaddingSchedule( - num_trials=padding.PaddingType.MULTIPLES_OF_10, # pyrefly: ignore[unexpected-keyword] - num_features=padding.PaddingType.POWERS_OF_2, # pyrefly: ignore[unexpected-keyword] + num_trials=padding.PaddingType.MULTIPLES_OF_10, + num_features=padding.PaddingType.POWERS_OF_2, ), ), dict( @@ -147,15 +147,15 @@ class GoogleGpBanditTest(parameterized.TestCase): batch_size=1, num_seed_trials=3, padding_schedule=padding.PaddingSchedule( - num_trials=padding.PaddingType.POWERS_OF_2, # pyrefly: ignore[unexpected-keyword] - num_features=padding.PaddingType.POWERS_OF_2, # pyrefly: ignore[unexpected-keyword] + num_trials=padding.PaddingType.POWERS_OF_2, + num_features=padding.PaddingType.POWERS_OF_2, ), acquisition_optimizer_factory=lbfgsb_optimizer_factory, ), dict( padding_schedule=padding.PaddingSchedule( - num_trials=padding.PaddingType.POWERS_OF_2, # pyrefly: ignore[unexpected-keyword] - num_features=padding.PaddingType.POWERS_OF_2, # pyrefly: ignore[unexpected-keyword] + num_trials=padding.PaddingType.POWERS_OF_2, + num_features=padding.PaddingType.POWERS_OF_2, ), ensemble_size=3, ), @@ -184,18 +184,18 @@ def test_on_flat_continuous_space( ) ) - designer = gp_bandit.VizierGPBandit( # pyrefly: ignore[missing-argument] - problem=problem, # pyrefly: ignore[unexpected-keyword] - acquisition_optimizer_factory=acquisition_optimizer_factory, # pyrefly: ignore[bad-argument-type, unexpected-keyword] - ard_optimizer=optimizers.JaxoptLbfgsB( # pyrefly: ignore[bad-argument-type, unexpected-keyword] + designer = gp_bandit.VizierGPBandit( + problem=problem, + acquisition_optimizer_factory=acquisition_optimizer_factory, # pyrefly: ignore[bad-argument-type] + ard_optimizer=optimizers.JaxoptLbfgsB( optimizers.LbfgsBOptions(maxiter=5, num_line_search_steps=5) ), - num_seed_trials=num_seed_trials, # pyrefly: ignore[unexpected-keyword] - ensemble_size=ensemble_size, # pyrefly: ignore[unexpected-keyword] - padding_schedule=padding_schedule, # pyrefly: ignore[unexpected-keyword] - use_trust_region=use_trust_region, # pyrefly: ignore[unexpected-keyword] - rng=jax.random.PRNGKey(0), # pyrefly: ignore[unexpected-keyword] - linear_coef=0.1, # pyrefly: ignore[unexpected-keyword] + num_seed_trials=num_seed_trials, + ensemble_size=ensemble_size, + padding_schedule=padding_schedule, + use_trust_region=use_trust_region, + rng=jax.random.PRNGKey(0), + linear_coef=0.1, ) with profiler.collect_events() as events: self.assertLen( @@ -241,8 +241,8 @@ def test_on_flat_continuous_space( batch_size=5, num_seed_trials=5, padding_schedule=padding.PaddingSchedule( - num_trials=padding.PaddingType.MULTIPLES_OF_10, # pyrefly: ignore[unexpected-keyword] - num_features=padding.PaddingType.POWERS_OF_2, # pyrefly: ignore[unexpected-keyword] + num_trials=padding.PaddingType.MULTIPLES_OF_10, + num_features=padding.PaddingType.POWERS_OF_2, ), ), ) @@ -260,12 +260,12 @@ def test_on_flat_mixed_space( name='metric', goal=vz.ObjectiveMetricGoal.MAXIMIZE ) ) - designer = gp_bandit.VizierGPBandit( # pyrefly: ignore[missing-argument] - problem=problem, # pyrefly: ignore[unexpected-keyword] - acquisition_optimizer_factory=vectorized_optimizer_factory, # pyrefly: ignore[unexpected-keyword] - num_seed_trials=num_seed_trials, # pyrefly: ignore[unexpected-keyword] - padding_schedule=padding_schedule, # pyrefly: ignore[unexpected-keyword] - use_trust_region=use_trust_region, # pyrefly: ignore[unexpected-keyword] + designer = gp_bandit.VizierGPBandit( + problem=problem, + acquisition_optimizer_factory=vectorized_optimizer_factory, + num_seed_trials=num_seed_trials, + padding_schedule=padding_schedule, + use_trust_region=use_trust_region, ) self.assertLen( test_runners.RandomMetricsRunner( @@ -318,22 +318,22 @@ def test_invariance_to_trials_padding_on_flat_mixed_space( # on ARD again. noop_ard_optimizer = optimizers.default_optimizer(maxiter=0) desinger_rng = jax.random.PRNGKey(0) - designer = gp_bandit.VizierGPBandit( # pyrefly: ignore[missing-argument] - problem=problem, # pyrefly: ignore[unexpected-keyword] - acquisition_optimizer_factory=acquisition_optimizer_factory, # pyrefly: ignore[unexpected-keyword] - ard_optimizer=noop_ard_optimizer, # pyrefly: ignore[unexpected-keyword] - num_seed_trials=num_seed_trials, # pyrefly: ignore[unexpected-keyword] - rng=desinger_rng, # pyrefly: ignore[unexpected-keyword] + designer = gp_bandit.VizierGPBandit( + problem=problem, + acquisition_optimizer_factory=acquisition_optimizer_factory, + ard_optimizer=noop_ard_optimizer, + num_seed_trials=num_seed_trials, + rng=desinger_rng, ) - padding_designer = gp_bandit.VizierGPBandit( # pyrefly: ignore[missing-argument] - problem=problem, # pyrefly: ignore[unexpected-keyword] - acquisition_optimizer_factory=acquisition_optimizer_factory, # pyrefly: ignore[unexpected-keyword] - ard_optimizer=noop_ard_optimizer, # pyrefly: ignore[unexpected-keyword] - num_seed_trials=num_seed_trials, # pyrefly: ignore[unexpected-keyword] - padding_schedule=padding.PaddingSchedule( # pyrefly: ignore[unexpected-keyword] - num_trials=padding.PaddingType.MULTIPLES_OF_10, # pyrefly: ignore[unexpected-keyword] + padding_designer = gp_bandit.VizierGPBandit( + problem=problem, + acquisition_optimizer_factory=acquisition_optimizer_factory, + ard_optimizer=noop_ard_optimizer, + num_seed_trials=num_seed_trials, + padding_schedule=padding.PaddingSchedule( + num_trials=padding.PaddingType.MULTIPLES_OF_10, ), - rng=desinger_rng, # pyrefly: ignore[unexpected-keyword] + rng=desinger_rng, ) metrics_runner_seed = 1 designer_suggestions = test_runners.RandomMetricsRunner( @@ -392,14 +392,14 @@ def test_jit_once(self, *args): ) def create_designer(problem): - return gp_bandit.VizierGPBandit( # pyrefly: ignore[missing-argument] - problem=problem, # pyrefly: ignore[unexpected-keyword] - acquisition_optimizer_factory=vectorized_optimizer_factory, # pyrefly: ignore[unexpected-keyword] - num_seed_trials=3, # pyrefly: ignore[unexpected-keyword] - ensemble_size=2, # pyrefly: ignore[unexpected-keyword] - padding_schedule=padding.PaddingSchedule( # pyrefly: ignore[unexpected-keyword] - num_trials=padding.PaddingType.MULTIPLES_OF_10, # pyrefly: ignore[unexpected-keyword] - num_features=padding.PaddingType.MULTIPLES_OF_10, # pyrefly: ignore[unexpected-keyword] + return gp_bandit.VizierGPBandit( + problem=problem, + acquisition_optimizer_factory=vectorized_optimizer_factory, + num_seed_trials=3, + ensemble_size=2, + padding_schedule=padding.PaddingSchedule( + num_trials=padding.PaddingType.MULTIPLES_OF_10, + num_features=padding.PaddingType.MULTIPLES_OF_10, ), ) @@ -445,19 +445,19 @@ def _qei_factory(data: types.ModelData) -> acquisitions.AcquisitionFunction: n_parallel = 4 iters = 3 - designer = gp_bandit.VizierGPBandit( # pyrefly: ignore[missing-argument] - problem=problem, # pyrefly: ignore[unexpected-keyword] - acquisition_optimizer_factory=vectorized_optimizer_factory, # pyrefly: ignore[unexpected-keyword] - ard_optimizer=optimizers.JaxoptLbfgsB( # pyrefly: ignore[bad-argument-type, unexpected-keyword] + designer = gp_bandit.VizierGPBandit( + problem=problem, + acquisition_optimizer_factory=vectorized_optimizer_factory, + ard_optimizer=optimizers.JaxoptLbfgsB( optimizers.LbfgsBOptions(maxiter=5, num_line_search_steps=5) ), - scoring_function_factory=scoring_fn_factory, # pyrefly: ignore[unexpected-keyword] - scoring_function_is_parallel=True, # pyrefly: ignore[unexpected-keyword] - use_trust_region=False, # pyrefly: ignore[unexpected-keyword] - num_seed_trials=n_parallel, # pyrefly: ignore[unexpected-keyword] - ensemble_size=3, # pyrefly: ignore[unexpected-keyword] - rng=jax.random.PRNGKey(0), # pyrefly: ignore[unexpected-keyword] - linear_coef=0.1, # pyrefly: ignore[unexpected-keyword] + scoring_function_factory=scoring_fn_factory, + scoring_function_is_parallel=True, + use_trust_region=False, + num_seed_trials=n_parallel, + ensemble_size=3, + rng=jax.random.PRNGKey(0), + linear_coef=0.1, ) self.assertLen( test_runners.RandomMetricsRunner( @@ -483,7 +483,7 @@ def test_multi_metrics(self, multitask_type: mt_type): problem = vz.ProblemStatement( search_space=search_space, metric_information=vz.MetricsConfig( - metrics=[ # pyrefly: ignore[unexpected-keyword] + metrics=[ vz.MetricInformation( 'obj1', goal=vz.ObjectiveMetricGoal.MAXIMIZE ), @@ -495,7 +495,7 @@ def test_multi_metrics(self, multitask_type: mt_type): ) iters = 2 - designer = gp_bandit.VizierGPBandit(problem, multitask_type=multitask_type) # pyrefly: ignore[unexpected-keyword] + designer = gp_bandit.VizierGPBandit(problem, multitask_type=multitask_type) self.assertLen( test_runners.RandomMetricsRunner( problem, @@ -528,9 +528,9 @@ def test_convergence( # pylint: disable=g-long-lambda lambda problem, seed: gp_bandit.VizierGPBandit( # pyrefly: ignore[bad-argument-type] problem, - rng=jax.random.PRNGKey(seed), # pyrefly: ignore[unexpected-keyword] - padding_schedule=padding.PaddingSchedule( # pyrefly: ignore[unexpected-keyword] - num_trials=padding.PaddingType.MULTIPLES_OF_10 # pyrefly: ignore[unexpected-keyword] + rng=jax.random.PRNGKey(seed), + padding_schedule=padding.PaddingSchedule( + num_trials=padding.PaddingType.MULTIPLES_OF_10 ), ) ), diff --git a/vizier/_src/algorithms/designers/gp_ucb_pe.py b/vizier/_src/algorithms/designers/gp_ucb_pe.py index dd6802ccf..4e6d5b1b6 100644 --- a/vizier/_src/algorithms/designers/gp_ucb_pe.py +++ b/vizier/_src/algorithms/designers/gp_ucb_pe.py @@ -600,13 +600,13 @@ def score_with_aux( def default_ard_optimizer() -> optimizers.Optimizer[types.ParameterDict]: - return optimizers.JaxoptScipyLbfgsB( # pyrefly: ignore[bad-return] - options=optimizers.LbfgsBOptions( # pyrefly: ignore[unexpected-keyword] + return optimizers.JaxoptScipyLbfgsB( + options=optimizers.LbfgsBOptions( num_line_search_steps=20, tol=1e-5, maxiter=500, ), - max_duration=datetime.timedelta(minutes=40), # pyrefly: ignore[unexpected-keyword] + max_duration=datetime.timedelta(minutes=40), ) @@ -660,7 +660,7 @@ class method that takes `ModelInput` and returns a kw_only=True, factory=lambda: VizierGPUCBPEBandit.default_acquisition_optimizer_factory, ) - _gp_model_class: GPModelClass = attr.field( # pyrefly: ignore[bad-assignment] + _gp_model_class: GPModelClass = attr.field( kw_only=True, factory=lambda: tuned_gp_models.VizierGaussianProcess, ) @@ -677,7 +677,7 @@ class method that takes `ModelInput` and returns a _ard_random_restarts: int = attr.field(default=4, kw_only=True) _use_trust_region: bool = attr.field(default=True, kw_only=True) _num_seed_trials: int = attr.field(default=1, kw_only=True) - _config: UCBPEConfig = attr.field( # pyrefly: ignore[bad-assignment] + _config: UCBPEConfig = attr.field( factory=UCBPEConfig, # pyrefly: ignore[bad-assignment] kw_only=True, ) diff --git a/vizier/_src/algorithms/designers/gp_ucb_pe_test.py b/vizier/_src/algorithms/designers/gp_ucb_pe_test.py index ae0beb91a..5f484dfa5 100644 --- a/vizier/_src/algorithms/designers/gp_ucb_pe_test.py +++ b/vizier/_src/algorithms/designers/gp_ucb_pe_test.py @@ -182,11 +182,11 @@ def test_on_flat_space( ) designer = gp_ucb_pe.VizierGPUCBPEBandit( problem, - acquisition_optimizer_factory=vectorized_optimizer_factory, # pyrefly: ignore[unexpected-keyword] - num_seed_trials=num_seed_trials, # pyrefly: ignore[unexpected-keyword] - ard_optimizer=ard_optimizer, # pyrefly: ignore[bad-argument-type, unexpected-keyword] - metadata_ns='gp_ucb_pe_bandit_test', # pyrefly: ignore[unexpected-keyword] - config=gp_ucb_pe.UCBPEConfig( # pyrefly: ignore[bad-argument-type, unexpected-keyword] + acquisition_optimizer_factory=vectorized_optimizer_factory, + num_seed_trials=num_seed_trials, + ard_optimizer=ard_optimizer, # pyrefly: ignore[bad-argument-type] + metadata_ns='gp_ucb_pe_bandit_test', + config=gp_ucb_pe.UCBPEConfig( # pyrefly: ignore[bad-argument-type] ucb_coefficient=10.0, explore_region_ucb_coefficient=0.5, # Sets the penalty coefficient to 0.0 so that the PE aquisition @@ -208,14 +208,14 @@ def test_on_flat_space( ), multitask_type=multitask_type, ), - ensemble_size=ensemble_size, # pyrefly: ignore[unexpected-keyword] - padding_schedule=padding.PaddingSchedule( # pyrefly: ignore[unexpected-keyword] - num_trials=padding.PaddingType.MULTIPLES_OF_10 # pyrefly: ignore[unexpected-keyword] + ensemble_size=ensemble_size, + padding_schedule=padding.PaddingSchedule( + num_trials=padding.PaddingType.MULTIPLES_OF_10 if applies_padding else padding.PaddingType.NONE, ), - rng=jax.random.PRNGKey(1), # pyrefly: ignore[unexpected-keyword] - mixes_linear_kernel=mixes_linear_kernel, # pyrefly: ignore[unexpected-keyword] + rng=jax.random.PRNGKey(1), + mixes_linear_kernel=mixes_linear_kernel, ) quasi_random_sampler = quasi_random.QuasiRandomDesigner( @@ -396,10 +396,10 @@ def test_ucb_overwrite(self): ) designer = gp_ucb_pe.VizierGPUCBPEBandit( problem, - acquisition_optimizer_factory=vectorized_optimizer_factory, # pyrefly: ignore[unexpected-keyword] - metadata_ns='gp_ucb_pe_bandit_test', # pyrefly: ignore[unexpected-keyword] - num_seed_trials=1, # pyrefly: ignore[unexpected-keyword] - config=gp_ucb_pe.UCBPEConfig( # pyrefly: ignore[bad-argument-type, unexpected-keyword] + acquisition_optimizer_factory=vectorized_optimizer_factory, + metadata_ns='gp_ucb_pe_bandit_test', + num_seed_trials=1, + config=gp_ucb_pe.UCBPEConfig( # pyrefly: ignore[bad-argument-type] ucb_coefficient=10.0, explore_region_ucb_coefficient=0.5, cb_violation_penalty_coefficient=10.0, @@ -407,10 +407,10 @@ def test_ucb_overwrite(self): pe_overwrite_probability=0.0, signal_to_noise_threshold=0.0, ), - padding_schedule=padding.PaddingSchedule( # pyrefly: ignore[unexpected-keyword] - num_trials=padding.PaddingType.MULTIPLES_OF_10 # pyrefly: ignore[unexpected-keyword] + padding_schedule=padding.PaddingSchedule( + num_trials=padding.PaddingType.MULTIPLES_OF_10 ), - rng=jax.random.PRNGKey(1), # pyrefly: ignore[unexpected-keyword] + rng=jax.random.PRNGKey(1), ) trial_id = 1 @@ -494,10 +494,10 @@ def dummy_prior_acquisition(xs: types.ModelInput): designer = gp_ucb_pe.VizierGPUCBPEBandit( problem, - acquisition_optimizer_factory=vectorized_optimizer_factory, # pyrefly: ignore[unexpected-keyword] - metadata_ns='gp_ucb_pe_bandit_test', # pyrefly: ignore[unexpected-keyword] - num_seed_trials=1, # pyrefly: ignore[unexpected-keyword] - config=gp_ucb_pe.UCBPEConfig( # pyrefly: ignore[bad-argument-type, unexpected-keyword] + acquisition_optimizer_factory=vectorized_optimizer_factory, + metadata_ns='gp_ucb_pe_bandit_test', + num_seed_trials=1, + config=gp_ucb_pe.UCBPEConfig( # pyrefly: ignore[bad-argument-type] ucb_coefficient=10.0, explore_region_ucb_coefficient=0.5, cb_violation_penalty_coefficient=10.0, @@ -508,11 +508,11 @@ def dummy_prior_acquisition(xs: types.ModelInput): optimize_set_acquisition_for_exploration ), ), - padding_schedule=padding.PaddingSchedule( # pyrefly: ignore[unexpected-keyword] - num_trials=padding.PaddingType.MULTIPLES_OF_10 # pyrefly: ignore[unexpected-keyword] + padding_schedule=padding.PaddingSchedule( + num_trials=padding.PaddingType.MULTIPLES_OF_10 ), - prior_acquisition=dummy_prior_acquisition, # pyrefly: ignore[unexpected-keyword] - rng=jax.random.PRNGKey(1), # pyrefly: ignore[unexpected-keyword] + prior_acquisition=dummy_prior_acquisition, + rng=jax.random.PRNGKey(1), ) trial_id = 1 @@ -631,13 +631,13 @@ def test_discrete_parameters_are_explored( ) designer = gp_ucb_pe.VizierGPUCBPEBandit( problem, - acquisition_optimizer_factory=vectorized_optimizer_factory, # pyrefly: ignore[unexpected-keyword] - metadata_ns='gp_ucb_pe_bandit_test', # pyrefly: ignore[unexpected-keyword] - num_seed_trials=1, # pyrefly: ignore[unexpected-keyword] - padding_schedule=padding.PaddingSchedule( # pyrefly: ignore[unexpected-keyword] - num_trials=padding.PaddingType.MULTIPLES_OF_10 # pyrefly: ignore[unexpected-keyword] + acquisition_optimizer_factory=vectorized_optimizer_factory, + metadata_ns='gp_ucb_pe_bandit_test', + num_seed_trials=1, + padding_schedule=padding.PaddingSchedule( + num_trials=padding.PaddingType.MULTIPLES_OF_10 ), - rng=jax.random.PRNGKey(1), # pyrefly: ignore[unexpected-keyword] + rng=jax.random.PRNGKey(1), ) all_trials = [] trial_id = 1 diff --git a/vizier/_src/algorithms/designers/grid.py b/vizier/_src/algorithms/designers/grid.py index ec56b0a65..ce6e48d85 100644 --- a/vizier/_src/algorithms/designers/grid.py +++ b/vizier/_src/algorithms/designers/grid.py @@ -171,7 +171,7 @@ def _grid_points_from_parameter_config( parameter_config, scale=True ) grid_scalars = np.linspace(0.0, 1.0, num=self._double_grid_resolution) - return converter.to_parameter_values(grid_scalars) # pytype:disable=bad-return-type + return converter.to_parameter_values(grid_scalars) # pyrefly: ignore[bad-return] elif parameter_config.type == pyvizier.ParameterType.INTEGER: min_value, max_value = parameter_config.bounds diff --git a/vizier/_src/algorithms/designers/grid_test.py b/vizier/_src/algorithms/designers/grid_test.py index 0d222e661..680238afd 100644 --- a/vizier/_src/algorithms/designers/grid_test.py +++ b/vizier/_src/algorithms/designers/grid_test.py @@ -128,8 +128,8 @@ def test_policy_wrapping(self, shuffle_seed): # Make sure we covered entire search space. all_suggestions = [] for _ in range(self.search_space_size): - request = pythia.SuggestRequest( # pyrefly: ignore[missing-argument] - study_descriptor=policy_supporter.study_descriptor(), count=1 # pyrefly: ignore[unexpected-keyword] + request = pythia.SuggestRequest( + study_descriptor=policy_supporter.study_descriptor(), count=1 ) decisions = policy.suggest(request) all_suggestions.extend(decisions.suggestions) diff --git a/vizier/_src/algorithms/designers/quasi_random_test.py b/vizier/_src/algorithms/designers/quasi_random_test.py index b25967ced..004070f6b 100644 --- a/vizier/_src/algorithms/designers/quasi_random_test.py +++ b/vizier/_src/algorithms/designers/quasi_random_test.py @@ -117,8 +117,8 @@ def test_policy_wrapping(self): # Make sure outputs are distinct. all_suggestions = [] for _ in range(1000): - request = pythia.SuggestRequest( # pyrefly: ignore[missing-argument] - study_descriptor=policy_supporter.study_descriptor(), count=1 # pyrefly: ignore[unexpected-keyword] + request = pythia.SuggestRequest( + study_descriptor=policy_supporter.study_descriptor(), count=1 ) decisions = policy.suggest(request) all_suggestions.extend(decisions.suggestions) diff --git a/vizier/_src/algorithms/designers/scheduled_designer_test.py b/vizier/_src/algorithms/designers/scheduled_designer_test.py index b0dda3c31..4eea5d415 100644 --- a/vizier/_src/algorithms/designers/scheduled_designer_test.py +++ b/vizier/_src/algorithms/designers/scheduled_designer_test.py @@ -83,16 +83,16 @@ def test_schedule_designer(self): param2 = scheduled_designer.ExponentialScheduledParam( init_value=1.5, final_value=20.1, rate=1.7 ) - mock_scheduled_designer = scheduled_designer.ScheduledDesigner( # pyrefly: ignore[missing-argument] + mock_scheduled_designer = scheduled_designer.ScheduledDesigner( problem, - designer_factory=MockParameterizedDesigner, # pyrefly: ignore[unexpected-keyword] - designer_state_updater=DirectDesignerStateUpdater(), # pyrefly: ignore[unexpected-keyword] - scheduled_params={"parameter1": param1, "parameter2": param2}, # pyrefly: ignore[unexpected-keyword] - expected_total_num_trials=expected_total_num_trials, # pyrefly: ignore[unexpected-keyword] + designer_factory=MockParameterizedDesigner, + designer_state_updater=DirectDesignerStateUpdater(), + scheduled_params={"parameter1": param1, "parameter2": param2}, + expected_total_num_trials=expected_total_num_trials, ) # Check initial values. - self.assertEqual(mock_scheduled_designer.designer.parameter1, 10.5) # pytype: disable=attribute-error - self.assertEqual(mock_scheduled_designer.designer.parameter2, 1.5) # pytype: disable=attribute-error + self.assertEqual(mock_scheduled_designer.designer.parameter1, 10.5) # pyrefly: ignore[missing-attribute] + self.assertEqual(mock_scheduled_designer.designer.parameter2, 1.5) # pyrefly: ignore[missing-attribute] # Check suggestions. self.assertLen( test_runners.RandomMetricsRunner( @@ -106,8 +106,8 @@ def test_schedule_designer(self): expected_total_num_trials, ) # Check final values. - self.assertAlmostEqual(mock_scheduled_designer.designer.parameter1, 2.1) # pytype: disable=attribute-error - self.assertAlmostEqual(mock_scheduled_designer.designer.parameter2, 20.1) # pytype: disable=attribute-error + self.assertAlmostEqual(mock_scheduled_designer.designer.parameter1, 2.1) # pyrefly: ignore[no-matching-overload] + self.assertAlmostEqual(mock_scheduled_designer.designer.parameter2, 20.1) # pyrefly: ignore[no-matching-overload] def test_validate_suggested_num_trials(self): # Test that updating the designer with trials changes the state. @@ -125,12 +125,12 @@ def test_validate_suggested_num_trials(self): param2 = scheduled_designer.ExponentialScheduledParam( init_value=1.5, final_value=20.1, rate=1.7 ) - mock_scheduled_designer = scheduled_designer.ScheduledDesigner( # pyrefly: ignore[missing-argument] + mock_scheduled_designer = scheduled_designer.ScheduledDesigner( problem, - designer_factory=MockParameterizedDesigner, # pyrefly: ignore[unexpected-keyword] - designer_state_updater=DirectDesignerStateUpdater(), # pyrefly: ignore[unexpected-keyword] - scheduled_params={"parameter1": param1, "parameter2": param2}, # pyrefly: ignore[unexpected-keyword] - expected_total_num_trials=10, # pyrefly: ignore[unexpected-keyword] + designer_factory=MockParameterizedDesigner, + designer_state_updater=DirectDesignerStateUpdater(), + scheduled_params={"parameter1": param1, "parameter2": param2}, + expected_total_num_trials=10, ) # Generate active and completed trials. active_trials = test_studies.flat_continuous_space_with_scaling_trials(2) @@ -165,12 +165,12 @@ def test_scheduled_designer_serialization(self): param2 = scheduled_designer.ExponentialScheduledParam( init_value=1.5, final_value=20.1, rate=1.7 ) - mock_scheduled_designer = scheduled_designer.ScheduledDesigner( # pyrefly: ignore[missing-argument] + mock_scheduled_designer = scheduled_designer.ScheduledDesigner( problem, - designer_factory=MockParameterizedDesigner, # pyrefly: ignore[unexpected-keyword] - designer_state_updater=DirectDesignerStateUpdater(), # pyrefly: ignore[unexpected-keyword] - scheduled_params={"parameter1": param1, "parameter2": param2}, # pyrefly: ignore[unexpected-keyword] - expected_total_num_trials=expected_total_num_trials, # pyrefly: ignore[unexpected-keyword] + designer_factory=MockParameterizedDesigner, + designer_state_updater=DirectDesignerStateUpdater(), + scheduled_params={"parameter1": param1, "parameter2": param2}, + expected_total_num_trials=expected_total_num_trials, ) # Making several suggestions so the state would change. mock_scheduled_designer.suggest(count=1) @@ -178,12 +178,12 @@ def test_scheduled_designer_serialization(self): # Store the state in metadata. state = mock_scheduled_designer.dump() # Create a new designer and load state. - new_mock_scheduled_designer = scheduled_designer.ScheduledDesigner( # pyrefly: ignore[missing-argument] + new_mock_scheduled_designer = scheduled_designer.ScheduledDesigner( problem, - designer_factory=MockParameterizedDesigner, # pyrefly: ignore[unexpected-keyword] - designer_state_updater=DirectDesignerStateUpdater(), # pyrefly: ignore[unexpected-keyword] - scheduled_params={"parameter1": param1, "parameter2": param2}, # pyrefly: ignore[unexpected-keyword] - expected_total_num_trials=expected_total_num_trials, # pyrefly: ignore[unexpected-keyword] + designer_factory=MockParameterizedDesigner, + designer_state_updater=DirectDesignerStateUpdater(), + scheduled_params={"parameter1": param1, "parameter2": param2}, + expected_total_num_trials=expected_total_num_trials, ) new_mock_scheduled_designer.load(state) self.assertEqual( @@ -206,12 +206,12 @@ def test_scheduled_gp_bandit(self): def _gp_bandit_factory(problem): return gp_bandit.VizierGPBandit(problem) - scheduled_desinger = scheduled_gp_bandit.ScheduledGPBanditFactory( # pyrefly: ignore[missing-argument] - gp_bandit_factory=_gp_bandit_factory, # pyrefly: ignore[unexpected-keyword] - expected_total_num_trials=2, # pyrefly: ignore[unexpected-keyword] - init_ucb_coefficient=4.0, # pyrefly: ignore[unexpected-keyword] - final_ucb_coefficient=1.0, # pyrefly: ignore[unexpected-keyword] - decay_ucb_coefficient=1.2, # pyrefly: ignore[unexpected-keyword] + scheduled_desinger = scheduled_gp_bandit.ScheduledGPBanditFactory( + gp_bandit_factory=_gp_bandit_factory, + expected_total_num_trials=2, + init_ucb_coefficient=4.0, + final_ucb_coefficient=1.0, + decay_ucb_coefficient=1.2, )(problem) self.assertLen( @@ -244,18 +244,18 @@ def _gp_ucb_pe_factory( ) -> gp_ucb_pe.VizierGPUCBPEBandit: return gp_ucb_pe.VizierGPUCBPEBandit(problem) - scheduled_desinger = scheduled_gp_ucb_pe.ScheduledGPUCBPEFactory( # pyrefly: ignore[missing-argument] - gp_ucb_pe_factory=_gp_ucb_pe_factory, # pyrefly: ignore[unexpected-keyword] - expected_total_num_trials=10, # pyrefly: ignore[unexpected-keyword] - init_ucb_coefficient=4.0, # pyrefly: ignore[unexpected-keyword] - final_ucb_coefficient=1.0, # pyrefly: ignore[unexpected-keyword] - decay_ucb_coefficient=1.2, # pyrefly: ignore[unexpected-keyword] - init_explore_region_ucb_coefficient=1.0, # pyrefly: ignore[unexpected-keyword] - final_explore_region_ucb_coefficient=0.5, # pyrefly: ignore[unexpected-keyword] - decay_explore_region_ucb_coefficient=1.2, # pyrefly: ignore[unexpected-keyword] - init_ucb_overwrite_probability=0.25, # pyrefly: ignore[unexpected-keyword] - final_ucb_overwrite_probability=0.0, # pyrefly: ignore[unexpected-keyword] - decay_ucb_overwrite_probability=1.0, # pyrefly: ignore[unexpected-keyword] + scheduled_desinger = scheduled_gp_ucb_pe.ScheduledGPUCBPEFactory( + gp_ucb_pe_factory=_gp_ucb_pe_factory, + expected_total_num_trials=10, + init_ucb_coefficient=4.0, + final_ucb_coefficient=1.0, + decay_ucb_coefficient=1.2, + init_explore_region_ucb_coefficient=1.0, + final_explore_region_ucb_coefficient=0.5, + decay_explore_region_ucb_coefficient=1.2, + init_ucb_overwrite_probability=0.25, + final_ucb_overwrite_probability=0.0, + decay_ucb_overwrite_probability=1.0, )(problem) self.assertLen( diff --git a/vizier/_src/algorithms/designers/scheduled_gp_bandit.py b/vizier/_src/algorithms/designers/scheduled_gp_bandit.py index 77b1c5866..4dfe46f2a 100644 --- a/vizier/_src/algorithms/designers/scheduled_gp_bandit.py +++ b/vizier/_src/algorithms/designers/scheduled_gp_bandit.py @@ -54,10 +54,10 @@ def _gp_bandit_state_updater(designer, params): rate=self._decay_ucb_coefficient, ) - return scheduled_designer.ScheduledDesigner( # pyrefly: ignore[missing-argument] + return scheduled_designer.ScheduledDesigner( problem, - designer_factory=self._gp_bandit_factory, # pyrefly: ignore[bad-argument-type, unexpected-keyword] - designer_state_updater=_gp_bandit_state_updater, # pyrefly: ignore[unexpected-keyword] - scheduled_params={'ucb_coefficient': ucb_coef_param}, # pyrefly: ignore[unexpected-keyword] - expected_total_num_trials=self._expected_total_num_trials, # pyrefly: ignore[unexpected-keyword] + designer_factory=self._gp_bandit_factory, # pyrefly: ignore[bad-argument-type] + designer_state_updater=_gp_bandit_state_updater, + scheduled_params={'ucb_coefficient': ucb_coef_param}, + expected_total_num_trials=self._expected_total_num_trials, ) diff --git a/vizier/_src/algorithms/designers/scheduled_gp_ucb_pe.py b/vizier/_src/algorithms/designers/scheduled_gp_ucb_pe.py index 4835472c9..51528bc51 100644 --- a/vizier/_src/algorithms/designers/scheduled_gp_ucb_pe.py +++ b/vizier/_src/algorithms/designers/scheduled_gp_ucb_pe.py @@ -97,10 +97,10 @@ def _gp_ucb_pe_state_updater(designer, params): 'ucb_overwrite_probability': ucb_overwrite_probability_param, } - return scheduled_designer.ScheduledDesigner( # pyrefly: ignore[missing-argument] + return scheduled_designer.ScheduledDesigner( problem, - designer_factory=self._gp_ucb_pe_factory, # pyrefly: ignore[bad-argument-type, unexpected-keyword] - designer_state_updater=_gp_ucb_pe_state_updater, # pyrefly: ignore[unexpected-keyword] - scheduled_params=scheduled_params, # pyrefly: ignore[bad-argument-type, unexpected-keyword] - expected_total_num_trials=self._expected_total_num_trials, # pyrefly: ignore[unexpected-keyword] + designer_factory=self._gp_ucb_pe_factory, # pyrefly: ignore[bad-argument-type] + designer_state_updater=_gp_ucb_pe_state_updater, + scheduled_params=scheduled_params, # pyrefly: ignore[bad-argument-type] + expected_total_num_trials=self._expected_total_num_trials, )