Как работает полиномиальное ядро ​​в tinygp? ⇐ Python

Программы на Python
Anonymous
Как работает полиномиальное ядро ​​в tinygp?

Сообщение Anonymous »

Я пытаюсь научиться использовать пакет tinygp (v 0.3.0) (Python версии 3.11.10 в macOS Sonoma 14.5), но столкнулся с проблемой с линейным ядром. Я следую одному из их руководств, и это код:

Код: Выделить всё

import numpy as np
import matplotlib.pyplot as plt
from statsmodels.datasets import co2

data = co2.load_pandas().data
t = 2000 + (np.array(data.index.to_julian_date()) - 2451545.0) / 365.25
y = np.array(data.co2)
m = np.isfinite(t) & np.isfinite(y) & (t < 1996)
t, y = t[m][::4], y[m][::4]

plt.plot(t, y, ".k")
plt.xlim(t.min(), t.max())
plt.xlabel("year")
_ = plt.ylabel("CO$_2$ in ppm")

import jax
import jax.numpy as jnp

from tinygp import kernels, GaussianProcess

jax.config.update("jax_enable_x64", True)

def build_gp(theta, X):
# We want most of our parameters to be positive so we take the `exp` here
# Note that we're using `jnp` instead of `np`
amps = jnp.exp(theta["log_amps"])
scales = jnp.exp(theta["log_scales"])

# Construct the kernel by multiplying and adding `Kernel` objects
k1 = amps[0] * kernels.ExpSquared(scales[0])

k2 = (amps[1] * kernels.ExpSquared(scales[1]) * kernels.ExpSineSquared(scale=jnp.exp(theta["log_period"]), gamma=jnp.exp(theta["log_gamma"])))

k3 = amps[2] * kernels.RationalQuadratic(alpha=jnp.exp(theta["log_alpha"]), scale=scales[2])

k4 = amps[3] * kernels.ExpSquared(scales[3])

kernel = k1 + k2 + k3 + k4

return GaussianProcess(
kernel, X, diag=jnp.exp(theta["log_diag"]), mean=theta["mean"])

def neg_log_likelihood(theta, X, y):
gp = build_gp(theta, X)
return -gp.log_probability(y)

theta_init = {
"mean": np.float64(340.0),
"log_diag": np.log(0.19),
"log_amps": np.log([66.0, 2.4, 0.66, 0.18, 2.0]),
"log_scales": np.log([67.0, 90.0, 0.78, 1.6, 3.0]),
"log_period": np.float64(0.0),
"log_gamma": np.log(4.3),
"log_alpha": np.log(1.2),
"sigma": jnp.array(1.0)
}

# `jax` can be used to differentiate functions, and also note that we're calling
# `jax.jit` for the best performance.
obj = jax.jit(jax.value_and_grad(neg_log_likelihood))

print(f"Initial negative log likelihood: {obj(theta_init, t, y)[0]}")
print(f"Gradient of the negative log likelihood, wrt the parameters:\n{obj(theta_init, t, y)[1]}")

import jaxopt

solver = jaxopt.ScipyMinimize(fun=neg_log_likelihood)
soln = solver.run(theta_init, X=t, y=y)
print(f"Final negative log likelihood: {soln.state.fun_val}")

t_final = 2055
x = np.linspace(max(t), t_final, 2000)
gp = build_gp(soln.params, t)
cond_gp = gp.condition(y, x).gp
mu, var = cond_gp.loc, cond_gp.variance

plt.plot(t, y, ".k")
plt.fill_between(x, mu + np.sqrt(var), mu - np.sqrt(var), color="C0", alpha=0.5)
plt.plot(x, mu, color="C0", lw=2)

plt.xlim(t.min(), t_final)
plt.xlabel("year")
_ = plt.ylabel("CO$_2$ in ppm")
На практике определяют 4 ядра с соответствующими параметрами, суммируют их и продолжают использовать эту сумму. Все идет нормально. Теперь я хотел попробовать добавить еще одно ядро, чтобы проверить все это, и какое бы ядро ​​я ни добавил, все продолжает работать нормально. За исключением полиномиального ядра (см. здесь). Итак, я добавляю следующую строку

Код: Выделить всё

k5 = 1.0 * kernels.Polynomial(scale=1.0, order = 2, sigma=1.0)
а затем подведите итоги остальным. На этом этапе, если я запущу код, я получу следующую ошибку:

Код: Выделить всё

ValueError
Traceback (most recent call last)
Cell In[6], line 56
52 # `jax` can be used to differentiate functions, and also note that we're calling
53 # `jax.jit` for the best performance.
54 obj = jax.jit(jax.value_and_grad(neg_log_likelihood))
---> 56 print(f"Initial negative log likelihood: {obj(theta_init, t, y)[0]}")
57 print(f"Gradient of the negative log likelihood, wrt the parameters:\n{obj(theta_init, t, y)[1]}")

[...  skipping hidden 20 frame]

Cell In[6], line 37, in neg_log_likelihood(theta, X, y)
36 def neg_log_likelihood(theta, X, y):
---> 37     gp = build_gp(theta, X)
38     return -gp.log_probability(y)

Cell In[6], line 32, in build_gp(theta, X)
28 print("Shape of X:", X.shape)
30 kernel = k1 + k2 + k3 + k4 + k5
---> 32 return GaussianProcess(
33     kernel, X, diag=jnp.exp(theta["log_diag"]), mean=theta["mean"])

