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