[WIP] simplify gp implementation

This commit is contained in:
Contramundum committed 2023-08-25 11:39:15 +09:00
1 parent 140ca3957b
commit ce38bf54f6
1 file changed
+214 -271
+214 -271
View File
@@ -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
)