mirror of
https://github.com/wassname/Volt.git
synced 2026-08-21 11:11:37 +08:00
28 lines
1.1 KiB
Python
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) |