mirror of
https://github.com/wassname/Volt.git
synced 2026-09-09 11:16:09 +08:00
62 lines
2.3 KiB
Python
62 lines
2.3 KiB
Python
import torch
|
|
|
|
from torch.distributions import Normal
|
|
from gpytorch.constraints import Positive, Interval
|
|
from gpytorch.likelihoods import Likelihood, _OneDimensionalLikelihood
|
|
|
|
|
|
class VolatilityGaussianLikelihood(_OneDimensionalLikelihood):
|
|
def __init__(self, K=5, batch_shape=torch.Size(), param="cv", *args, **kwargs):
|
|
"""
|
|
parameterization of gaussian likelihood for volatility models like in
|
|
wilson & ghahramani, copula processes, eq. 21.
|
|
|
|
we also consider the gp-exp parameterization
|
|
"""
|
|
|
|
super().__init__()
|
|
if param == "cv":
|
|
self.raw_a = torch.nn.Parameter(torch.rand(*batch_shape, K, requires_grad=True))
|
|
raw_b_init = 0.1 * torch.rand(*batch_shape, K)
|
|
self.raw_b = torch.nn.Parameter(raw_b_init.detach_().requires_grad_())
|
|
self.raw_c = torch.nn.Parameter(torch.rand(*batch_shape, K, requires_grad=True))
|
|
|
|
self.register_constraint("raw_a", Positive())
|
|
self.register_constraint("raw_b", Interval(0.0, 3.0))
|
|
self.register_constraint("raw_c", Interval(-3.0, 3.0))
|
|
# elif param == "exp":
|
|
# print("Using gp-exp parameterization.")
|
|
self.param = param
|
|
|
|
@property
|
|
def trans_a(self):
|
|
return self.raw_a_constraint.transform(self.raw_a)
|
|
|
|
@property
|
|
def trans_b(self):
|
|
return self.raw_b_constraint.transform(self.raw_b)
|
|
|
|
@property
|
|
def trans_c(self):
|
|
return self.raw_c_constraint.transform(self.raw_c)
|
|
|
|
def forward(self, function_samples, *args, **kwargs):
|
|
if self.param == "cv":
|
|
transform = (
|
|
(self.trans_b * function_samples.unsqueeze(-1) + self.trans_c).exp() + 1
|
|
).log() * self.trans_a
|
|
summed_transform = transform.sum(-1)
|
|
else:
|
|
summed_transform = function_samples.exp()
|
|
return Normal(torch.zeros_like(summed_transform), summed_transform.clamp(min=1e-3))
|
|
|
|
def expected_log_prob(self, target, input, *params, **kwargs):
|
|
res = super().expected_log_prob(target, input, *params, **kwargs)
|
|
num_event_dim = len(input.event_shape)
|
|
if num_event_dim > 1:
|
|
res = res.sum(-1)
|
|
return res
|
|
|
|
|
|
# TODO: use a multitask Gaussian likelihood somehow in the multitask setting
|