import torch import torch.nn as nn import torch.nn.functional as F import math class KANLinear(nn.Module): def __init__( self, in_features: int, out_features: int, grid_size: int = 5, spline_order: int = 3, scale_noise: float = 0.1, scale_base: float = 1.0, scale_spline: float = 1.0, enable_standalone_scale_spline: bool = True, base_activation: nn.Module = nn.SiLU(), grid_eps: float = 0.02, grid_range: list = [-1, 1], ): super(KANLinear, self).__init__() self.in_features = in_features self.out_features = out_features self.grid_size = grid_size self.spline_order = spline_order self.grid_eps = grid_eps h = (grid_range[1] - grid_range[0]) / grid_size grid = ( ( torch.arange(-spline_order, grid_size + spline_order + 1) * h + grid_range[0] ) .expand(in_features, -1) .contiguous() ) self.register_buffer("grid", grid) self.base_weight = nn.Parameter(torch.Tensor(out_features, in_features)) self.spline_weight = nn.Parameter( torch.Tensor(out_features, in_features, grid_size + spline_order) ) if enable_standalone_scale_spline: self.spline_scaler = nn.Parameter( torch.Tensor(out_features, in_features) ) else: self.register_parameter("spline_scaler", None) self.scale_noise = scale_noise self.scale_base = scale_base self.scale_spline = scale_spline self.enable_standalone_scale_spline = enable_standalone_scale_spline self.base_activation = base_activation self.reset_parameters() def reset_parameters(self): nn.init.kaiming_uniform_(self.base_weight, a=math.sqrt(5) * self.scale_base) with torch.no_grad(): noise = ( ( torch.rand(self.grid_size + 1, self.in_features, self.out_features) - 1 / 2 ) * self.scale_noise / self.grid_size ) self.spline_weight.data.copy_( (self.scale_spline if not self.enable_standalone_scale_spline else 1.0) * self.curve2coeff( self.grid.T[self.spline_order : -self.spline_order], noise, ) ) if self.enable_standalone_scale_spline: nn.init.kaiming_uniform_(self.spline_scaler, a=math.sqrt(5) * self.scale_spline) def b_splines(self, x: torch.Tensor): assert x.dim() == 2 and x.size(1) == self.in_features grid = self.grid x = x.unsqueeze(-1) bases = ((x >= grid[:, :-1]) & (x < grid[:, 1:])).to(x.dtype) for k in range(1, self.spline_order + 1): bases = ( (x - grid[:, : -(k + 1)]) / (grid[:, k:-1] - grid[:, : -(k + 1)]) * bases[:, :, :-1] ) + ( (grid[:, k + 1 :] - x) / (grid[:, k + 1 :] - grid[:, 1:(-k)]) * bases[:, :, 1:] ) assert bases.size() == (x.size(0), self.in_features, self.grid_size + self.spline_order) return bases.contiguous() def curve2coeff(self, x: torch.Tensor, y: torch.Tensor): assert x.dim() == 2 and x.size(1) == self.in_features assert y.size() == (x.size(0), self.in_features, self.out_features) A = self.b_splines(x).transpose(0, 1) B = y.transpose(0, 1) solution = torch.linalg.lstsq(A, B).solution result = solution.permute(2, 0, 1) assert result.size() == (self.out_features, self.in_features, self.grid_size + self.spline_order) return result.contiguous() def forward(self, x: torch.Tensor): assert x.dim() == 2 and x.size(1) == self.in_features base_output = F.linear(self.base_activation(x), self.base_weight) spline_output = F.linear( self.b_splines(x).view(x.size(0), -1), self.spline_weight.view(self.out_features, -1), ) if self.enable_standalone_scale_spline: spline_output = spline_output * self.spline_scaler.unsqueeze(0).mean(dim=2) return base_output + spline_output