mirror of
https://github.com/wassname/Volt.git
synced 2026-08-21 11:11:37 +08:00
32 lines
1.1 KiB
Python
32 lines
1.1 KiB
Python
from gpytorch.kernels import Kernel
|
|
from gpytorch.functions import RBFCovariance
|
|
from gpytorch.settings import trace_mode
|
|
|
|
def postprocess_rbf(dist_mat):
|
|
return dist_mat.div_(-2).exp_()
|
|
|
|
class OUKernel(Kernel):
|
|
has_lengthscale = True
|
|
|
|
def forward(self, x1, x2, diag=False, **params):
|
|
if (
|
|
x1.requires_grad
|
|
or x2.requires_grad
|
|
or (self.ard_num_dims is not None and self.ard_num_dims > 1)
|
|
or diag
|
|
or params.get("last_dim_is_batch", False)
|
|
or trace_mode.on()
|
|
):
|
|
x1_ = x1.div(self.lengthscale)
|
|
x2_ = x2.div(self.lengthscale)
|
|
return self.covar_dist(
|
|
x1_, x2_, square_dist=False, diag=diag, dist_postprocess_func=postprocess_rbf, postprocess=True, **params
|
|
)
|
|
return RBFCovariance.apply(
|
|
x1,
|
|
x2,
|
|
self.lengthscale,
|
|
lambda x1, x2: self.covar_dist(
|
|
x1, x2, square_dist=False, diag=False, dist_postprocess_func=postprocess_rbf, postprocess=False, **params
|
|
),
|
|
) |