mirror of
https://github.com/wassname/optuna-dashboard.git
synced 2026-10-04 12:50:44 +08:00
format
This commit is contained in:
1 parent
28062b648a
commit
81ea33368e
1 file changed
+46
-20
@@ -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):
|
||||
|
||||
Reference in new issue
Block a user