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
4 changes: 2 additions & 2 deletions vizier/_src/algorithms/optimizers/vectorized_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
10 changes: 5 additions & 5 deletions vizier/_src/benchmarks/analyzers/convergence_curve.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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,
Expand All @@ -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,
Expand All @@ -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,
Expand Down
2 changes: 1 addition & 1 deletion vizier/_src/jax/models/gaussian_process_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
2 changes: 1 addition & 1 deletion vizier/_src/jax/models/hebo_gp_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
2 changes: 1 addition & 1 deletion vizier/_src/jax/models/multitask_tuned_gp_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
2 changes: 1 addition & 1 deletion vizier/_src/jax/models/tuned_gp_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down
Loading