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

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 

8 

9 

10class ddegp: 

11 """ 

12 Sparse Cholesky variant of the Directional DEGP 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 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 """ 

49 

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

54 

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 = [] 

68 

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) 

83 

84 self.flattened_der_indices = utils.flatten_der_indices(der_indices) 

85 

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) 

95 

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) 

99 

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 ) 

109 

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) 

126 

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

132 

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) 

144 

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 } 

152 

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 

157 

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 

176 

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 

181 

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] 

189 

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 

195 

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) 

201 

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 

220 

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 

229 

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 ) 

236 

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

250 

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 

258 

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 

267 

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 

273 

274 if self.normalize: 

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

276 

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 

282 

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) 

285 

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 ) 

302 

303 f_mean = K_s.T @ alpha 

304 

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 

312 

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) 

318 

319 if not calc_cov: 

320 return reshaped_mean 

321 

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) 

324 

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 

328 

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 ) 

342 

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 

348 

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

358 

359 reshaped_var = f_var.reshape(num_derivs, n) 

360 return reshaped_mean, reshaped_var