Spaces:
Sleeping
Sleeping
| import torch.nn as nn | |
| from torchvision.models import resnet50, ResNet50_Weights | |
| from kan1 import KANLinear | |
| class ResNetKAN(nn.Module): | |
| def __init__(self, num_classes=10, freeze_backbone=True): | |
| super().__init__() | |
| weights = ResNet50_Weights.DEFAULT | |
| self.resnet = resnet50(weights=weights) | |
| if freeze_backbone: | |
| for p in self.resnet.parameters(): | |
| p.requires_grad = False | |
| for p in self.resnet.layer3.parameters(): | |
| p.requires_grad = True | |
| for p in self.resnet.layer4.parameters(): | |
| p.requires_grad = True | |
| num_features = self.resnet.fc.in_features | |
| self.resnet.fc = nn.Identity() | |
| self.kan1 = KANLinear(num_features, 512) | |
| self.bn1 = nn.BatchNorm1d(512) | |
| self.act1 = nn.ReLU() | |
| self.kan2 = KANLinear(512, 512) | |
| self.bn2 = nn.BatchNorm1d(512) | |
| self.act2 = nn.ReLU() | |
| self.kan3 = KANLinear(512, num_classes) | |
| def forward(self, x): | |
| x = self.resnet(x) | |
| x = x.view(x.size(0), -1) | |
| x = self.kan1(x) | |
| x = self.bn1(x) | |
| x = self.act1(x) | |
| x = self.kan2(x) | |
| x = self.bn2(x) | |
| x = self.act2(x) | |
| x = self.kan3(x) | |
| return x |