[... skipping hidden 3 frame]

File /opt/anaconda3/envs/testgp/lib/python3.11/site-packages/tinygp/gp.py:107, in GaussianProcess.__init__(self, kernel, X, diag, noise, mean, solver, mean_value, covariance_value, **solver_kwargs)
105     else:
106         solver = DirectSolver
--> 107 self.solver = solver(
108     kernel,
109     self.X,
110     self.noise,
111     covariance=covariance_value,
112     **solver_kwargs,
113 )

[... skipping hidden 3 frame]

File /opt/anaconda3/envs/testgp/lib/python3.11/site-packages/tinygp/solvers/direct.py:49, in DirectSolver.__init__(self, kernel, X, noise, covariance)
38 """Build a :class:`DirectSolver` for a given kernel and coordinates
39
40 Args:
(...)
46         and adding ``diag``, but that is not checked.
47 """
48 self.X = X
---> 49 self.variance_value = kernel(X) + noise.diagonal()
50 if covariance is None:
51     covariance = kernel(X, X) + noise

File /opt/anaconda3/envs/testgp/lib/python3.11/site-packages/tinygp/kernels/base.py:86, in Kernel.__call__(self, X1, X2)
84 def __call__(self, X1: JAXArray, X2: JAXArray | None = None) -> JAXArray:
85     if X2 is None:
---> 86         k = jax.vmap(self.evaluate_diag, in_axes=0)(X1)
87         if k.ndim != 1:
88             raise ValueError(
89                 "Invalid kernel diagonal shape: "
90                 f"expected ndim = 1, got ndim={k.ndim} "
91                 "check the dimensions of parameters and custom kernels"
92             )

[... skipping hidden 4 frame]

File /opt/anaconda3/envs/testgp/lib/python3.11/site-packages/tinygp/kernels/base.py:66, in Kernel.evaluate_diag(self, X)
59 def evaluate_diag(self, X: JAXArray) -> JAXArray:
60     """Evaluate the kernel on its diagonal
61
62     The default implementation simply calls :func:`Kernel.evaluate` with
63     ``X`` as both arguments, but subclasses can use this to make diagonal
64     calcuations more efficient.
65     """
---> 66     return self.evaluate(X, X)

[... skipping hidden 1 frame]

File /opt/anaconda3/envs/testgp/lib/python3.11/site-packages/tinygp/kernels/base.py:177, in Sum.evaluate(self, X1, X2)
176 def evaluate(self, X1: JAXArray, X2: JAXArray) -> JAXArray:
--> 177     return self.kernel1.evaluate(X1, X2) + self.kernel2.evaluate(X1, X2)

[... skipping hidden 1 frame]

File /opt/anaconda3/envs/testgp/lib/python3.11/site-packages/tinygp/kernels/base.py:187, in Product.evaluate(self, X1, X2)
186 def evaluate(self, X1: JAXArray, X2: JAXArray) -> JAXArray:
--> 187     return self.kernel1.evaluate(X1, X2) * self.kernel2.evaluate(X1, X2)

[... skipping hidden 1 frame]

File /opt/anaconda3/envs/testgp/lib/python3.11/site-packages/tinygp/kernels/base.py:248, in Polynomial.evaluate(self, X1, X2)
246 def evaluate(self, X1: JAXArray, X2: JAXArray) -> JAXArray:
247     return (
--> 248         (X1 / self.scale) @ (X2 / self.scale) + jnp.square(self.sigma)
249     ) ** self.order

File /opt/anaconda3/envs/testgp/lib/python3.11/site-packages/jax/_src/numpy/array_methods.py:743, in _forward_operator_to_aval..op(self, *args)
742 def op(self, *args):
--> 743   return getattr(self.aval, f"_{name}")(self, *args)

File /opt/anaconda3/envs/testgp/lib/python3.11/site-packages/jax/_src/numpy/array_methods.py:271, in _defer_to_unrecognized_arg..deferring_binary_op(self, other)
269 args = (other, self) if swap else (self, other)
270 if isinstance(other, _accepted_binop_types):
--> 271   return binary_op(*args)
272 # Note: don't use isinstance here, because we don't want to raise for
273 # subclasses, e.g. NamedTuple objects that may override operators.
274 if type(other) in _rejected_binop_types:

[...  skipping hidden 12 frame]

File /opt/anaconda3/envs/testgp/lib/python3.11/site-packages/jax/_src/numpy/lax_numpy.py:3308, in matmul(a, b, precision, preferred_element_type)
3305   if ndim(x) < 1:
3306     msg = (f"matmul input operand {i} must have ndim at least 1, "
3307            f"but it has ndim {ndim(x)}")
-> 3308     raise ValueError(msg)
3309 if preferred_element_type is None:
3310   preferred_element_type, output_weak_type = dtypes.result_type(a, b, return_weak_type_flag=True)

ValueError: matmul input operand 0 must have ndim at least 1, but it has ndim 0
Но я не понимаю, что делаю не так. Если я определяю k5, но не суммирую его с остальными, ошибок не будет, поэтому я думаю, что синтаксис должен быть правильным. Я пытался изменить параметры, но ничего не помогло.

Подробнее здесь: https://stackoverflow.com/questions/790 ... inygp-work

Вернуться в «Python»