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

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 

8 

9 

10class gddegp: 

11 """ 

12 Sparse Cholesky variant of the GDDEGP model. 

13 

14 Adds sparse inverse-Cholesky acceleration for NLML evaluation during 

15 hyperparameter optimisation. Prediction always uses the exact dense 

16 Cholesky solve. 

17 

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 """ 

51 

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): 

57 

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 ) 

67 

68 elif der_indices is None and n_order == 0: 

69 der_indices = [] 

70 derivative_locations = [] 

71 

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 

85 

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) 

92 

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) 

102 

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 ) 

109 

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 ) 

119 

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) 

135 

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() 

141 

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) 

153 

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 } 

161 

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 

166 

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 

185 

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 

190 

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] 

198 

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) 

213 

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 ) 

220 

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 ) 

244 

245 length_scales = params[:-1] 

246 sigma_n = params[-1] 

247 

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 

253 

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) 

259 

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 

283 

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 ) 

290 

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) 

306 

307 powers = [0] * (len(self.flattened_der_indices) + 1) 

308 

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 

316 

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 

325 

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() 

331 

332 rays_test = rays_predict 

333 

334 if self.normalize: 

335 X_test = utils.normalize_x_data_test(X_test, self.sigmas_x, self.mus_x) 

336 

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) 

345 

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 ) 

352 

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, :, :] 

361 

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 ) 

369 

370 f_mean = K_s @ alpha 

371 

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 

379 

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) 

385 

386 if not calc_cov: 

387 return reshaped_mean 

388 

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 ) 

395 

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, :, :] 

401 

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 ) 

410 

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 

417 

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)) 

427 

428 reshaped_var = f_var.reshape(num_derivs, n) 

429 return reshaped_mean, reshaped_var