代码语法测试

技术笔记 #代码 #测试 约 1 分钟 · 10 字

代码语法测试

测试代码

from typing import List
import torch.nn
import torch.nn as nn
from .util import RadialBasisFunction, SplineLinear
from utils import L1
class FastKANLayer(nn.Module):
def __init__(
self,
input_dim: int,
output_dim: int,
grid_min: float = -2.,
grid_max: float = 2.,
num_grids: int = 8,
use_base_update: bool = True,
base_activation=nn.SiLU,
spline_weight_init_scale: float = 0.1,
) -> None:
super().__init__()
self.layernorm = nn.LayerNorm(input_dim)
self.rbf = RadialBasisFunction(grid_min, grid_max, num_grids)
self.spline_linear = SplineLinear(input_dim * num_grids, output_dim, spline_weight_init_scale)
self.use_base_update = use_base_update
if use_base_update:
self.base_activation = base_activation()
self.base_linear = nn.Linear(input_dim, output_dim)
def forward(self, x, time_benchmark=True):
if not time_benchmark:
spline_basis = self.rbf(self.layernorm(x))
else:
spline_basis = self.rbf(x)
ret = self.spline_linear(spline_basis.view(*spline_basis.shape[:-2], -1))
if self.use_base_update:
base = self.base_linear(self.base_activation(x))
ret = ret + base
return ret
class FastKAN(nn.Module):
def __init__(
self,
layers_hidden: List[int],
dropout: float = 0.0,
l1_decay: float = 0.0,
grid_range: List[float] = [-2, 2],
grid_size: int = 8,
use_base_update: bool = True,
base_activation=nn.SiLU,
spline_weight_init_scale: float = 0.1,
first_dropout: bool = True, **kwargs
) -> None:
super().__init__()
self.layers_hidden = layers_hidden
self.grid_min = grid_range[0]
self.grid_max = grid_range[1]
self.use_base_update = use_base_update
self.base_activation = base_activation
self.spline_weight_init_scale = spline_weight_init_scale
self.num_layers = len(layers_hidden[:-1])
self.layers = nn.ModuleList([])
if dropout > 0 and first_dropout:
self.layers.append(nn.Dropout(p=dropout))
for i, (in_features, out_features) in enumerate(zip(layers_hidden[:-1], layers_hidden[1:])):
self.layers.append(torch.nn.BatchNorm1d(in_features))
# Base weight for linear transformation in each layer
layer = FastKANLayer(in_features, out_features,
grid_min=self.grid_min,
grid_max=self.grid_max,
num_grids=grid_size,
use_base_update=use_base_update,
base_activation=base_activation,
spline_weight_init_scale=spline_weight_init_scale)
if l1_decay > 0 and i != self.num_layers - 1:
layer = L1(layer, l1_decay)
self.layers.append(layer)
if dropout > 0 and i != self.num_layers - 1:
self.layers.append(nn.Dropout(p=dropout))
def forward(self, x):
for layer in self.layers:
x = layer(x)
return x