This commit is contained in:
Contramundum committed 2023-08-30 10:05:43 +09:00
1 parent 28062b648a
commit 81ea33368e
1 file changed
+46 -20
+46 -20
View File
@@ -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):