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
36 changes: 18 additions & 18 deletions vizier/_src/pyglove/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,7 +72,7 @@ class VizierBackend(pg.tuning.Backend):
# Class-level variables.
default_owner: str = getpass.getuser()
default_study_prefix: Optional[str] = None
tuner_cls: Type[client.VizierTuner] = attrs.field() # pyrefly: ignore[bad-class-definition]
tuner_cls: Type[client.VizierTuner] = attrs.field()

# Instance-level variables.

Expand All @@ -82,40 +82,40 @@ class VizierBackend(pg.tuning.Backend):

# Worker group - workers that belong to the same group will share the same
# Vizier client ID.
_group: Union[None, int, str] = attrs.field() # pyrefly: ignore[bad-class-definition]
_group: Union[None, int, str] = attrs.field()

# Max number of examples to sample.
_num_examples: Optional[int] = attrs.field() # pyrefly: ignore[bad-class-definition]
_num_examples: Optional[int] = attrs.field()

# Prior study IDs for transfer learning.
_prior_study_ids: Optional[Sequence[str]] = attrs.field() # pyrefly: ignore[bad-class-definition]
_prior_study_ids: Optional[Sequence[str]] = attrs.field()

# If True, add the completed trials from prior studies. Otherwise, simply
# warm up the algorithm using these trials without adding them to the current
# study.
_add_prior_trials: bool = attrs.field() # pyrefly: ignore[bad-class-definition]
_add_prior_trials: bool = attrs.field()

#
# Internal states.
#

