Coverage for jetgp/full_gddegp_sparse/gddegp.py: 87%
201 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
2import jetgp.utils as utils
3from jetgp.kernel_funcs.kernel_funcs import KernelFactory, get_oti_module
4from jetgp.full_gddegp_sparse.optimizer import Optimizer
5from jetgp.full_gddegp_sparse import gddegp_utils
6from scipy.linalg import cho_solve, cho_factor, solve_triangular
7import warnings
10class gddegp:
11 """
12 Sparse Cholesky variant of the GDDEGP 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 rays_list : list of ndarray
27 List of ray arrays. rays_list[i] has shape (d, len(derivative_locations[i])).
28 der_indices : list of lists
29 Derivative multi-indices corresponding to each derivative term.
30 derivative_locations : list of lists
31 Which training points have which derivatives.
32 n_bases : int, optional
33 Override the OTI space size. By default ``2 * n_direction_types``.
34 normalize : bool, default=True
35 Whether to normalize inputs and outputs.
36 sigma_data : float or array-like, optional
37 Observation noise standard deviation or diagonal noise values.
38 kernel : str, default='SE'
39 Kernel type.
40 kernel_type : str, default='anisotropic'
41 Kernel anisotropy.
42 smoothness_parameter : float, optional
43 Smoothness parameter for Matern kernel.
44 rho : float, default=3.0
45 Sparsity radius multiplier.
46 use_supernodes : bool, default=True
47 If True, aggregate columns into supernodes.
48 supernode_lam : float, default=1.5
49 Lambda parameter for supernode merging.
50 """
52 def __init__(self, x_train, y_train, n_order, rays_list, der_indices,
53 derivative_locations=None, n_bases=None, normalize=True,
54 sigma_data=None, kernel="SE", kernel_type="anisotropic",
55 smoothness_parameter=None,
56 rho=3.0, use_supernodes=True, supernode_lam=1.5):
58 if n_order > 0 and derivative_locations is None:
59 n_derivs = sum(len(order_derivs) for order_derivs in der_indices)
60 n_train = len(x_train)
61 derivative_locations = [[i for i in range(n_train)] for _ in range(n_derivs)]
62 warnings.warn(
63 f"derivative_locations not provided. Assuming all {n_derivs} derivative(s) "
64 f"are available at all {n_train} training point(s).",
65 UserWarning
66 )
68 elif der_indices is None and n_order == 0:
69 der_indices = []
70 derivative_locations = []
72 self.x_train = x_train
73 self.y_train = y_train
74 self.sigma_data = sigma_data
75 self.n_order = n_order
76 self.max_order = n_order
77 self.rays_list = rays_list
78 self.dim = x_train.shape[1]
79 self.num_points = x_train.shape[0]
80 self.kernel = kernel
81 self.kernel_type = kernel_type
82 self.normalize = normalize
83 self.derivative_locations = derivative_locations
84 self.der_indices = der_indices
86 self.flattened_der_indices = utils.flatten_der_indices(der_indices)
87 if n_bases is not None:
88 self.n_bases = n_bases
89 else:
90 self.n_bases = 2 * len(self.flattened_der_indices)
91 self.oti = get_oti_module(self.n_bases, n_order)
93 if normalize:
94 self.y_train, self.mu_y, self.sigma_y, self.sigmas_x, self.mus_x, sigma_data = \
95 utils.normalize_y_data_directional(
96 x_train, y_train, sigma_data, self.flattened_der_indices)
97 self.rays_list = utils.normalize_directions_2(self.sigmas_x, self.rays_list)
98 self.x_train = utils.normalize_x_data_train(x_train)
99 else:
100 self.x_train = x_train
101 self.y_train = utils.reshape_y_train(y_train)
103 self.differences_by_dim = gddegp_utils.differences_by_dim_func(
104 self.x_train, self.x_train,
105 self.rays_list, self.rays_list,
106 self.derivative_locations, self.derivative_locations,
107 n_order, self.oti, return_deriv=True
108 )
110 self.sigma_data = (
111 np.zeros((self.y_train.shape[0], self.y_train.shape[0]))
112 if sigma_data is None else np.diag(sigma_data)
113 )
114 self.sigma_data_sq_diag = (
115 np.zeros(self.y_train.shape[0])
116 if sigma_data is None
117 else np.asarray(sigma_data) ** 2
118 )
120 self.kernel_factory = KernelFactory(
121 dim=self.dim,
122 normalize=self.normalize,
123 n_order=self.max_order,
124 differences_by_dim=self.differences_by_dim,
125 smoothness_parameter=smoothness_parameter,
126 oti_module=self.oti,
127 sparse_diffs=False
128 )
129 self.kernel_func = self.kernel_factory.create_kernel(
130 kernel_name=self.kernel,
131 kernel_type=self.kernel_type
132 )
133 self.bounds = self.kernel_factory.bounds
134 self.optimizer = Optimizer(self)
136 # Sparse Cholesky setup
137 self.rho = rho
138 self.use_supernodes = use_supernodes
139 self.supernode_lam = supernode_lam
140 self._setup_sparse_cholesky()
142 def _setup_sparse_cholesky(self):
143 """Precompute MMD ordering, fill-distances, and sparsity pattern."""
144 from jetgp.full_gddegp_sparse.sparse_cholesky import (
145 mmd_ordering, build_sparsity_pattern, build_supernodes,
146 expand_mmd_permutation, expand_sparsity_to_blocks,
147 expand_supernodes_to_blocks,
148 )
149 X = self.x_train
150 self.mmd_P, self.mmd_l = mmd_ordering(X)
151 X_ord = X[self.mmd_P]
152 self.sparse_S = build_sparsity_pattern(X_ord, self.mmd_l, self.rho)
154 self.mmd_P_full, self._phys_to_rows = expand_mmd_permutation(
155 self.mmd_P, self.num_points, self.derivative_locations
156 )
157 self.sparse_S_full = expand_sparsity_to_blocks(self.sparse_S, self._phys_to_rows)
158 self.sparse_S_full_arr = {
159 j: np.asarray(s, dtype=np.intp) for j, s in self.sparse_S_full.items()
160 }
162 N_total = len(self.mmd_P_full)
163 total_nb = sum(len(s) for s in self.sparse_S_full.values())
164 self.sparse_fill_fraction = total_nb / (N_total * N_total)
165 self._use_dense_factor = self.sparse_fill_fraction > 0.25
167 if self.use_supernodes:
168 phys_sns = build_supernodes(
169 X_ord, self.mmd_l, self.sparse_S, lam=self.supernode_lam
170 )
171 self.sparse_supernodes = phys_sns
172 self.sparse_supernodes_full = expand_supernodes_to_blocks(
173 phys_sns, self._phys_to_rows
174 )
175 for sn in self.sparse_supernodes_full:
176 sn['children_arr'] = np.asarray(sn['children'])
177 ch_pos = {c: i for i, c in enumerate(sn['children'])}
178 sn['ch_pos'] = ch_pos
179 sn['parent_positions'] = np.array(
180 [ch_pos[p] for p in sn['parents']]
181 )
182 else:
183 self.sparse_supernodes = None
184 self.sparse_supernodes_full = None
186 def optimize_hyperparameters(self, *args, **kwargs):
187 """Run the optimizer. Returns optimized hyperparameter vector."""
188 self.params = self.optimizer.optimize_hyperparameters(*args, **kwargs)
189 return self.params
191 def predict(self, X_test, params, rays_predict=None, calc_cov=False,
192 return_deriv=False, derivs_to_predict=None):
193 """
194 Predict posterior mean and optional variance at test points.
195 Uses exact dense Cholesky solve (not sparse approximation).
196 """
197 n_predict = X_test.shape[0]
199 # Handle missing rays_predict when derivatives are requested
200 if return_deriv and rays_predict is None:
201 n_rays = len(self.flattened_der_indices)
202 warnings.warn(
203 f"No rays_predict provided for derivative predictions. "
204 f"Predictions will be made along coordinate axes.",
205 UserWarning
206 )
207 rays_predict = []
208 for i in range(n_rays):
209 axis_idx = i % self.dim
210 ray_array = np.zeros((self.dim, n_predict))
211 ray_array[axis_idx, :] = 1.0
212 rays_predict.append(ray_array)
214 if not return_deriv and rays_predict is not None:
215 warnings.warn(
216 "rays_predict was provided but return_deriv=False. "
217 "The provided rays will be ignored.",
218 UserWarning
219 )
221 if return_deriv and rays_predict is not None:
222 if len(self.rays_list) > 0 and len(rays_predict) > len(self.rays_list):
223 raise ValueError(
224 f"Number of prediction rays ({len(rays_predict)}) exceeds the number of "
225 f"training rays ({len(self.rays_list)})."
226 )
227 for i, ray_array in enumerate(rays_predict):
228 if not isinstance(ray_array, np.ndarray):
229 raise TypeError(
230 f"Ray array {i} must be a numpy ndarray, got {type(ray_array).__name__}."
231 )
232 if ray_array.ndim != 2:
233 raise ValueError(
234 f"Ray array {i} must be 2-dimensional, got {ray_array.ndim} dimensions."
235 )
236 if ray_array.shape[0] != self.dim:
237 raise ValueError(
238 f"Ray array {i} has {ray_array.shape[0]} rows, expected {self.dim}."
239 )
240 if ray_array.shape[1] != n_predict:
241 raise ValueError(
242 f"Ray array {i} has {ray_array.shape[1]} columns, expected {n_predict}."
243 )
245 length_scales = params[:-1]
246 sigma_n = params[-1]
248 if return_deriv:
249 if derivs_to_predict is not None:
250 common_derivs = derivs_to_predict
251 else:
252 common_derivs = self.flattened_der_indices
254 required_order = max(
255 sum(pair[1] for pair in deriv_spec)
256 for deriv_spec in common_derivs
257 )
258 predict_order = max(required_order, self.n_order)
260 if predict_order > self.n_order:
261 predict_oti = get_oti_module(self.n_bases, predict_order)
262 smoothness_param = getattr(self.kernel_factory, 'alpha', None)
263 predict_kernel_factory = KernelFactory(
264 dim=self.dim,
265 normalize=self.normalize,
266 differences_by_dim=self.differences_by_dim,
267 n_order=predict_order,
268 smoothness_parameter=smoothness_param,
269 oti_module=predict_oti,
270 sparse_diffs=False
271 )
272 predict_kernel_func = predict_kernel_factory.create_kernel(
273 kernel_name=self.kernel, kernel_type=self.kernel_type
274 )
275 else:
276 predict_oti = self.oti
277 predict_kernel_func = self.kernel_func
278 else:
279 common_derivs = []
280 predict_order = self.n_order
281 predict_oti = self.oti
282 predict_kernel_func = self.kernel_func
284 _cache_hit = (
285 hasattr(self, '_cached_params')
286 and self._cached_params is not None
287 and np.array_equal(self._cached_params, params)
288 and getattr(self, '_cached_L', None) is not None
289 )
291 if _cache_hit:
292 L = self._cached_L
293 low = self._cached_low
294 alpha = self._cached_alpha
295 self.n_bases = self._cached_n_bases
296 cho_solve_failed = False
297 else:
298 phi_train = self.kernel_func(self.differences_by_dim, length_scales)
299 if self.n_order == 0:
300 self.n_bases = 0
301 phi_exp_train = phi_train.real[np.newaxis, :, :]
302 else:
303 active = phi_train.get_active_bases()
304 self.n_bases = max(self.n_bases, active[-1] if active else 0)
305 phi_exp_train = phi_train.get_all_derivs(self.n_bases, 2 * self.n_order)
307 powers = [0] * (len(self.flattened_der_indices) + 1)
309 K = gddegp_utils.rbf_kernel(
310 phi_train, phi_exp_train, self.n_order, self.n_bases,
311 self.flattened_der_indices,
312 index=self.derivative_locations
313 )
314 K.flat[::K.shape[0] + 1] += (10 ** sigma_n) ** 2
315 K += self.sigma_data ** 2
317 try:
318 cho_solve_failed = False
319 L, low = cho_factor(K, lower=True)
320 alpha = cho_solve((L, low), self.y_train)
321 except Exception:
322 cho_solve_failed = True
323 alpha = np.linalg.solve(K, self.y_train)
324 L, low = None, None
326 self._cached_L = L
327 self._cached_low = low
328 self._cached_alpha = alpha
329 self._cached_n_bases = self.n_bases
330 self._cached_params = params.copy()
332 rays_test = rays_predict
334 if self.normalize:
335 X_test = utils.normalize_x_data_test(X_test, self.sigmas_x, self.mus_x)
337 if not return_deriv:
338 rays_test = None
339 derivative_locations_test = None
340 else:
341 derivative_locations_test = [
342 list(range(X_test.shape[0])) for _ in range(len(common_derivs))]
343 if self.normalize:
344 rays_test = utils.normalize_directions_2(self.sigmas_x, rays_test)
346 diff_x_train_x_test = gddegp_utils.differences_by_dim_func(
347 self.x_train, X_test,
348 self.rays_list, rays_test,
349 self.derivative_locations, derivative_locations_test,
350 predict_order, predict_oti, return_deriv=return_deriv
351 )
353 phi_train_test = predict_kernel_func(diff_x_train_x_test, length_scales)
354 if predict_order > 0:
355 if return_deriv:
356 phi_exp_train_test = phi_train_test.get_all_derivs(self.n_bases, 2 * predict_order)
357 else:
358 phi_exp_train_test = phi_train_test.get_all_derivs(self.n_bases, predict_order)
359 else:
360 phi_exp_train_test = phi_train_test.real[np.newaxis, :, :]
362 K_s = gddegp_utils.rbf_kernel_predictions(
363 phi_train_test, phi_exp_train_test, predict_order, self.n_bases,
364 self.flattened_der_indices,
365 return_deriv=return_deriv,
366 index=self.derivative_locations,
367 common_derivs=common_derivs
368 )
370 f_mean = K_s @ alpha
372 if self.normalize:
373 if return_deriv:
374 f_mean = utils.transform_predictions_directional(
375 f_mean, self.mu_y, self.sigma_y, self.sigmas_x,
376 common_derivs, X_test)
377 else:
378 f_mean = self.mu_y + f_mean * self.sigma_y
380 f_mean = f_mean.reshape(-1, 1)
381 n = X_test.shape[0]
382 m = f_mean.shape[0]
383 num_derivs = m // n
384 reshaped_mean = f_mean.reshape(num_derivs, n)
386 if not calc_cov:
387 return reshaped_mean
389 diff_x_test_x_test = gddegp_utils.differences_by_dim_func(
390 X_test, X_test,
391 rays_test, rays_test,
392 derivative_locations_test, derivative_locations_test,
393 predict_order, predict_oti, return_deriv=return_deriv
394 )
396 phi_test_test = predict_kernel_func(diff_x_test_x_test, length_scales)
397 if predict_order > 0:
398 phi_exp_test_test = phi_test_test.get_all_derivs(self.n_bases, 2 * predict_order)
399 else:
400 phi_exp_test_test = phi_test_test.real[np.newaxis, :, :]
402 K_ss = gddegp_utils.rbf_kernel_predictions(
403 phi_test_test, phi_exp_test_test, predict_order, self.n_bases,
404 self.flattened_der_indices,
405 return_deriv=return_deriv,
406 index=derivative_locations_test,
407 common_derivs=common_derivs,
408 calc_cov=True,
409 )
411 if cho_solve_failed:
412 v_fallback = np.linalg.solve(K, K_s.T)
413 f_cov = K_ss - K_s @ v_fallback
414 else:
415 v = solve_triangular(L, K_s.T, lower=low)
416 f_cov = K_ss - v.T @ v
418 if self.normalize:
419 if return_deriv:
420 f_var = utils.transform_cov_directional(
421 f_cov, self.sigma_y, self.sigmas_x,
422 common_derivs, X_test)
423 else:
424 f_var = self.sigma_y ** 2 * np.diag(np.abs(f_cov))
425 else:
426 f_var = np.diag(np.abs(f_cov))
428 reshaped_var = f_var.reshape(num_derivs, n)
429 return reshaped_mean, reshaped_var