Код: Выделить всё
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")
Код: Выделить всё
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
Подробнее здесь: https://stackoverflow.com/questions/790 ... inygp-work