Skip to content

Commit e91584d

Browse files
authored
Merge pull request #151 from PyAutoLabs/feature/lazy-heavy-imports
refactor: lazy astropy + workspace version-warning dedupe
2 parents 14887c3 + c6a25ad commit e91584d

2 files changed

Lines changed: 23 additions & 9 deletions

File tree

autonerves/fitsable.py

Lines changed: 7 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -9,11 +9,6 @@
99
except ImportError:
1010
pass
1111

12-
try:
13-
from astropy.io import fits
14-
except ImportError:
15-
pass
16-
1712
import numpy as np
1813
from pathlib import Path
1914
from typing import Dict, Optional, Union, List
@@ -58,8 +53,10 @@ def hdu_list_for_output_from(
5853
ext_name_list=["data", "noise_map"]
5954
)
6055
"""
56+
from astropy.io import fits
57+
6158
hdu_list = []
62-
59+
6360
header = fits.Header()
6461

6562
if header_dict is not None:
@@ -207,6 +204,8 @@ def ndarray_via_fits_from(
207204
--------
208205
array_2d = ndarray_via_fits_from(file_path='/path/to/file/filename.fits', hdu=0)
209206
"""
207+
from astropy.io import fits
208+
210209
with fits.open(
211210
file_path, do_not_scale_image_data=do_not_scale_image_data
212211
) as hdu_list:
@@ -235,6 +234,8 @@ def header_obj_from(file_path: Union[Path, str], hdu: int) -> Dict:
235234
--------
236235
array_2d = ndarray_via_fits_from(file_path='/path/to/file/filename.fits', hdu=0)
237236
"""
237+
from astropy.io import fits
238+
238239
with fits.open(file_path) as hdu_list:
239240
return hdu_list[hdu].header
240241

autonerves/workspace.py

Lines changed: 16 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,19 @@ class WorkspaceVersionMismatchError(exc.ConfigException):
1717
# enough to suggest the clone is genuinely stale.
1818
_STALENESS_WINDOW_DAYS = 30
1919

20+
# Every library init (autofit, autogalaxy, autolens, ...) calls check_version,
21+
# so without dedup a byte-identical warning prints once per library in the
22+
# import chain. Python's own warning registry does not dedupe here because
23+
# third-party imports between the calls invalidate it.
24+
_warned_messages = set()
25+
26+
27+
def _warn_once(message):
28+
if message in _warned_messages:
29+
return
30+
_warned_messages.add(message)
31+
warnings.warn(message)
32+
2033

2134
def _read_general_yaml(workspace_root):
2235
"""
@@ -224,7 +237,7 @@ def check_version(library_version, workspace_root=None):
224237
if floor_version is None or floor_version == "":
225238
if _is_source_checkout(root):
226239
return
227-
warnings.warn(_missing_version_warning(root, library_version))
240+
_warn_once(_missing_version_warning(root, library_version))
228241
return
229242

230243
if floor_version == library_version:
@@ -234,7 +247,7 @@ def check_version(library_version, workspace_root=None):
234247
library_parsed = _parse_version(library_version)
235248

236249
if floor_parsed is None or library_parsed is None:
237-
warnings.warn(
250+
_warn_once(
238251
_unparseable_mismatch_message(floor_version, library_version, root)
239252
)
240253
return
@@ -252,4 +265,4 @@ def check_version(library_version, workspace_root=None):
252265
and library_date is not None
253266
and (library_date - floor_date).days > _STALENESS_WINDOW_DAYS
254267
):
255-
warnings.warn(_stale_workspace_message(floor_version, library_version, root))
268+
_warn_once(_stale_workspace_message(floor_version, library_version, root))

0 commit comments

Comments
 (0)