diff --git a/vizier/_src/algorithms/optimizers/vectorized_base.py b/vizier/_src/algorithms/optimizers/vectorized_base.py index 2ecfe2329..0375b3601 100644 --- a/vizier/_src/algorithms/optimizers/vectorized_base.py +++ b/vizier/_src/algorithms/optimizers/vectorized_base.py @@ -480,7 +480,7 @@ def _optimization_one_step(_, args): self.strategy.init_state( init_seed, n_parallel=parallel_dim, - prior_features=prior_features, + prior_features=prior_features, # pyrefly: ignore[bad-argument-type] prior_rewards=prior_rewards, ), init_best_results, @@ -658,7 +658,7 @@ def trials_to_sorted_array( ) -> Optional[types.ModelInput]: """Sorts trials by the order they were created and converts to array.""" if prior_trials: - prior_trials = sorted(prior_trials, key=lambda x: x.creation_time) + prior_trials = sorted(prior_trials, key=lambda x: x.creation_time) # pyrefly: ignore[no-matching-overload] prior_features = converter.to_features(prior_trials) else: prior_features = None diff --git a/vizier/_src/benchmarks/analyzers/convergence_curve.py b/vizier/_src/benchmarks/analyzers/convergence_curve.py index 40f0e387c..eb50a0ee5 100644 --- a/vizier/_src/benchmarks/analyzers/convergence_curve.py +++ b/vizier/_src/benchmarks/analyzers/convergence_curve.py @@ -995,7 +995,7 @@ def curve(self) -> ConvergenceCurve: class OptimalityGapWinRateComparatorFactory(ConvergenceComparatorFactory): """Factory class for OptimalityGapWinRateComparator.""" - def __call__( + def __call__( # pyrefly: ignore[bad-override] self, baseline_curve: ConvergenceCurve, compared_curve: ConvergenceCurve, @@ -1016,7 +1016,7 @@ def __call__( class OptimalityGapGainComparatorFactory(ConvergenceComparatorFactory): """Factory class for OptimalityGapGainComparator.""" - def __call__( + def __call__( # pyrefly: ignore[bad-override] self, baseline_curve: ConvergenceCurve, compared_curve: ConvergenceCurve, @@ -1040,7 +1040,7 @@ class WinRateConvergenceCurveComparatorFactory(ConvergenceComparatorFactory): comparison_mode: Literal['pairwise', 'quantiles'] = 'pairwise' - def __call__( + def __call__( # pyrefly: ignore[bad-override] self, baseline_curve: ConvergenceCurve, compared_curve: ConvergenceCurve, @@ -1064,7 +1064,7 @@ class LogEfficiencyConvergenceCurveComparatorFactory( ): """Factory class for LogEfficiencyConvergenceCurveComparator.""" - def __call__( + def __call__( # pyrefly: ignore[bad-override] self, baseline_curve: ConvergenceCurve, compared_curve: ConvergenceCurve, @@ -1087,7 +1087,7 @@ class PercentageBetterConvergenceCurveComparatorFactory( ): """Factory class for PercentageBetterConvergenceCurveComparator.""" - def __call__( + def __call__( # pyrefly: ignore[bad-override] self, baseline_curve: ConvergenceCurve, compared_curve: ConvergenceCurve, diff --git a/vizier/_src/jax/models/gaussian_process_model.py b/vizier/_src/jax/models/gaussian_process_model.py index 4555c70b0..473c652a6 100644 --- a/vizier/_src/jax/models/gaussian_process_model.py +++ b/vizier/_src/jax/models/gaussian_process_model.py @@ -97,7 +97,7 @@ def __init__( ' True) to your main' ) - def __call__( + def __call__( # pyrefly: ignore[bad-override] self, inputs: Optional[types.ModelInput] = None ) -> Generator[sp_model.ModelParameter, Array, tfd.GaussianProcess]: # TODO: Remove the following line when the linter bug is fixed. diff --git a/vizier/_src/jax/models/hebo_gp_model.py b/vizier/_src/jax/models/hebo_gp_model.py index 8bb5db8cb..fc4bfac48 100644 --- a/vizier/_src/jax/models/hebo_gp_model.py +++ b/vizier/_src/jax/models/hebo_gp_model.py @@ -51,7 +51,7 @@ def build_model( gp_coroutine = VizierHeboGaussianProcess() return sp.StochasticProcessModel(gp_coroutine) - def __call__( + def __call__( # pyrefly: ignore[bad-override] self, inputs: Optional[types.ModelInput] = None ) -> Generator[sp.ModelParameter, jax.Array, tfd.GaussianProcess]: """Creates a generator. diff --git a/vizier/_src/jax/models/multitask_tuned_gp_models.py b/vizier/_src/jax/models/multitask_tuned_gp_models.py index d6b926753..7b3d84c6e 100644 --- a/vizier/_src/jax/models/multitask_tuned_gp_models.py +++ b/vizier/_src/jax/models/multitask_tuned_gp_models.py @@ -217,7 +217,7 @@ def sample(key: Any) -> jnp.ndarray: return sample # pyrefly: ignore[bad-return] - def __call__( + def __call__( # pyrefly: ignore[bad-override] self, inputs: Optional[types.ModelInput] = None ) -> Generator[ sp.ModelParameter, diff --git a/vizier/_src/jax/models/tuned_gp_models.py b/vizier/_src/jax/models/tuned_gp_models.py index 99006ba9f..1e60c1db4 100644 --- a/vizier/_src/jax/models/tuned_gp_models.py +++ b/vizier/_src/jax/models/tuned_gp_models.py @@ -129,7 +129,7 @@ def build_model( ) return sp.StochasticProcessModel(gp_coroutine) - def __call__( + def __call__( # pyrefly: ignore[bad-override] self, inputs: Optional[types.ModelInput] = None, ) -> Generator[sp.ModelParameter, jax.Array, tfd.GaussianProcess]: