Coverage for jetgp/full_ddegp_sparse/ddegp.py: 87%
173 statements
« prev ^ index » next coverage.py v7.10.7, created at 2026-04-10 23:11 -0500
« prev ^ index » next coverage.py v7.10.7, created at 2026-04-10 23:11 -0500
1import numpy as np
2from numpy.linalg import cholesky, solve
3import jetgp.utils as utils
4from jetgp.kernel_funcs.kernel_funcs import KernelFactory, get_oti_module
5from jetgp.full_ddegp_sparse.optimizer import Optimizer
6from jetgp.full_ddegp_sparse import ddegp_utils
7from scipy.linalg import cho_solve, cho_factor, solve_triangular
10class ddegp:
11 """
12 Sparse Cholesky variant of the Directional DEGP model.
14 Adds sparse inverse-Cholesky acceleration for NLML evaluation during
15 hyperparameter optimisation. Prediction always uses the exact dense
16 Cholesky solve.
18 Parameters
19 ----------
20 x_train : ndarray
21 Training input data of shape (n_samples, n_features).
22 y_train : list or ndarray
23 Training targets or list of directional derivatives.
24 n_order : int
25 Maximum derivative order.
26 der_indices : list of lists
27 Derivative multi-indices corresponding to each derivative term.
28 rays : ndarray
29 Array of shape (d, n_rays), where each column is a direction vector.
30 derivative_locations : list of lists
31 Which training points have which derivatives.
32 normalize : bool, default=True
33 Whether to normalize inputs and outputs.
34 sigma_data : float or array-like, optional
35 Observation noise standard deviation or diagonal noise values.
36 kernel : str, default='SE'
37 Kernel type.
38 kernel_type : str, default='anisotropic'
39 Kernel anisotropy.
40 smoothness_parameter : float, optional
41 Smoothness parameter for Matern kernel.
42 rho : float, default=3.0
43 Sparsity radius multiplier.
44 use_supernodes : bool, default=True
45 If True, aggregate columns into supernodes.
46 supernode_lam : float, default=1.5
47 Lambda parameter for supernode merging.
48 """
50 def __init__(self, x_train, y_train, n_order, der_indices, rays,
51 derivative_locations=None, normalize=True, sigma_data=None,
52 kernel="SE", kernel_type="anisotropic", smoothness_parameter=None,
53 rho=3.0, use_supernodes=True, supernode_lam=1.5):
55 if n_order > 0 and derivative_locations is None:
56 import warnings
57 n_derivs = sum(len(order_derivs) for order_derivs in der_indices)
58 n_train = len(x_train)
59 derivative_locations = [[i for i in range(n_train)] for _ in range(n_derivs)]
60 warnings.warn(
61 f"derivative_locations not provided. Assuming all {n_derivs} derivative(s) "
62 f"are available at all {n_train} training point(s).",
63 UserWarning
64 )
65 elif der_indices is None and n_order == 0:
66 der_indices = []
67 derivative_locations = []
69 self.x_train = x_train
70 self.y_train = y_train
71 self.sigma_data = sigma_data
72 self.n_order = n_order
73 self.rays = rays
74 self.n_rays = rays.shape[1]
75 self.dim = x_train.shape[1]
76 self.num_points = x_train.shape[0]
77 self.kernel = kernel
78 self.kernel_type = kernel_type
79 self.der_indices = der_indices
80 self.normalize = normalize
81 self.derivative_locations = derivative_locations
82 self.oti = get_oti_module(self.n_rays, n_order)
84 self.flattened_der_indices = utils.flatten_der_indices(der_indices)
86 if normalize:
87 self.y_train, self.mu_y, self.sigma_y, self.sigmas_x, self.mus_x, sigma_data = \
88 utils.normalize_y_data_directional(
89 x_train, y_train, sigma_data, self.flattened_der_indices)
90 self.rays = utils.normalize_directions(self.sigmas_x, self.rays)
91 self.x_train = utils.normalize_x_data_train(x_train)
92 else:
93 self.x_train = x_train
94 self.y_train = utils.reshape_y_train(y_train)
96 self.powers = utils.build_companion_array(self.n_rays, n_order, der_indices)
97 self.differences_by_dim = ddegp_utils.differences_by_dim_func(
98 self.x_train, self.x_train, self.rays, n_order, self.oti)
100 self.sigma_data = (
101 np.zeros((self.y_train.shape[0], self.y_train.shape[0]))
102 if sigma_data is None else np.diag(sigma_data)
103 )
104 self.sigma_data_sq_diag = (
105 np.zeros(self.y_train.shape[0])
106 if sigma_data is None
107 else np.asarray(sigma_data) ** 2
108 )
110 self.kernel_factory = KernelFactory(
111 dim=self.dim,
112 normalize=self.normalize,
113 n_order=self.n_order,
114 differences_by_dim=self.differences_by_dim,
115 smoothness_parameter=smoothness_parameter,
116 oti_module=self.oti,
117 sparse_diffs=False
118 )
119 self.kernel_func = self.kernel_factory.create_kernel(
120 kernel_name=self.kernel,
121 kernel_type=self.kernel_type
122 )
123 self.bounds = self.kernel_factory.bounds
124 self.n_bases = self.n_rays
125 self.optimizer = Optimizer(self)
127 # Sparse Cholesky setup
128 self.rho = rho
129 self.use_supernodes = use_supernodes
130 self.supernode_lam = supernode_lam
131 self._setup_sparse_cholesky()
133 def _setup_sparse_cholesky(self):
134 """Precompute MMD ordering, fill-distances, and sparsity pattern."""
135 from jetgp.full_ddegp_sparse.sparse_cholesky import (
136 mmd_ordering, build_sparsity_pattern, build_supernodes,
137 expand_mmd_permutation, expand_sparsity_to_blocks,
138 expand_supernodes_to_blocks,
139 )
140 X = self.x_train
141 self.mmd_P, self.mmd_l = mmd_ordering(X)
142 X_ord = X[self.mmd_P]
143 self.sparse_S = build_sparsity_pattern(X_ord, self.mmd_l, self.rho)
145 self.mmd_P_full, self._phys_to_rows = expand_mmd_permutation(
146 self.mmd_P, self.num_points, self.derivative_locations
147 )
148 self.sparse_S_full = expand_sparsity_to_blocks(self.sparse_S, self._phys_to_rows)
149 self.sparse_S_full_arr = {
150 j: np.asarray(s, dtype=np.intp) for j, s in self.sparse_S_full.items()
151 }
153 N_total = len(self.mmd_P_full)
154 total_nb = sum(len(s) for s in self.sparse_S_full.values())
155 self.sparse_fill_fraction = total_nb / (N_total * N_total)
156 self._use_dense_factor = self.sparse_fill_fraction > 0.25
158 if self.use_supernodes:
159 phys_sns = build_supernodes(
160 X_ord, self.mmd_l, self.sparse_S, lam=self.supernode_lam
161 )
162 self.sparse_supernodes = phys_sns
163 self.sparse_supernodes_full = expand_supernodes_to_blocks(
164 phys_sns, self._phys_to_rows
165 )
166 for sn in self.sparse_supernodes_full:
167 sn['children_arr'] = np.asarray(sn['children'])
168 ch_pos = {c: i for i, c in enumerate(sn['children'])}
169 sn['ch_pos'] = ch_pos
170 sn['parent_positions'] = np.array(
171 [ch_pos[p] for p in sn['parents']]
172 )
173 else:
174 self.sparse_supernodes = None
175 self.sparse_supernodes_full = None
177 def optimize_hyperparameters(self, *args, **kwargs):
178 """Run the optimizer. Returns optimized hyperparameter vector."""
179 self.params = self.optimizer.optimize_hyperparameters(*args, **kwargs)
180 return self.params
182 def predict(self, X_test, params, calc_cov=False, return_deriv=False, derivs_to_predict=None):
183 """
184 Predict posterior mean and optional variance at test points.
185 Uses exact dense Cholesky solve (not sparse approximation).
186 """
187 length_scales = params[:-1]
188 sigma_n = params[-1]
190 if return_deriv:
191 if derivs_to_predict is not None:
192 common_derivs = derivs_to_predict
193 else:
194 common_derivs = self.flattened_der_indices
196 required_order = max(
197 sum(pair[1] for pair in deriv_spec)
198 for deriv_spec in common_derivs
199 )
200 predict_order = max(required_order, self.n_order)
202 if predict_order > self.n_order:
203 predict_oti = get_oti_module(self.n_rays, predict_order)
204 smoothness_param = getattr(self.kernel_factory, 'alpha', None)
205 predict_kernel_factory = KernelFactory(
206 dim=self.dim,
207 normalize=self.normalize,
208 differences_by_dim=self.differences_by_dim,
209 n_order=predict_order,
210 smoothness_parameter=smoothness_param,
211 oti_module=predict_oti,
212 sparse_diffs=False
213 )
214 predict_kernel_func = predict_kernel_factory.create_kernel(
215 kernel_name=self.kernel, kernel_type=self.kernel_type
216 )
217 else:
218 predict_oti = self.oti
219 predict_kernel_func = self.kernel_func
221 self.powers_predict = utils.build_companion_array_predict(
222 self.n_rays, predict_order, common_derivs)
223 else:
224 common_derivs = []
225 self.powers_predict = None
226 predict_order = self.n_order
227 predict_oti = self.oti
228 predict_kernel_func = self.kernel_func
230 _cache_hit = (
231 hasattr(self, '_cached_params')
232 and self._cached_params is not None
233 and np.array_equal(self._cached_params, params)
234 and getattr(self, '_cached_L', None) is not None
235 )
237 if _cache_hit:
238 L = self._cached_L
239 low = self._cached_low
240 alpha = self._cached_alpha
241 self.n_bases_rays = self._cached_n_bases_rays
242 cho_solve_failed = False
243 else:
244 phi_train = self.kernel_func(self.differences_by_dim, length_scales)
245 self.n_bases_rays = phi_train.get_active_bases()[-1]
246 if self.n_order > 0:
247 phi_exp_train = phi_train.get_all_derivs(self.n_bases_rays, 2 * self.n_order)
248 else:
249 phi_exp_train = phi_train.real[np.newaxis, :, :]
251 K = ddegp_utils.rbf_kernel(
252 phi_train, phi_exp_train, self.n_order, self.n_bases_rays,
253 self.flattened_der_indices, self.powers,
254 index=self.derivative_locations
255 )
256 K.flat[::K.shape[0] + 1] += (10 ** sigma_n) ** 2
257 K += self.sigma_data ** 2
259 try:
260 L, low = cho_factor(K, lower=True)
261 alpha = cho_solve((L, low), self.y_train)
262 cho_solve_failed = False
263 except Exception:
264 alpha = np.linalg.solve(K, self.y_train)
265 L, low = None, None
266 cho_solve_failed = True
268 self._cached_L = L
269 self._cached_low = low
270 self._cached_alpha = alpha
271 self._cached_params = params.copy()
272 self._cached_n_bases_rays = self.n_bases_rays
274 if self.normalize:
275 X_test = utils.normalize_x_data_test(X_test, self.sigmas_x, self.mus_x)
277 if return_deriv:
278 derivative_locations_test = [
279 list(range(X_test.shape[0])) for _ in range(len(common_derivs))]
280 else:
281 derivative_locations_test = None
283 diff_x_test_x_train = ddegp_utils.differences_by_dim_func(
284 self.x_train, X_test, self.rays, predict_order, predict_oti, return_deriv=return_deriv)
286 phi_train_test = predict_kernel_func(diff_x_test_x_train, length_scales)
287 if predict_order > 0:
288 if return_deriv:
289 phi_exp_train_test = phi_train_test.get_all_derivs(self.n_bases_rays, 2 * predict_order)
290 else:
291 phi_exp_train_test = phi_train_test.get_all_derivs(self.n_bases_rays, predict_order)
292 else:
293 phi_exp_train_test = phi_train_test.real[np.newaxis, :, :]
294 K_s = ddegp_utils.rbf_kernel_predictions(
295 phi_train_test, phi_exp_train_test, predict_order, self.n_bases_rays,
296 self.flattened_der_indices, self.powers,
297 return_deriv=return_deriv,
298 index=self.derivative_locations,
299 common_derivs=common_derivs,
300 powers_predict=self.powers_predict
301 )
303 f_mean = K_s.T @ alpha
305 if self.normalize:
306 if return_deriv:
307 f_mean = utils.transform_predictions_directional(
308 f_mean, self.mu_y, self.sigma_y, self.sigmas_x,
309 common_derivs, X_test)
310 else:
311 f_mean = self.mu_y + f_mean * self.sigma_y
313 f_mean = f_mean.reshape(-1, 1)
314 n = X_test.shape[0]
315 m = f_mean.shape[0]
316 num_derivs = m // n
317 reshaped_mean = f_mean.reshape(num_derivs, n)
319 if not calc_cov:
320 return reshaped_mean
322 diff_x_test_x_test = ddegp_utils.differences_by_dim_func(
323 X_test, X_test, self.rays, predict_order, predict_oti, return_deriv=return_deriv)
325 phi_test_test = predict_kernel_func(diff_x_test_x_test, length_scales)
326 bases = phi_test_test.get_active_bases()
327 n_bases = bases[-1] if len(bases) > 0 else 0
329 if predict_order > 0:
330 phi_exp_test_test = phi_test_test.get_all_derivs(n_bases, 2 * predict_order)
331 else:
332 phi_exp_test_test = phi_test_test.real[np.newaxis, :, :]
333 K_ss = ddegp_utils.rbf_kernel_predictions(
334 phi_test_test, phi_exp_test_test, predict_order, n_bases,
335 self.flattened_der_indices, self.powers,
336 return_deriv=return_deriv,
337 index=derivative_locations_test,
338 common_derivs=common_derivs,
339 calc_cov=True,
340 powers_predict=self.powers_predict
341 )
343 if cho_solve_failed:
344 f_cov = K_ss - K_s.T @ np.linalg.inv(K) @ K_s
345 else:
346 v = solve_triangular(L, K_s, lower=low)
347 f_cov = K_ss - v.T @ v
349 if self.normalize:
350 if return_deriv:
351 f_var = utils.transform_cov_directional(
352 f_cov, self.sigma_y, self.sigmas_x,
353 common_derivs, X_test)
354 else:
355 f_var = self.sigma_y ** 2 * np.diag(np.abs(f_cov))
356 else:
357 f_var = np.diag(np.abs(f_cov))
359 reshaped_var = f_var.reshape(num_derivs, n)
360 return reshaped_mean, reshaped_var