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
38 changes: 38 additions & 0 deletions src/harv/_optional_deps.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,38 @@
"""On-demand importers for optional dependencies.

harv's plotting and diagnostic layers need packages that are not required to
run a sampler: matplotlib and arviz (which itself imports matplotlib). Importing
either one is expensive, and on a cluster with node-local cache directories
every rank that imports matplotlib rebuilds its font list. Importing them here,
at the point of use, keeps ``import harv`` free of both so a sampling-only rank
never pays for them. See "Operational notes" in ``docs/at-scale.md``.

Each importer takes the name of the calling function, used only to write the
``ImportError`` message, and is private to the package.
"""

__all__ = ("get_arviz", "get_mpl")

from typing import Any


def get_mpl(func_name: str) -> tuple[Any, Any]:
"""Import matplotlib, returning ``(matplotlib, matplotlib.pyplot)``."""
try:
import matplotlib as mpl # noqa: PLC0415 (optional dependency)
import matplotlib.pyplot as plt # noqa: PLC0415
except ImportError as e:
msg = f"matplotlib is required for {func_name}."
raise ImportError(msg) from e
return mpl, plt


def get_arviz(func_name: str) -> tuple[Any, Any]:
"""Import arviz, returning ``(arviz, arviz_base.labels.MapLabeller)``."""
try:
import arviz as az # noqa: PLC0415 (optional dependency)
from arviz_base.labels import MapLabeller # noqa: PLC0415
except ImportError as e:
msg = f"arviz is required for {func_name}."
raise ImportError(msg) from e
return az, MapLabeller
6 changes: 0 additions & 6 deletions src/harv/data/datasets.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,12 +20,6 @@

from harv.custom_types import NAngle, NFloatArray, NTime, NVelocity, ScalarQTime

# Optional dependency:
try:
import matplotlib.pyplot as plt
except ImportError:
plt: Any = None


class AbstractData(eqx.Module):
"""Abstract base class for observational data time series."""
Expand Down
23 changes: 6 additions & 17 deletions src/harv/plot.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,18 +17,12 @@
from unxt import Q, ustrip
from unxt.quantity import AllowValue

from harv._optional_deps import get_mpl
from harv.custom_types import BatchQTime, NQAny, NTime, ScalarQTime
from harv.data import GaiaAstrometryData, RVData, SourceData, SystemData
from harv.models.extensions.multi_survey import MultiSurveyOffset
from harv.samplers import Samples

try:
import matplotlib as mpl
import matplotlib.pyplot as plt
except ImportError:
plt: Any = None


# Default styles:
_DEFAULT_ERRORBAR_STYLE: dict[str, Any] = {
"linestyle": "none",
Expand Down Expand Up @@ -93,6 +87,7 @@ def plot_timeseries_errorbar(
Forwarded to ``ax.errorbar()``, overriding defaults.
"""
if ax is None:
_, plt = get_mpl("plot_timeseries_errorbar")
_, ax = plt.subplots()

if time_unit is None:
Expand Down Expand Up @@ -455,9 +450,7 @@ def plot_rv( # noqa: C901 -- plotting code is inherently complex
>>> ax = plot_rv(samples, rv_data, model=sampler.model) # doctest: +SKIP
>>> ax = plot_rv(samples, rv_data, phase_fold_median=True) # doctest: +SKIP
"""
if plt is None:
msg = "matplotlib is required for plot_rv."
raise ImportError(msg)
mpl, plt = get_mpl("plot_rv")

if model is None:
from harv.models.rv import RVModel as _RVModel # noqa: PLC0415
Expand Down Expand Up @@ -929,7 +922,7 @@ def _curve_for_sample(
return ax


def plot_gaia_sky_orbit( # noqa: C901 -- plotting code is inherently complex
def plot_gaia_sky_orbit(
model: Any,
samples: Samples,
*,
Expand Down Expand Up @@ -990,9 +983,7 @@ def plot_gaia_sky_orbit( # noqa: C901 -- plotting code is inherently complex
ValueError
If *samples* does not contain exactly one posterior sample.
"""
if plt is None:
msg = "matplotlib is required for plot_gaia_sky_orbit."
raise ImportError(msg)
_, plt = get_mpl("plot_gaia_sky_orbit")

if len(samples) != 1:
msg = (
Expand Down Expand Up @@ -1184,9 +1175,7 @@ def plot_gaia_astrometry(
... samples.map_sample(), data=gaia_data, model=sampler.model
... )
"""
if plt is None:
msg = "matplotlib is required for plot_gaia_astrometry."
raise ImportError(msg)
_, plt = get_mpl("plot_gaia_astrometry")

if len(samples) != 1:
msg = (
Expand Down
17 changes: 3 additions & 14 deletions src/harv/samplers/samples.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
import quaxed.numpy as jnp
from unxt import AbstractQuantity, Q, ustrip

from harv._optional_deps import get_arviz
from harv.data.datasets import AbstractData
from harv.kepler import masses
from harv.models.parameterizations._base import AbstractParameterization
Expand All @@ -25,14 +26,6 @@
)
from harv.samplers.conversion import convert_parameterization

try:
import arviz as az
from arviz_base.labels import MapLabeller

HAS_ARVIZ = True
except ImportError:
HAS_ARVIZ = False

__all__ = ("Samples", "pad_and_stack_samples")

# Default minimum evidence effective sample size (logZ_int_ess) below which a
Expand Down Expand Up @@ -1665,9 +1658,7 @@ def to_arviz(
--------
>>> idata = samples.to_arviz(["period", "eccentricity"]) # doctest: +SKIP
"""
if not HAS_ARVIZ:
msg = "arviz is required for to_arviz()."
raise ImportError(msg)
az, _ = get_arviz("to_arviz()")

if params is None:
params = self.keys()
Expand Down Expand Up @@ -1745,9 +1736,7 @@ def plot_corner(
... truths={"period": Q(100, "day"), "eccentricity": 0.3},
... )
"""
if not HAS_ARVIZ:
msg = "arviz is required for corner plots."
raise ImportError(msg)
az, MapLabeller = get_arviz("corner plots")

# Select default parameters based on available params
if params is None:
Expand Down
2 changes: 1 addition & 1 deletion tests/unit/test_samplers/test_samples_methods.py
Original file line number Diff line number Diff line change
Expand Up @@ -1408,7 +1408,7 @@ def _capture_truth_overlay(self, samples, truths, params):
backend="matplotlib",
viz={"plot": SimpleNamespace(values=fake_axes)},
)
with patch("harv.samplers.samples.az.plot_pair", return_value=fake_plot_matrix):
with patch("arviz.plot_pair", return_value=fake_plot_matrix):
result = samples.plot_corner(params=params, truths=truths)
return result, fake_axes

Expand Down