diff --git a/optuna_dashboard/preferential/samplers/_gp.py b/optuna_dashboard/preferential/samplers/_gp.py index 2f8bdba6..7f1f4dda 100644 --- a/optuna_dashboard/preferential/samplers/_gp.py +++ b/optuna_dashboard/preferential/samplers/_gp.py @@ -1,15 +1,16 @@ -#%% +# %% from __future__ import annotations import math -from typing import Any, Callable +from typing import Any +from typing import Callable import botorch.acquisition.analytic import botorch.models.model import botorch.optim import botorch.posteriors.gpytorch -import gpytorch.kernels import gpytorch.constraints +import gpytorch.kernels from gpytorch.likelihoods.gaussian_likelihood import Prior import numpy as np import optuna @@ -19,6 +20,7 @@ import torch from .._system_attrs import get_preferences + def _orthants_MVN_Gibbs_sampling( cov_inv: torch.Tensor, cycles: int, @@ -38,9 +40,7 @@ def _orthants_MVN_Gibbs_sampling( 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] - ) + _one_side_trunc_norm_sampling(lower=-conditional_mean / conditional_std[j]) * conditional_std[j] + conditional_mean ) @@ -48,6 +48,7 @@ def _orthants_MVN_Gibbs_sampling( return out + def _one_side_trunc_norm_sampling(lower: torch.Tensor) -> torch.Tensor: if lower > 4.0: r = torch.max(torch.tensor(1e-300), torch.rand(torch.Size(()), dtype=torch.float64)) @@ -61,7 +62,8 @@ def _one_side_trunc_norm_sampling(lower: torch.Tensor) -> torch.Tensor: _orthants_MVN_Gibbs_sampling_jit = torch.jit.script(_orthants_MVN_Gibbs_sampling) - + + def _compute_cov_diff_diff_inv_and_logdet( preferences: torch.Tensor, cov_x_x: torch.Tensor, @@ -88,12 +90,13 @@ def _compute_cov_diff_diff_inv_and_logdet( cov_diff_diff_inv = ( cov_diff_diff_inv[preferences[:, 0], :] - cov_diff_diff_inv[preferences[:, 1], :] ) - cov_diff_diff_inv *= -1 / obs_noise_var ** 2 + cov_diff_diff_inv *= -1 / obs_noise_var**2 idx_M = torch.arange(M) cov_diff_diff_inv[idx_M, idx_M] += 1.0 / obs_noise_var return cov_diff_diff_inv, logdet + class _SampledGP(botorch.models.model.Model): def __init__( self, @@ -133,7 +136,9 @@ class _SampledGP(botorch.models.model.Model): cov_X_diff = cov_X_x[..., self.preferences[:, 0]] - cov_X_x[..., self.preferences[:, 1]] mean = cov_X_diff @ (self._cov_diff_diff_inv @ self.diff) - cov = self.kernel_func(X, X) - cov_X_diff @ self._cov_diff_diff_inv @ cov_X_diff.transpose(-1, -2) + cov = self.kernel_func(X, X) - cov_X_diff @ self._cov_diff_diff_inv @ cov_X_diff.transpose( + -1, -2 + ) if observation_noise: idx = torch.arange(cov.shape[-1]) cov[..., idx, idx] += self.obs_noise_var @@ -155,7 +160,9 @@ class _SampledGP(botorch.models.model.Model): class _PreferentialGP: - def _kernel_func(self, x1: torch.Tensor, x2: torch.Tensor, lengthscale: torch.Tensor, nu: float) -> torch.Tensor: + def _kernel_func( + self, x1: torch.Tensor, x2: torch.Tensor, lengthscale: torch.Tensor, nu: float + ) -> torch.Tensor: x1_ = x1.div(lengthscale) x2_ = x2.div(lengthscale) distance = torch.cdist(x1_, x2_) @@ -183,7 +190,9 @@ class _PreferentialGP: noise = torch.exp(log_noise) + self.minimum_noise log_transform_jacobian = torch.sum(log_lengthscale) + log_noise - log_prior = torch.sum(self.lengthscale_prior.log_prob(lengthscale)) + self.noise_prior.log_prob(noise) + log_prior = torch.sum( + self.lengthscale_prior.log_prob(lengthscale) + ) + self.noise_prior.log_prob(noise) cov_x_x = self._kernel_func(x, x, lengthscale, 2.5) cov_inv, cov_inv_logdet = _compute_cov_diff_diff_inv_and_logdet( preferences=preferences, @@ -219,7 +228,14 @@ class _PreferentialGP: } 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"]), + self._potential_func, + ( + self._x, + self._preferences, + self._diff, + initial_raw_params["log_lengthscale"], + initial_raw_params["log_noise"], + ), check_trace=False, ) @@ -233,7 +249,7 @@ class _PreferentialGP: diff=self._diff, log_lengthscale=z["log_lengthscale"], log_noise=z["log_noise"], - ), + ), adapt_step_size=True, adapt_mass_matrix=False, target_accept_prob=0.5, @@ -244,7 +260,6 @@ class _PreferentialGP: self._nuts.setup(warmup_steps=1e15) # Use default step size self._last_params = initial_raw_params - def sample_gp(self, x: torch.Tensor, preferences: torch.Tensor, cycles: int) -> _SampledGP: if len(preferences) == 0: lengthscale = self.lengthscale_prior.sample() + self.minimum_lengthscale @@ -272,16 +287,27 @@ class _PreferentialGP: with torch.no_grad(): cov_diff_diff_inv, _ = _compute_cov_diff_diff_inv_and_logdet( preferences=preferences, - cov_x_x=self._kernel_func(x, x, torch.exp(self._last_params["log_lengthscale"]) + self.minimum_lengthscale, nu=2.5), - obs_noise_var=torch.exp(self._last_params["log_noise"]) + self.minimum_noise, + cov_x_x=self._kernel_func( + x, + x, + torch.exp(self._last_params["log_lengthscale"]) + + self.minimum_lengthscale, + nu=2.5, + ), + obs_noise_var=torch.exp(self._last_params["log_noise"]) + + self.minimum_noise, ) self._diff = _orthants_MVN_Gibbs_sampling_jit( - cov_inv=cov_diff_diff_inv, initial_sample=self._diff, cycles=10, + cov_inv=cov_diff_diff_inv, + initial_sample=self._diff, + cycles=10, )[-1] self._nuts.clear_cache() self._last_params = self._nuts.sample(self._last_params) - lengthscale = torch.exp(self._last_params["log_lengthscale"]) + self.minimum_lengthscale + lengthscale = ( + torch.exp(self._last_params["log_lengthscale"]) + self.minimum_lengthscale + ) noise = torch.exp(self._last_params["log_noise"]) + self.minimum_noise return _SampledGP( kernel_func=lambda x1, x2: self._kernel_func(x1, x2, lengthscale, 2.5), @@ -303,7 +329,7 @@ class PreferentialGPSampler(optuna.samplers.BaseSampler): ) -> None: self.lengthscale_prior = lengthscale_prior or gpytorch.priors.GammaPrior(5.0, 10.0) self.noise_prior = noise_prior or gpytorch.priors.GammaPrior(5.0, 50.0) - + self._rng = np.random.RandomState(seed) self.independent_sampler = independent_sampler or optuna.samplers.RandomSampler( seed=self._rng.randint(2**32) @@ -351,7 +377,7 @@ class PreferentialGPSampler(optuna.samplers.BaseSampler): lengthscale_prior=self.lengthscale_prior, minimum_lengthscale=0.1, noise_prior=self.noise_prior, - minimum_noise=1e-6, # To avoid NaN + minimum_noise=1e-6, # To avoid NaN dims=len(trans.bounds), ) if self._gp.dims != len(trans.bounds):