diff --git a/optuna_dashboard/preferential/samplers/_gp.py b/optuna_dashboard/preferential/samplers/_gp.py index ff4001c1..e7e6085d 100644 --- a/optuna_dashboard/preferential/samplers/_gp.py +++ b/optuna_dashboard/preferential/samplers/_gp.py @@ -5,19 +5,16 @@ from math import erfc from typing import Any from botorch.acquisition.analytic import LogExpectedImprovement -from botorch.models.gpytorch import GPyTorchModel from botorch.optim import optimize_acqf +import botorch.models.model +import botorch.posteriors.gpytorch + +import botorch import gpytorch.constraints import gpytorch.kernels import gpytorch.likelihoods.gaussian_likelihood -from gpytorch.likelihoods.gaussian_likelihood import GaussianLikelihood -from gpytorch.likelihoods.gaussian_likelihood import Interval from gpytorch.likelihoods.gaussian_likelihood import Prior -from gpytorch.models.exact_gp import ExactGP import gpytorch.module -from linear_operator.operators import DiagLinearOperator -from linear_operator.operators import LinearOperator -from linear_operator.utils.errors import NotPSDError import numpy as np import optuna from optuna import distributions @@ -26,8 +23,6 @@ from optuna._transform import _SearchSpaceTransform from optuna.distributions import BaseDistribution from optuna.search_space import IntersectionSearchSpace from optuna.trial import FrozenTrial -import pyro -import pyro.infer.autoguide import pyro.infer.mcmc from scipy.special import erfcinv import torch @@ -36,288 +31,252 @@ from torch import Tensor from .._system_attrs import get_preferences -class _WeightedGaussianLikelihood(GaussianLikelihood): - def __init__( - self, - weights: torch.Tensor | None = None, - noise_prior: Prior | None = None, - noise_constraint: Interval | None = None, - batch_shape: torch.Size = torch.Size(), - **kwargs: Any, - ) -> None: - super().__init__( - noise_prior=noise_prior, - noise_constraint=noise_constraint, - batch_shape=batch_shape, - **kwargs, - ) - self.weights = weights - - def _shaped_noise_covar( - self, base_shape: torch.Size, *params: Any, **kwargs: Any - ) -> Tensor | LinearOperator: - assert self.weights is not None - assert base_shape[-1] == self.weights.shape[-1] - return DiagLinearOperator(1.0 / self.weights) * super()._shaped_noise_covar( - base_shape, *params, **kwargs - ) - - -def _sample_y( - preferences: np.ndarray, - cov_X_X: np.ndarray, - obs_noise_var: float, - cycles: int, - initial_sample: np.ndarray, - rng: np.random.RandomState, -) -> np.ndarray: - # TODO: Refactor and write tests for this function. - - N = cov_X_X.shape[0] - M = len(preferences) - cov_X_X = cov_X_X + np.eye(N) * 1e-6 # Add jitter - cov_X_X_chol = np.linalg.cholesky(cov_X_X) - cov_X_X_inv = np.linalg.inv(cov_X_X) - - # (sI + A K A^T)^-1 = s^-1 I - s^-2 A(K^-1 + s^-1 A^T A)^-1 A^T - - schur = cov_X_X_inv.copy() - np.add.at(schur, (preferences[:, 0], preferences[:, 0]), 1.0 / (2 * obs_noise_var)) - np.add.at(schur, (preferences[:, 1], preferences[:, 1]), 1.0 / (2 * obs_noise_var)) - np.add.at(schur, (preferences[:, 0], preferences[:, 1]), -1.0 / (2 * obs_noise_var)) - np.add.at(schur, (preferences[:, 1], preferences[:, 0]), -1.0 / (2 * obs_noise_var)) - idx_M = np.arange(M) - - schur_inv = np.linalg.inv(schur) - - cov_diff_inv = schur_inv[:, preferences[:, 0]] - schur_inv[:, preferences[:, 1]] - cov_diff_inv = cov_diff_inv[preferences[:, 0], :] - cov_diff_inv[preferences[:, 1], :] - cov_diff_inv *= -1 / (2 * obs_noise_var) ** 2 - cov_diff_inv[idx_M, idx_M] += 1.0 / (2 * obs_noise_var) - - diffs = _orthants_MVN_Gibbs_sampling( - cov_diff_inv, - cycles=cycles, - initial_sample=initial_sample[:, 0] - initial_sample[:, 1], - rng=rng, - )[-1] - - random_ys = (cov_X_X_chol @ rng.randn(N))[preferences] + np.sqrt(obs_noise_var) * rng.randn( - M, 2 - ) - errors = diffs - (random_ys[:, 0] - random_ys[:, 1]) - cov_diff_inv_errors = cov_diff_inv @ errors - - AT_cov_diff_inv_errors = np.zeros((N,)) - np.add.at(AT_cov_diff_inv_errors, preferences[:, 0], cov_diff_inv_errors) - np.add.at(AT_cov_diff_inv_errors, preferences[:, 1], -cov_diff_inv_errors) - - return ( - random_ys - + (cov_X_X @ AT_cov_diff_inv_errors)[preferences] - + obs_noise_var * np.array([[1, -1]]) * cov_diff_inv_errors[:, None] - ) - - _SQRT2 = math.sqrt(2) - def _orthants_MVN_Gibbs_sampling( - cov_inv: np.ndarray, + cov_inv: torch.Tensor, cycles: int, - initial_sample: np.ndarray, - rng: np.random.RandomState, -) -> np.ndarray: + initial_sample: torch.Tensor, +) -> torch.Tensor: dim = cov_inv.shape[0] assert cov_inv.shape == (dim, dim) - - if initial_sample is None: - sample_chain = np.zeros(dim) - else: + with torch.no_grad(): sample_chain = initial_sample + conditional_std = 1 / torch.sqrt(torch.diag(cov_inv)) + scaled_cov_inv = cov_inv / torch.diag(cov_inv)[:, None] - conditional_std = 1 / np.sqrt(np.diag(cov_inv)) + out = torch.empty((cycles + 1, dim), dtype=torch.float64) + out[0, :] = sample_chain - scaled_cov_inv = cov_inv / np.c_[np.diag(cov_inv)] - - out = np.empty((cycles + 1, dim)) - out[0, :] = sample_chain - - for i in range(cycles): - for j in range(dim): - conditional_mean = sample_chain[j] - scaled_cov_inv[j] @ sample_chain - sample_chain[j] = ( - _one_side_trunc_norm_sampling( - lower=-conditional_mean / conditional_std[j], rng=rng + for i in range(cycles): + for j in range(dim): + conditional_mean = sample_chain[j] - scaled_cov_inv[j] @ sample_chain + sample_chain[j] = ( + _one_side_trunc_norm_sampling(lower=float(-conditional_mean / conditional_std[j])) + * conditional_std[j] + + conditional_mean ) - * conditional_std[j] - + conditional_mean - ) - out[i + 1, :] = sample_chain + out[i + 1, :] = sample_chain - return out + return out -def _one_side_trunc_norm_sampling(lower: float, rng: np.random.RandomState) -> float: - return erfcinv(rng.rand() * erfc(lower / _SQRT2)) * _SQRT2 +def _one_side_trunc_norm_sampling(lower: float) -> float: + return erfcinv(torch.rand() * erfc(lower / _SQRT2)) * _SQRT2 +def _compute_cov_diff_diff_inv( + preferences: torch.Tensor, + cov_x_x: torch.Tensor, + obs_noise_var: float, +) -> torch.Tensor: + N = cov_x_x.shape[0] + M = preferences.shape[0] -class _PreferentialGP(GPyTorchModel, ExactGP): - _num_outputs = 1 + # (sI + A K A^T)^-1 = s^-1 I - s^-2 A(K^-1 + s^-1 A^T A)^-1 A^T + # (K^-1 + s^-1 A^T A)^-1 = K (I + s^-1 A^T A K)^-1 (To avoid computing K^-1) + I_plus_sinv_AT_A_K = torch.eye(N) + A_K = cov_x_x[preferences[:, 0], :] - cov_x_x[preferences[:, 1], :] + I_plus_sinv_AT_A_K.index_add_(0, preferences[:, 0], A_K, alpha=1 / obs_noise_var) + I_plus_sinv_AT_A_K.index_add_(0, preferences[:, 1], A_K, alpha=-1/obs_noise_var) + schur_inv: torch.Tensor = torch.linalg.solve(I_plus_sinv_AT_A_K, cov_x_x, left=False) + cov_diff_diff_inv = schur_inv[:, preferences[:, 0]] - schur_inv[:, preferences[:, 1]] + cov_diff_diff_inv = cov_diff_diff_inv[preferences[:, 0], :] - cov_diff_diff_inv[preferences[:, 1], :] + cov_diff_diff_inv *= -1 / (2 * obs_noise_var) ** 2 + idx_M = torch.arange(M) + cov_diff_diff_inv[idx_M, idx_M] += 1.0 / (2 * obs_noise_var) + + return cov_diff_diff_inv + +def _multinormal_logpdf(Sigma_inv: torch.Tensor, x: torch.Tensor): + return -0.5 * x @ Sigma_inv @ x + 0.5 * torch.logdet(Sigma_inv) - 0.5 * x.shape[0] * math.log(2 * math.pi) + +class _SampledGP(botorch.models.model.Model): def __init__( self, kernel: gpytorch.kernels.Kernel, - noise_prior: Prior | None = None, - noise_constraint: Interval | None = None, + x: torch.Tensor, + preferences: torch.Tensor, + obs_noise_var: torch.Tensor, + diff: torch.Tensor, ) -> None: - GPyTorchModel.__init__(self) - likelihood = _WeightedGaussianLikelihood( - noise_prior=noise_prior, noise_constraint=noise_constraint + self.kernel = kernel + self.x = x + self.preferences = preferences + self.diff = diff + self.obs_noise_var = obs_noise_var + self._cov_diff_diff_inv = _compute_cov_diff_diff_inv( + preferences=preferences, + cov_x_x=self.kernel(x).to_dense(), + obs_noise_var=float(obs_noise_var), ) - ExactGP.__init__(self, train_inputs=None, train_targets=None, likelihood=likelihood) - self.covar_module = kernel - self._last_params: dict[str, torch.Tensor] | None = None - self._last_mcmc_step_size: float | None = None + def posterior( + self, + x2: Tensor, + output_indices: list[int] | None = None, + observation_noise: bool = False, + posterior_transform: Any | None = None, + **kwargs: Any, + ) -> botorch.posteriors.gpytorch.GPyTorchPosterior: + assert posterior_transform is None + assert output_indices is None + assert self.x.shape[-1] == x2.shape[-1] - def _pyro_model(self, train_x: torch.Tensor, train_y: torch.Tensor) -> None: - # with gpytorch.settings.fast_computations(False, False, False): - sampled_model = self.pyro_sample_from_prior() + x_expanded = self.x.expand(x2.shape[:-2] + (self.x.shape[-2], x2.shape[-1])) - ys = sampled_model.likelihood(sampled_model.forward(train_x)) + cov_x2_x: torch.Tensor = self.kernel(x2, x_expanded).to_dense() + cov_x2_diff: torch.Tensor = cov_x2_x[..., self.preferences[:, 0]] - cov_x2_x[..., self.preferences[:, 1]] - pyro.sample("y", ys, obs=train_y) + mean: torch.Tensor = cov_x2_diff @ (self._cov_diff_diff_inv @ self.diff) + cov: torch.Tensor = self.kernel(x2).to_dense() - cov_x2_diff @ self._cov_diff_diff_inv @ cov_x2_diff.transpose(-1, -2) + if observation_noise: + idx = torch.arange(cov.shape[-1]) + cov[..., idx, idx] += self.obs_noise_var - def fit_mcmc( - self, X: torch.Tensor, preferences: torch.Tensor, cycles: int, rng: np.random.RandomState + return botorch.posteriors.gpytorch.GPyTorchPosterior( + distribution=gpytorch.distributions.MultivariateNormal( + mean=mean, + covariance_matrix=cov, + ) + ) + + @property + def batch_shape(self) -> torch.Size: + return torch.Size() + + @property + def num_outputs(self) -> int: + return 1 + + +class _PreferentialGP: + def _kernel_factory(self, lengthscale: torch.Tensor) -> gpytorch.kernels.Kernel: + kernel = gpytorch.kernels.MaternKernel( + nu=2.5, + ard_num_dims=lengthscale.shape[0], + ) + kernel.lengthscale = lengthscale + return kernel + + def _potential_func( + self, + x: torch.Tensor, + preferences: torch.Tensor, + diff: torch.Tensor, + log_lengthscale: torch.Tensor, + log_noise: torch.Tensor, + ) -> torch.Tensor: + + lengthscale = torch.exp(log_lengthscale) + noise = torch.exp(log_noise) + log_transform_jacobian = torch.sum(log_lengthscale) + log_noise + log_prior = self.lengthscale_prior.log_prob(lengthscale) + self.noise_prior.log_prob(noise) + cov_inv = _compute_cov_diff_diff_inv( + preferences=preferences, + cov_x_x=self._kernel_factory(lengthscale)(x).to_dense(), + obs_noise_var=noise, + ) + log_likelihood = _multinormal_logpdf(cov_inv, diff) + return log_prior + log_likelihood + log_transform_jacobian + + def __init__( + self, + lengthscale_prior: Prior, + noise_prior: Prior, + dims: int, ) -> None: + self.lengthscale_prior: Prior = lengthscale_prior.expand((dims,)) + self.noise_prior = noise_prior + self.dims = dims + + self._x = torch.empty((0, dims), dtype=torch.float64) + self._preferences = torch.empty((0, 2), dtype=torch.float64) + self._diff = torch.empty((0,), dtype=torch.float64) + + initial_raw_params = { + "log_lengthscale": torch.log(self.lengthscale_prior.sample()), + "log_noise": torch.log(self.noise_prior.sample()), + } + + self._potential_func_jit = torch.jit.trace( + self._potential_func, + (self._x, self._preferences, self._diff, initial_raw_params["log_lengthscale"], initial_raw_params["log_noise"]), + ) + + # HMC-Gibbs workarounds + # https://github.com/pyro-ppl/pyro/issues/1926 + + self._nuts = pyro.infer.mcmc.NUTS(potential_fn=lambda z:self._potential_func_jit( + x=self._x, + preferences=self._preferences, + diff=self._diff, + log_lengthscale=z["log_lengthscale"], + log_noise=z["log_noise"], + )) + + self._nuts.initial_params = initial_raw_params + self._nuts.setup(warmup_steps=1e15) # Infinite warmup + self._last_params = initial_raw_params + + def sample_gp( + self, x: torch.Tensor, preferences: torch.Tensor, cycles: int, rng: np.random.RandomState + ) -> _SampledGP: if len(preferences) == 0: - # Skip actual MCMC computation - self.set_train_data( - inputs=torch.empty((0, X.shape[-1])), - targets=torch.empty((0,)), - strict=False, + return _SampledGP( + kernel=self._kernel_factory(self.lengthscale_prior.sample()), + x=x, + preferences=preferences, + obs_noise_var=self.noise_prior.sample(), + diff=torch.empty((0,), dtype=torch.float64), ) - self.likelihood.weights = torch.empty((0,)) else: - dtype = torch.float64 - cnt = torch.bincount(preferences.reshape(-1)) - mask = cnt > 0 - train_x = X[mask] - weights = cnt[mask] + original_diff_size = len(self._diff) + self._diff.resize_(len(preferences)) + self._diff[original_diff_size:] = 0.0 - assert isinstance(self.likelihood, _WeightedGaussianLikelihood) - self.likelihood.weights = weights - - preferences_np = preferences.detach().numpy() - - all_ys_np = np.zeros((len(preferences), 2)) - train_y = torch.zeros( - ( - len( - train_x, - ) - ), - dtype=dtype, - ) - - nuts = pyro.infer.mcmc.NUTS( - model=self._pyro_model, - init_strategy=pyro.infer.autoguide.init_to_sample, - step_size=self._last_mcmc_step_size or 1.0, - ) - warmup_steps = max(0, cycles - 2) - nuts.setup(warmup_steps=warmup_steps, train_x=train_x, train_y=train_y) - - raw_params = self._last_params or nuts.initial_params - for i in range(cycles): - params = { - name: nuts.transforms[name].inv(value) for name, value in raw_params.items() - } - _set_params(self, params) - self.set_train_data(train_x, train_y, strict=False) - all_ys_np = _sample_y( - preferences=preferences_np, - cov_X_X=self.covar_module(train_x).detach().numpy(), - obs_noise_var=float(self.likelihood.noise_covar.noise), - cycles=10, - initial_sample=all_ys_np, - rng=rng, + for _ in range(cycles): + kernel = self._kernel_factory(torch.exp(self._last_params["log_lengthscale"])) + cov_diff_diff_inv = _compute_cov_diff_diff_inv( + preferences=preferences, + cov_x_x=kernel(x).to_dense(), + obs_noise_var=float(torch.exp(self._last_params["log_noise"])), ) - ys_sum_np = np.zeros((len(X),)) - np.add.at(ys_sum_np, preferences_np.reshape(-1), all_ys_np.reshape(-1)) - ys_sum = torch.from_numpy(ys_sum_np) - train_y[:] = ys_sum[mask] / cnt[mask] - nuts.clear_cache() - try: - raw_params = nuts.sample(raw_params) - except NotPSDError: - nuts.cleanup() - nuts = pyro.infer.mcmc.NUTS( - model=self._pyro_model, - init_strategy=pyro.infer.autoguide.init_to_sample, - step_size=self._last_mcmc_step_size or 1.0, - ) - nuts.setup(warmup_steps=warmup_steps, train_x=train_x, train_y=train_y) - raw_params = nuts.initial_params - - params = {name: nuts.transforms[name].inv(value) for name, value in raw_params.items()} - self.set_train_data(train_x, train_y, strict=False) - _set_params(self, params) - - self._last_params = raw_params - self._last_mcmc_step_size = nuts.step_size - nuts.cleanup() - - def forward(self, x: torch.Tensor) -> gpytorch.distributions.MultivariateNormal: - mean_module = gpytorch.means.ZeroMean() - return gpytorch.distributions.MultivariateNormal( - mean_module(x), - self.covar_module(x), - ) - - -def _set_params( - module: gpytorch.Module, - params_dict: dict[str, torch.Tensor], - memo: set | None = None, - prefix: str = "", -) -> None: - if memo is None: - memo = set() - if hasattr(module, "_priors"): - for name, (prior, closure, setting_closure) in module._priors.items(): - if prior is not None and prior not in memo: - memo.add(prior) - setting_closure(module, params_dict[prefix + ("." if prefix else "") + name]) - - for mname, module_ in module.named_children(): - submodule_prefix = prefix + ("." if prefix else "") + mname - _set_params(module_, params_dict, memo=memo, prefix=submodule_prefix) + self._diff = _orthants_MVN_Gibbs_sampling( + cov_inv=cov_diff_diff_inv, + cycles=10, + initial_sample=self._diff + )[:-1] + self._nuts.clear_cache() + self._last_raw_params = self._nuts.sample(self._last_raw_params) + + return _SampledGP( + kernel=self._kernel_factory(torch.exp(self._last_params["log_lengthscale"])), + x=x, + preferences=preferences, + obs_noise_var=torch.exp(self._last_params["log_noise"]), + diff=self._diff, + ) class PreferentialGPSampler(optuna.samplers.BaseSampler): def __init__( self, *, - kernel: gpytorch.kernels.Kernel | None = None, + # kernel_factory: typing.Callable[[int], gpytorch.kernels.Kernel] | None = None, + lengthscale_prior: Prior | None = None, noise_prior: Prior | None = None, independent_sampler: optuna.samplers.BaseSampler | None = None, seed: int | None = None, - device: torch.device | None = None, + # device: torch.device | None = None, ) -> None: + self.lengthscale_prior = lengthscale_prior or gpytorch.priors.GammaPrior(3.0, 6.0) + self.noise_prior = noise_prior or gpytorch.priors.GammaPrior(1.1, 10.0) + self.independent_sampler = independent_sampler or optuna.samplers.RandomSampler(seed=self._rng.randint(2**32)) + self._rng = np.random.RandomState(seed) self._search_space = IntersectionSearchSpace() - - self.kernel = kernel - self.noise_prior = noise_prior - self.independent_sampler = independent_sampler or optuna.samplers.RandomSampler( - seed=self._rng.randint(2**32), - ) - self.device = device or torch.device("cpu") - self._gp: _PreferentialGP | None = None def reseed_rng(self) -> None: @@ -352,15 +311,9 @@ class PreferentialGPSampler(optuna.samplers.BaseSampler): ) dims = len(trans.bounds) self._gp = self._gp or _PreferentialGP( - kernel=self.kernel - or gpytorch.kernels.MaternKernel( - nu=2.5, - ard_num_dims=dims, - lengthscale_prior=gpytorch.priors.GammaPrior(3.0, 6.0), - lengthscale_constraint=gpytorch.constraints.Positive(), - ), - noise_prior=self.noise_prior or gpytorch.priors.GammaPrior(1.1, 2.0), - noise_constraint=gpytorch.constraints.Positive(), + lengthscale_prior=self.lengthscale_prior, + noise_prior=self.noise_prior, + dims=dims, ) ids: dict[int, int] = {} @@ -373,23 +326,12 @@ class PreferentialGPSampler(optuna.samplers.BaseSampler): ids[t] = len(ids) params.append(trans.transform(trials[t].params)) pref_ids.append((ids[better], ids[worse])) - dtype = torch.float64 - - params_torch = torch.tensor(np.array(params), dtype=dtype, device=self.device) - pref_ids_torch = torch.tensor( - np.array(pref_ids), - dtype=torch.int32, - device=self.device, - ) - self._gp.fit_mcmc(params_torch, pref_ids_torch, cycles=10, rng=self._rng) - self._gp.eval() - scores = self._gp(params_torch).mean - - best_f = torch.max(scores) - + params_torch = torch.tensor(np.array(params), dtype=torch.float64) + pref_ids_torch = torch.tensor(np.array(pref_ids), dtype=torch.int32) + sampled_gp = self._gp.sample_gp(params_torch, pref_ids_torch, cycles=10, rng=self._rng) acqf = LogExpectedImprovement( - model=self._gp, - best_f=best_f, + model=sampled_gp, + best_f=torch.max(sampled_gp.posterior(params_torch[:, None, :]).mean), ) # TODO: Make it possible to apply it on categorical variables @@ -415,3 +357,4 @@ class PreferentialGPSampler(optuna.samplers.BaseSampler): return self.independent_sampler.sample_independent( study, trial, param_name, param_distribution ) +