Deploy_ConvNext / convnext_kan.py
Sirius16's picture
Upload 28 files
4dc60af verified
Raw
History Blame Contribute Delete
4.77 kB
import torch
import torch.nn as nn
from torchvision.models import convnext_tiny, ConvNeXt_Tiny_Weights
from kan import KANLinear
class ConvNeXtFeatureExtractor(nn.Module):
def __init__(
self,
freeze_backbone: bool = True,
unfreeze_last_stage: bool = False,
):
super().__init__()
weights = ConvNeXt_Tiny_Weights.DEFAULT
self.backbone = convnext_tiny(weights=weights)
if freeze_backbone:
for p in self.backbone.parameters():
p.requires_grad = False
if unfreeze_last_stage:
for p in self.backbone.features[-1].parameters():
p.requires_grad = True
self.output_dim = self.backbone.classifier[2].in_features
self.backbone.classifier = nn.Identity()
def forward(self, x):
x = self.backbone(x)
x = x.view(x.size(0), -1)
return x
class KANHead(nn.Module):
def __init__(
self,
input_dim: int,
num_classes: int,
head_depth: int = 2,
hidden_dim_1: int = 512,
hidden_dim_2: int = 256,
head_style: str = "standard",
):
super().__init__()
self.head_depth = head_depth
self.head_style = head_style
if head_style == "rakyan":
if head_depth == 2:
self.kan1 = KANLinear(input_dim, hidden_dim_1)
self.bn1 = nn.BatchNorm1d(hidden_dim_1)
self.act1 = nn.ReLU(inplace=True)
self.kan2 = KANLinear(hidden_dim_1, num_classes)
elif head_depth == 3:
self.kan1 = KANLinear(input_dim, hidden_dim_1)
self.bn1 = nn.BatchNorm1d(hidden_dim_1)
self.act1 = nn.ReLU(inplace=True)
self.drop1 = nn.Dropout(p=0.3)
self.kan2 = KANLinear(hidden_dim_1, hidden_dim_2)
self.bn2 = nn.BatchNorm1d(hidden_dim_2)
self.act2 = nn.ReLU(inplace=True)
self.drop2 = nn.Dropout(p=0.3)
self.kan3 = KANLinear(hidden_dim_2, num_classes)
else:
raise ValueError("head_depth must be 2 or 3")
elif head_style == "standard":
if head_depth == 2:
self.pre_norm = nn.LayerNorm(input_dim)
self.kan1 = KANLinear(input_dim, hidden_dim_1)
self.norm1 = nn.LayerNorm(hidden_dim_1)
self.kan2 = KANLinear(hidden_dim_1, num_classes)
elif head_depth == 3:
self.pre_norm = nn.LayerNorm(input_dim)
self.kan1 = KANLinear(input_dim, hidden_dim_1)
self.norm1 = nn.LayerNorm(hidden_dim_1)
self.kan2 = KANLinear(hidden_dim_1, hidden_dim_2)
self.norm2 = nn.LayerNorm(hidden_dim_2)
self.kan3 = KANLinear(hidden_dim_2, num_classes)
else:
raise ValueError("head_depth must be 2 or 3")
else:
raise ValueError("head_style must be 'standard' or 'rakyan'")
def forward(self, x):
if self.head_style == "rakyan":
x = self.kan1(x); x = self.bn1(x); x = self.act1(x); x = self.drop1(x)
if self.head_depth == 2:
x = self.kan2(x)
else:
x = self.kan2(x); x = self.bn2(x); x = self.act2(x); x = self.drop2(x)
x = self.kan3(x)
return x
x = self.pre_norm(x)
x = self.kan1(x)
x = self.norm1(x)
if self.head_depth == 2:
x = self.kan2(x)
else:
x = self.kan2(x)
x = self.norm2(x)
x = self.kan3(x)
return x
class ConvNextKAN(nn.Module):
def __init__(
self,
num_classes: int = 10,
head_depth: int = 2,
hidden_dim_1: int = 512,
hidden_dim_2: int = 256,
freeze_backbone: bool = True,
unfreeze_last_stage: bool = False,
head_style: str = "standard"
):
super().__init__()
self.feature_extractor = ConvNeXtFeatureExtractor(
freeze_backbone=freeze_backbone,
unfreeze_last_stage=unfreeze_last_stage
)
self.head = KANHead(
input_dim=self.feature_extractor.output_dim,
num_classes=num_classes,
head_depth=head_depth,
hidden_dim_1=hidden_dim_1,
hidden_dim_2=hidden_dim_2,
head_style=head_style
)
def forward(self, x):
features = self.feature_extractor(x)
out = self.head(features)
return out