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
22 changes: 6 additions & 16 deletions docs/conf.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,7 +83,7 @@
"examples/_*.py",
"examples/**/_*.py",
# GPJax's notebook helper predates that naming convention and is imported by
# name (`from utils import use_mpl_style`), so it cannot simply be renamed.
# name (`from utils import clean_legend`), so it cannot simply be renamed.
"examples/utils.py",
# Image/style assets, plus one stray legacy module (static/jaxkern/main.py)
# that source_suffix would otherwise read as a notebook.
Expand Down Expand Up @@ -317,21 +317,11 @@
# GitHub Pages used to send.
html_extra_path = ["_redirects", "_headers"]
html_theme_options = {
# `accent_color` only accepts a radix ramp *name*. shibuya writes the value
# verbatim into `<html data-accent-color="...">` and its stylesheet carries
# one `[data-accent-color=<name>]` block per radix ramp, each mapping
# `--accent-1..12` onto that ramp. A hex value matches no block, so the whole
# accent ramp goes undefined -- confirmed in a browser: with
# `data-accent-color="#7a2e2a"` both `--accent-9` and the `--sy-c-link` it
# feeds compute to the empty string, the active sidebar entry drops back to
# body-text grey and code blocks lose their tint entirely.
#
# So the named ramp stays, and supplies the derived tints (code-block and
# admonition surfaces, hover states). `red` replaces the previous `crimson`
# because the brand is now #7a2e2a, a true red at hue 3 degrees; crimson is a
# pink-red at hue 348 and its tints read pink against it. The exact brand hex
# is pinned over `--accent-9` in stylesheets/extra.css.
"accent_color": "red",
# `accent_color` only accepts a radix ramp *name*: a hex value matches none of
# shibuya's `[data-accent-color=<name>]` blocks and leaves `--accent-*`
# undefined. As in impulso, it names the token family, and
# stylesheets/extra.css re-tones the crimson scale to ledger oxblood.
"accent_color": "crimson",
"color_mode": "auto", # follow the reader's light/dark preference
"github_url": "https://github.com/QuantClimate/GPJax",
"nav_links": [
Expand Down
10 changes: 8 additions & 2 deletions docs/examples/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,11 +27,17 @@
# [official documentation](https://docs.kidger.site/equinox/).
#

# %% tags=["remove-cell"]
import logging

# The build host may not have the ledger fonts (Public Sans, Spectral); hide
# matplotlib's font fallback messages.
logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %%
import typing as tp

import equinox as eqx
from utils import use_mpl_style
from gpjax.mean_functions import (
AbstractMeanFunction,
Constant,
Expand Down Expand Up @@ -72,7 +78,7 @@ def glue(*args, **kwargs):


# set the default style for plotting
use_mpl_style()
gpx.plotting.use_style()

cols = mpl.rcParams["axes.prop_cycle"].by_key()["color"]

Expand Down
10 changes: 8 additions & 2 deletions docs/examples/barycentres.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,10 +30,16 @@
# significantly more favourable uncertainty estimation.
#

# %% tags=["remove-cell"]
import logging

# The build host may not have the ledger fonts (Public Sans, Spectral); hide
# matplotlib's font fallback messages.
logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %%
import typing as tp

from utils import use_mpl_style
import jax

# Enable Float64 for more stable matrix inversions.
Expand All @@ -60,7 +66,7 @@ def glue(*args, **kwargs):
key = jr.key(123)

# set the default style for plotting
use_mpl_style()
gpx.plotting.use_style()

cols = plt.rcParams["axes.prop_cycle"].by_key()["color"]

Expand Down
10 changes: 8 additions & 2 deletions docs/examples/classification.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,9 +25,15 @@
# closed form, is covered in the [regression notebook](regression.py); everything that
# follows is a consequence of losing that closed form.

# %% tags=["remove-cell"]
import logging

# The build host may not have the ledger fonts (Public Sans, Spectral); hide
# matplotlib's font fallback messages.
logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %%
import equinox as eqx
from utils import use_mpl_style
from gpjax.linalg import add_jitter, cholesky_factor
from gpjax.parameters import val
import jax
Expand Down Expand Up @@ -57,7 +63,7 @@
identity_matrix = jnp.eye

# set the default style for plotting
use_mpl_style()
gpx.plotting.use_style()

key = jr.key(42)
cols = plt.rcParams["axes.prop_cycle"].by_key()["color"]
Expand Down
10 changes: 8 additions & 2 deletions docs/examples/collapsed_vi.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,9 +32,15 @@
# the uncollapsed bound of the
# [sparse stochastic variational inference notebook](uncollapsed_vi.py).

# %% tags=["remove-cell"]
import logging

# The build host may not have the ledger fonts (Public Sans, Spectral); hide
# matplotlib's font fallback messages.
logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %%
# Enable Float64 for more stable matrix inversions.
from utils import use_mpl_style
from jax import (
config,
jit,
Expand All @@ -54,7 +60,7 @@


# set the default style for plotting
use_mpl_style()
gpx.plotting.use_style()

key = jr.key(42)

Expand Down
10 changes: 8 additions & 2 deletions docs/examples/constructing_new_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,9 +24,15 @@
# to a GP prior; if not, our [introduction to kernels](intro_to_kernels.py) builds that
# intuition first.

# %% tags=["remove-cell"]
import logging

# The build host may not have the ledger fonts (Public Sans, Spectral); hide
# matplotlib's font fallback messages.
logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %%
# Enable Float64 for more stable matrix inversions.
from utils import use_mpl_style
from gpjax.kernels.base import val
from gpjax.kernels.computations import DenseKernelComputation
from gpjax.parameters import PositiveReal
Expand All @@ -49,7 +55,7 @@


# set the default style for plotting
use_mpl_style()
gpx.plotting.use_style()

key = jr.key(42)

Expand Down
10 changes: 8 additions & 2 deletions docs/examples/deep_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,9 +26,15 @@
# {cite:t}`wilson2016deep`, transforming the inputs to our
# Gaussian process model's kernel through a neural network can offer a solution to this.

# %% tags=["remove-cell"]
import logging

# The build host may not have the ledger fonts (Public Sans, Spectral); hide
# matplotlib's font fallback messages.
logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %%
import equinox as eqx
from utils import use_mpl_style
from gpjax.kernels.computations import (
AbstractKernelComputation,
DenseKernelComputation,
Expand Down Expand Up @@ -58,7 +64,7 @@


# set the default style for plotting
use_mpl_style()
gpx.plotting.use_style()
cols = mpl.rcParams["axes.prop_cycle"].by_key()["color"]

key = jr.key(42)
Expand Down
11 changes: 9 additions & 2 deletions docs/examples/dual_svgp.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,13 @@
# denotes the stored sites; $\mathbf{a}_i =
# \mathbf{K}_{zz}^{-1}\mathbf{k}_z(x_i)$.

# %% tags=["remove-cell"]
import logging

# The build host may not have the ledger fonts (Public Sans, Spectral); hide
# matplotlib's font fallback messages.
logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %%
# Enable Float64 for more stable matrix inversions.
import time
Expand All @@ -45,7 +52,7 @@
import matplotlib.pyplot as plt
import optax as ox
import paramax
from utils import clean_legend, use_mpl_style
from utils import clean_legend

config.update("jax_enable_x64", True)

Expand All @@ -67,7 +74,7 @@
key = jr.key(123)

# set the default style for plotting
use_mpl_style()
gpx.plotting.use_style()
cols = mpl.rcParams["axes.prop_cycle"].by_key()["color"]


Expand Down
10 changes: 8 additions & 2 deletions docs/examples/graph_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,10 +25,16 @@
# kernels supported within GPJax, see the
# [kernels notebook](constructing_new_kernels.py).

# %% tags=["remove-cell"]
import logging

# The build host may not have the ledger fonts (Public Sans, Spectral); hide
# matplotlib's font fallback messages.
logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %%
import random

from utils import use_mpl_style

# Enable Float64 for more stable matrix inversions.
from jax import config
Expand All @@ -52,7 +58,7 @@ def glue(*args, **kwargs):


# set the default style for plotting
use_mpl_style()
gpx.plotting.use_style()

key = jr.key(42)

Expand Down
10 changes: 8 additions & 2 deletions docs/examples/heteroscedastic_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,13 @@
# moments of each into an ELBO. For non-Gaussian likelihoods the same structure
# remains; only the expected log-likelihood changes.

# %% tags=["remove-cell"]
import logging

# The build host may not have the ledger fonts (Public Sans, Spectral); hide
# matplotlib's font fallback messages.
logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %% tags=["remove-cell"]
import os

Expand All @@ -60,7 +67,6 @@
import matplotlib.pyplot as plt
import optax as ox

from utils import use_mpl_style
import gpjax as gpx
from gpjax.likelihoods import (
HeteroscedasticGaussian,
Expand All @@ -76,7 +82,7 @@
config.update("jax_enable_x64", True)


use_mpl_style()
gpx.plotting.use_style()
key = jr.key(123)
cols = mpl.rcParams["axes.prop_cycle"].by_key()["color"]

Expand Down
15 changes: 10 additions & 5 deletions docs/examples/intro_to_gps.py
Original file line number Diff line number Diff line change
Expand Up @@ -120,6 +120,13 @@
#
# We can plot three different parameterisations of this density.

# %% tags=["remove-cell"]
import logging

# The build host may not have the ledger fonts (Public Sans, Spectral); hide
# matplotlib's font fallback messages.
logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %%
import warnings

Expand All @@ -131,13 +138,11 @@
import pandas as pd
import seaborn as sns

from utils import (
confidence_ellipse,
use_mpl_style,
)
from gpjax import plotting
from utils import confidence_ellipse

# set the default style for plotting
use_mpl_style()
plotting.use_style()

key = jr.key(42)

Expand Down
10 changes: 8 additions & 2 deletions docs/examples/intro_to_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,11 +21,17 @@
#
# In this guide we provide an introduction to kernels, and the role they play in Gaussian process models.

# %% tags=["remove-cell"]
import logging

# The build host may not have the ledger fonts (Public Sans, Spectral); hide
# matplotlib's font fallback messages.
logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %%
# Enable Float64 for more stable matrix inversions.
from pathlib import Path

from utils import use_mpl_style
from gpjax.typing import Array
from jax import config
import jax.numpy as jnp
Expand All @@ -50,7 +56,7 @@
key = jr.key(42)

# set the default style for plotting
use_mpl_style()
gpx.plotting.use_style()
cols = mpl.rcParams["axes.prop_cycle"].by_key()["color"]

# %% [markdown]
Expand Down
10 changes: 8 additions & 2 deletions docs/examples/likelihoods_guide.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,6 +72,13 @@
# these methods in the forthcoming sections, but first, we will show how to instantiate
# a likelihood object. To do this, we'll need a dataset.

# %% tags=["remove-cell"]
import logging

# The build host may not have the ledger fonts (Public Sans, Spectral); hide
# matplotlib's font fallback messages.
logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %% mystnb={"figure": {"caption": "Fifty noisy observations of a sinusoid, shown alongside the latent function that generated them.", "name": "fig-likelihoods-guide-data"}}
import jax

Expand All @@ -81,14 +88,13 @@
import jax.random as jr
import matplotlib.pyplot as plt

from utils import use_mpl_style
import gpjax as gpx

config.update("jax_enable_x64", True)


# set the default style for plotting
use_mpl_style()
gpx.plotting.use_style()
cols = plt.rcParams["axes.prop_cycle"].by_key()["color"]

key = jr.key(42)
Expand Down
10 changes: 8 additions & 2 deletions docs/examples/multioutput.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,8 +31,14 @@
# coregionalization matrix to see what the model has discovered about the output
# structure.

# %% tags=["remove-cell"]
import logging

# The build host may not have the ledger fonts (Public Sans, Spectral); hide
# matplotlib's font fallback messages.
logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR)

# %%
from utils import use_mpl_style
from jax import config
import jax.numpy as jnp
import jax.random as jr
Expand All @@ -46,7 +52,7 @@
import gpjax as gpx

key = jr.key(42)
use_mpl_style()
gpx.plotting.use_style()
cols = mpl.rcParams["axes.prop_cycle"].by_key()["color"]

# %% [markdown]
Expand Down
Loading
Loading