_tuner: client.VizierTuner # pyrefly: ignore[bad-class-definition]
_dna_spec: pg.DNASpec # pyrefly: ignore[bad-class-definition]
_algorithm: pg.geno.DNAGenerator = attrs.field() # pyrefly: ignore[bad-class-definition]
_early_stopping_policy: Optional[pg.tuning.EarlyStoppingPolicy] = ( # pyrefly: ignore[bad-class-definition]
_tuner: client.VizierTuner
_dna_spec: pg.DNASpec
_algorithm: pg.geno.DNAGenerator = attrs.field()
_early_stopping_policy: Optional[pg.tuning.EarlyStoppingPolicy] = (
attrs.field()
)

_study_owner: str = attrs.field() # pyrefly: ignore[bad-class-definition]
_study_name: ExpandedStudyName = attrs.field() # pyrefly: ignore[bad-class-definition]
_converter: converters.VizierConverter = attrs.field() # pyrefly: ignore[bad-class-definition]
_study_owner: str = attrs.field()
_study_name: ExpandedStudyName = attrs.field()
_converter: converters.VizierConverter = attrs.field()

_study: client_abc.StudyInterface = attrs.field() # pyrefly: ignore[bad-class-definition]
_suggestion_generator: Any = attrs.field() # pyrefly: ignore[bad-class-definition]
_study: client_abc.StudyInterface = attrs.field()
_suggestion_generator: Any = attrs.field()

_run_mode: TunerMode = attrs.field() # pyrefly: ignore[bad-class-definition]
_auto_election_thread: Optional[threading.Thread] = attrs.field() # pyrefly: ignore[bad-class-definition]
_is_active: bool = attrs.field() # pyrefly: ignore[bad-class-definition]
_run_mode: TunerMode = attrs.field()
_auto_election_thread: Optional[threading.Thread] = attrs.field()
_is_active: bool = attrs.field()

def __init__(
self,
Expand Down Expand Up @@ -468,7 +468,7 @@ def _auto_elect_primary_if_needed(self) -> None:
def next(self) -> pg.tuning.Feedback:
"""Gets the next tuning feedback object."""
try:
trial = next(self._suggestion_generator) # pytype: disable=wrong-arg-types
trial = next(self._suggestion_generator)
return core.Feedback(self._study.get_trial(trial.id), self._converter)
except StopIteration as e:
self._is_active = False
Expand Down
2 changes: 1 addition & 1 deletion vizier/_src/pyglove/oss_vizier.py
Original file line number Diff line number Diff line change
Expand Up @@ -234,7 +234,7 @@ def get_group_id(self, group_id: Union[None, int, str] = None) -> str:
elif isinstance(group_id, int):
return f'group:{group_id}'
elif isinstance(group_id, str):
return group_id # pytype: disable=bad-return-type
return group_id

def ping_tuner(self, tuner_id: str) -> bool:
# We treat `tuner_id` as the Pythia endpoint.
Expand Down
6 changes: 3 additions & 3 deletions vizier/_src/pyglove/oss_vizier_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,9 +33,9 @@ class OSSVizierSampleTest(vizier_test_lib.SampleTest):
def setUpClass(cls):
super().setUpClass()
server = vizier_server.DefaultVizierServer(
host=os.uname()[1], # pyrefly: ignore[unexpected-keyword]
database_url=constants.SQL_MEMORY_URL, # pyrefly: ignore[unexpected-keyword]
early_stop_recycle_period=datetime.timedelta(seconds=0.0), # pyrefly: ignore[unexpected-keyword]
host=os.uname()[1],
database_url=constants.SQL_MEMORY_URL,
early_stop_recycle_period=datetime.timedelta(seconds=0.0),
)
logging.info(server.endpoint)
vizier._services.reset_for_testing()
Expand Down
2 changes: 1 addition & 1 deletion vizier/_src/pyglove/performance_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,7 @@ class PerformanceTest(parameterized.TestCase):
def setUpClass(cls):
super().setUpClass()
server = vizier_server.DefaultVizierServer(
host=os.uname()[1], database_url=constants.SQL_MEMORY_URL # pyrefly: ignore[unexpected-keyword]
host=os.uname()[1], database_url=constants.SQL_MEMORY_URL
)
logging.info(server.endpoint)
vizier._services.reset_for_testing()
Expand Down
4 changes: 2 additions & 2 deletions vizier/_src/pyglove/pythia.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,10 +50,10 @@ def __attrs_post_init__(self):
# study (unless the binary crashes). We can therefore use an in-ram cache
# and avoid re-loading the same trials over and over again.
self._suggestion_cache = trial_caches.IdDeduplicatingTrialLoader(
self.supporter, include_intermediate_measurements=False # pyrefly: ignore[unexpected-keyword]
self.supporter, include_intermediate_measurements=False
)
self._stopping_cache = trial_caches.IdDeduplicatingTrialLoader(
self.supporter, include_intermediate_measurements=True # pyrefly: ignore[unexpected-keyword]
self.supporter, include_intermediate_measurements=True
)

@property
Expand Down
4 changes: 2 additions & 2 deletions vizier/_src/pyglove/pythia_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,13 +74,13 @@ def test_stopping_policy(self):

m.should_stop_early.return_value = False
res = policy.early_stop(
pythia.EarlyStopRequest(study_descriptor=descriptor, trial_ids=[4]) # pyrefly: ignore[missing-argument, unexpected-keyword]
pythia.EarlyStopRequest(study_descriptor=descriptor, trial_ids=[4])
)
self.assertFalse(res.decisions[0].should_stop)

m.should_stop_early.return_value = True
res = policy.early_stop(
pythia.EarlyStopRequest(study_descriptor=descriptor, trial_ids=[4]) # pyrefly: ignore[missing-argument, unexpected-keyword]
pythia.EarlyStopRequest(study_descriptor=descriptor, trial_ids=[4])
)
self.assertTrue(res.decisions[0].should_stop)
self.assertEqual(m.should_stop_early.call_count, 5) # 3 + 1 + 1
Expand Down
4 changes: 2 additions & 2 deletions vizier/_src/service/clients_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,7 +80,7 @@ def setUpClass(cls):
logging.info('Test setup started.')
super().setUpClass()
cls._server = vizier_server.DefaultVizierServer(
database_url=constants.SQL_MEMORY_URL # pyrefly: ignore[unexpected-keyword]
database_url=constants.SQL_MEMORY_URL
)
clients.environment_variables.server_endpoint = cls._server.endpoint
logging.info('Test setup finished.')
Expand All @@ -99,7 +99,7 @@ def setUpClass(cls):
logging.info('Test setup started.')
super().setUpClass()
cls._server = vizier_server.DistributedPythiaVizierServer(
database_url=constants.SQL_MEMORY_URL # pyrefly: ignore[unexpected-keyword]
database_url=constants.SQL_MEMORY_URL
)
clients.environment_variables.server_endpoint = cls._server.endpoint
logging.info('Test setup finished.')
Expand Down
2 changes: 1 addition & 1 deletion vizier/_src/service/performance_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ class PerformanceTest(parameterized.TestCase):
def setUpClass(cls):
super().setUpClass()
cls.server = vizier_server.DefaultVizierServer(
database_url=constants.SQL_MEMORY_URL # pyrefly: ignore[unexpected-keyword]
database_url=constants.SQL_MEMORY_URL
)
vizier_client.environment_variables.server_endpoint = cls.server.endpoint

Expand Down
Loading
Loading