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
14 changes: 4 additions & 10 deletions vizier/_src/algorithms/optimizers/eagle_strategy.py
Original file line number Diff line number Diff line change
Expand Up @@ -389,7 +389,6 @@ def __call__(
# This configuration updates all the fireflies in each iteration.
suggestion_batch_size = pool_size
# Use priors to populate Eagle state
# pytype: disable=wrong-arg-types # jnp-type
return VectorizedEagleStrategy(
n_feature_dimensions=feature_dimensions.n_feature_dimensions,
n_feature_dimensions_with_padding=(
Expand All @@ -402,9 +401,8 @@ def __call__(
pool_size=pool_size,
categorical_sizes=jnp.array(feature_dimensions.categorical_sizes),
max_categorical_size=feature_dimensions.max_categorical_size,
dtype=converter._impl.dtype,
dtype=converter._impl.dtype, # pyrefly: ignore[bad-argument-type]
)
# pytype: enable=wrong-arg-types


@struct.dataclass
Expand Down Expand Up @@ -555,15 +553,13 @@ def init_state(
n_parallel=n_parallel,
seed=seed,
)
# pytype: disable=wrong-arg-types # jnp-type
return VectorizedEagleStrategyState(
iterations=jnp.array(0),
features=init_features,
rewards=jnp.ones(self.pool_size) * -jnp.inf,
best_reward=-jnp.inf,
best_reward=-jnp.inf, # pyrefly: ignore[bad-argument-type]
perturbations=jnp.ones(self.pool_size) * self.config.perturbation,
)
# pytype: enable=wrong-arg-types

def _populate_pool_with_prior_trials(
self,
Expand Down Expand Up @@ -847,14 +843,12 @@ def _create_features(
if self.config.mutate_normalization_type == MutateNormalizationType.MEAN:
# Divide the push / pull forces by the number of participating fireflies.
# Also multiply by normalization_scale.
# pytype: disable=wrong-arg-types # jnp-type
norm_scaled_pulls = self.config.normalization_scale * jnp.nan_to_num(
scaled_pulls / jnp.sum(scaled_pulls > 0.0, axis=1, keepdims=True), 0
scaled_pulls / jnp.sum(scaled_pulls > 0.0, axis=1, keepdims=True), 0 # pyrefly: ignore[bad-argument-type]
)
norm_scaled_push = self.config.normalization_scale * jnp.nan_to_num(
scaled_push / jnp.sum(scaled_push < 0.0, axis=1, keepdims=True), 0
scaled_push / jnp.sum(scaled_push < 0.0, axis=1, keepdims=True), 0 # pyrefly: ignore[bad-argument-type]
)
# pytype: enable=wrong-arg-types
elif self.config.mutate_normalization_type == (
MutateNormalizationType.RANDOM
):
Expand Down
14 changes: 7 additions & 7 deletions vizier/_src/algorithms/optimizers/eagle_strategy_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -548,8 +548,8 @@ def test_optimize_with_eagle_padding(self):
converter = converters.TrialToModelInputConverter.from_problem(
problem,
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,
),
)
eagle_factory = eagle_strategy.VectorizedEagleStrategyFactory()
Expand All @@ -575,8 +575,8 @@ def test_singleton_constraints_are_respected_with_padding(self):
converter = converters.TrialToModelInputConverter.from_problem(
problem,
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,
),
)
eagle_factory = eagle_strategy.VectorizedEagleStrategyFactory()
Expand Down Expand Up @@ -618,8 +618,8 @@ def test_compute_feature_dimensions_from_converter(
converter = converters.TrialToModelInputConverter.from_problem(
problem,
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,
),
)
feature_dimensions = (
Expand Down Expand Up @@ -685,7 +685,7 @@ def score_fn(x, _):
)

padding_schedule = padding.PaddingSchedule(
num_trials=padding.PaddingType.MULTIPLES_OF_10, # pyrefly: ignore[unexpected-keyword]
num_trials=padding.PaddingType.MULTIPLES_OF_10,
)
padding_converter = converters.TrialToModelInputConverter.from_problem(
problem,
Expand Down
2 changes: 1 addition & 1 deletion vizier/_src/algorithms/optimizers/vectorized_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -313,7 +313,7 @@ class VectorizedOptimizer(Generic[_S]):
max_evaluations: int = struct.field(pytree_node=False, default=75_000)
dtype: types.ContinuousAndCategorical[jnp.dtype] = struct.field(
pytree_node=False,
default=types.ContinuousAndCategorical[jnp.dtype]( # pytype: disable=wrong-arg-types # jnp-type
default=types.ContinuousAndCategorical[jnp.dtype](
jnp.float64, types.INT_DTYPE # pyrefly: ignore[bad-argument-type]
),
)
Expand Down
88 changes: 44 additions & 44 deletions vizier/_src/benchmarks/analyzers/convergence_curve.py
Original file line number Diff line number Diff line change
Expand Up @@ -811,16 +811,16 @@ def score(self) -> float:
extended_compared = ConvergenceCurve.extrapolate_ys(
compared_curve, extend_steps
)
baseline_comparator = LogEfficiencyConvergenceCurveComparator( # pyrefly: ignore[missing-argument]
baseline_curve=combined_curve, # pyrefly: ignore[unexpected-keyword]
compared_curve=extended_baseline, # pyrefly: ignore[unexpected-keyword]
compared_quantile=self._baseline_quantile, # pyrefly: ignore[unexpected-keyword]
baseline_comparator = LogEfficiencyConvergenceCurveComparator(
baseline_curve=combined_curve,
compared_curve=extended_baseline,
compared_quantile=self._baseline_quantile,
)
efficiency_baseline = baseline_comparator.curve()
compared_comparator = LogEfficiencyConvergenceCurveComparator( # pyrefly: ignore[missing-argument]
baseline_curve=combined_curve, # pyrefly: ignore[unexpected-keyword]
compared_curve=extended_compared, # pyrefly: ignore[unexpected-keyword]
compared_quantile=self._compared_quantile, # pyrefly: ignore[unexpected-keyword]
compared_comparator = LogEfficiencyConvergenceCurveComparator(
baseline_curve=combined_curve,
compared_curve=extended_compared,
compared_quantile=self._compared_quantile,
)
efficiency_compared = compared_comparator.curve()

Expand Down Expand Up @@ -878,7 +878,7 @@ def _compute_directional_score(
(len(compared) - i) / len(compared) if i != float('inf') else 0
for i in convergence_curve
]
return np.mean(pct_baseline_compared) # pyrefly: ignore[bad-return]
return np.mean(pct_baseline_compared)

def score(self) -> float:
"""Computes the percentage better score.
Expand Down Expand Up @@ -1003,13 +1003,13 @@ def __call__( # pyrefly: ignore[bad-override]
compared_quantile: float = 0.5,
steps_cutoff: Optional[int] = None,
) -> ConvergenceComparator:
return OptimalityGapWinRateComparator( # pyrefly: ignore[missing-argument]
baseline_curve=baseline_curve, # pyrefly: ignore[unexpected-keyword]
compared_curve=compared_curve, # pyrefly: ignore[unexpected-keyword]
baseline_quantile=baseline_quantile, # pyrefly: ignore[unexpected-keyword]
compared_quantile=compared_quantile, # pyrefly: ignore[unexpected-keyword]
name='optimality_gap_win_rate', # pyrefly: ignore[unexpected-keyword]
steps_cutoff=steps_cutoff, # pyrefly: ignore[unexpected-keyword]
return OptimalityGapWinRateComparator(
baseline_curve=baseline_curve,
compared_curve=compared_curve,
baseline_quantile=baseline_quantile,
compared_quantile=compared_quantile,
name='optimality_gap_win_rate',
steps_cutoff=steps_cutoff,
)


Expand All @@ -1024,13 +1024,13 @@ def __call__( # pyrefly: ignore[bad-override]
compared_quantile: float = 0.5,
steps_cutoff: Optional[int] = None,
) -> ConvergenceComparator:
return OptimalityGapGainComparator( # pyrefly: ignore[missing-argument]
baseline_curve=baseline_curve, # pyrefly: ignore[unexpected-keyword]
compared_curve=compared_curve, # pyrefly: ignore[unexpected-keyword]
baseline_quantile=baseline_quantile, # pyrefly: ignore[unexpected-keyword]
compared_quantile=compared_quantile, # pyrefly: ignore[unexpected-keyword]
name='optimality_gap_gain', # pyrefly: ignore[unexpected-keyword]
steps_cutoff=steps_cutoff, # pyrefly: ignore[unexpected-keyword]
return OptimalityGapGainComparator(
baseline_curve=baseline_curve,
compared_curve=compared_curve,
baseline_quantile=baseline_quantile,
compared_quantile=compared_quantile,
name='optimality_gap_gain',
steps_cutoff=steps_cutoff,
)


Expand All @@ -1048,14 +1048,14 @@ def __call__( # pyrefly: ignore[bad-override]
compared_quantile: float = 0.5,
steps_cutoff: Optional[int] = None,
) -> ConvergenceComparator:
return WinRateConvergenceCurveComparator( # pyrefly: ignore[missing-argument]
baseline_curve=baseline_curve, # pyrefly: ignore[unexpected-keyword]
compared_curve=compared_curve, # pyrefly: ignore[unexpected-keyword]
baseline_quantile=baseline_quantile, # pyrefly: ignore[unexpected-keyword]
compared_quantile=compared_quantile, # pyrefly: ignore[unexpected-keyword]
name='convergence_curve_win_rate', # pyrefly: ignore[unexpected-keyword]
return WinRateConvergenceCurveComparator(
baseline_curve=baseline_curve,
compared_curve=compared_curve,
baseline_quantile=baseline_quantile,
compared_quantile=compared_quantile,
name='convergence_curve_win_rate',
comparison_mode=self.comparison_mode,
steps_cutoff=steps_cutoff, # pyrefly: ignore[unexpected-keyword]
steps_cutoff=steps_cutoff,
)


Expand All @@ -1072,13 +1072,13 @@ def __call__( # pyrefly: ignore[bad-override]
compared_quantile: float = 0.5,
steps_cutoff: Optional[int] = None,
) -> ConvergenceComparator:
return LogEfficiencyConvergenceCurveComparator( # pyrefly: ignore[missing-argument]
baseline_curve=baseline_curve, # pyrefly: ignore[unexpected-keyword]
compared_curve=compared_curve, # pyrefly: ignore[unexpected-keyword]
baseline_quantile=baseline_quantile, # pyrefly: ignore[unexpected-keyword]
compared_quantile=compared_quantile, # pyrefly: ignore[unexpected-keyword]
name='log_eff', # pyrefly: ignore[unexpected-keyword]
steps_cutoff=steps_cutoff, # pyrefly: ignore[unexpected-keyword]
return LogEfficiencyConvergenceCurveComparator(
baseline_curve=baseline_curve,
compared_curve=compared_curve,
baseline_quantile=baseline_quantile,
compared_quantile=compared_quantile,
name='log_eff',
steps_cutoff=steps_cutoff,
)


Expand All @@ -1095,13 +1095,13 @@ def __call__( # pyrefly: ignore[bad-override]
compared_quantile: float = 0.5,
steps_cutoff: Optional[int] = None,
) -> ConvergenceComparator:
return PercentageBetterConvergenceCurveComparator( # pyrefly: ignore[missing-argument]
baseline_curve=baseline_curve, # pyrefly: ignore[unexpected-keyword]
compared_curve=compared_curve, # pyrefly: ignore[unexpected-keyword]
baseline_quantile=baseline_quantile, # pyrefly: ignore[unexpected-keyword]
compared_quantile=compared_quantile, # pyrefly: ignore[unexpected-keyword]
name='pct_better', # pyrefly: ignore[unexpected-keyword]
steps_cutoff=steps_cutoff, # pyrefly: ignore[unexpected-keyword]
return PercentageBetterConvergenceCurveComparator(
baseline_curve=baseline_curve,
compared_curve=compared_curve,
baseline_quantile=baseline_quantile,
compared_quantile=compared_quantile,
name='pct_better',
steps_cutoff=steps_cutoff,
)


Expand Down
Loading
Loading