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
6 changes: 3 additions & 3 deletions vizier/_src/algorithms/designers/bocs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -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))
Expand Down
2 changes: 1 addition & 1 deletion vizier/_src/algorithms/designers/gp_bandit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
128 changes: 64 additions & 64 deletions vizier/_src/algorithms/designers/gp_bandit_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
]
),
Expand All @@ -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


Expand Down Expand Up @@ -138,24 +138,24 @@ 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(
iters=5,
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,
),
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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,
),
),
)
Expand All @@ -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(
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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,
),
)

Expand Down Expand Up @@ -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(
Expand All @@ -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
),
Expand All @@ -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,
Expand Down Expand Up @@ -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
),
)
),
Expand Down
10 changes: 5 additions & 5 deletions vizier/_src/algorithms/designers/gp_ucb_pe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
)


Expand Down Expand Up @@ -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,
)
Expand All @@ -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,
)
Expand Down
Loading
Loading