Files
2022-06-12 11:16:58 -04:00

28 lines
1.1 KiB
Python

import math
import torch
import gpytorch
import numpy as np
class MaternGP(gpytorch.models.ExactGP):
def __init__(self, train_x, train_y, likelihood):
super(MaternGP, self).__init__(train_x, train_y, likelihood)
self.mean_module = gpytorch.means.ConstantMean()
self.covar_module = gpytorch.kernels.ScaleKernel(gpytorch.kernels.MaternKernel())
def forward(self, x):
mean_x = self.mean_module(x)
covar_x = self.covar_module(x)
return gpytorch.distributions.MultivariateNormal(mean_x, covar_x)
class SMGP(gpytorch.models.ExactGP):
def __init__(self, train_x, train_y, likelihood, num_mixtures=10):
super(SMGP, self).__init__(train_x, train_y, likelihood)
self.mean_module = gpytorch.means.ConstantMean()
self.covar_module = gpytorch.kernels.SpectralMixtureKernel(num_mixtures=num_mixtures)
self.covar_module.initialize_from_data(train_x, train_y)
def forward(self,x):
mean_x = self.mean_module(x)
covar_x = self.covar_module(x)
return gpytorch.distributions.MultivariateNormal(mean_x, covar_x)