Deploy_ConvNext / kan.py
Sirius16's picture
Upload 28 files
4dc60af verified
Raw
History Blame Contribute Delete
4.62 kB
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