Fatihaybasn's picture
Publish trained checkpoint, model card, metrics, and provenance
63f5b1f verified
Raw
History Blame Contribute Delete
3.73 kB
from __future__ import annotations
import torch
import torch.nn as nn
import timm
class HybridDN121EffB0(nn.Module):
def __init__(self, num_classes=2, head_dim=256, dropout=0.2):
super().__init__()
self.bb1 = timm.create_model("densenet121", pretrained=False, num_classes=0, global_pool="avg")
self.bb2 = timm.create_model("efficientnet_b0", pretrained=False, num_classes=0, global_pool="avg")
self.head = nn.Sequential(nn.Linear(self.bb1.num_features + self.bb2.num_features, head_dim), nn.GELU(), nn.Dropout(dropout), nn.Linear(head_dim, num_classes))
def forward(self, x):
return self.head(torch.cat([self.bb1(x), self.bb2(x)], dim=1))
class HybridSwinTEffB0(nn.Module):
def __init__(self, num_classes=2, head_dim=256, dropout=0.2):
super().__init__()
self.bb1 = timm.create_model("swin_tiny_patch4_window7_224", pretrained=False, num_classes=0, global_pool="avg")
self.bb2 = timm.create_model("efficientnet_b0", pretrained=False, num_classes=0, global_pool="avg")
self.head = nn.Sequential(nn.Linear(self.bb1.num_features + self.bb2.num_features, head_dim), nn.GELU(), nn.Dropout(dropout), nn.Linear(head_dim, num_classes))
def forward(self, x):
return self.head(torch.cat([self.bb1(x), self.bb2(x)], dim=1))
class SEBlock(nn.Module):
def __init__(self, channels, reduction=16):
super().__init__()
hidden = max(8, channels // reduction)
self.pool = nn.AdaptiveAvgPool2d(1)
self.fc = nn.Sequential(nn.Linear(channels, hidden, bias=False), nn.GELU(), nn.Linear(hidden, channels, bias=False), nn.Sigmoid())
def forward(self, x):
b, c, _, _ = x.shape
return x * self.fc(self.pool(x).view(b, c)).view(b, c, 1, 1)
class Conv1x1BNAct(nn.Module):
def __init__(self, in_ch, out_ch):
super().__init__()
self.net = nn.Sequential(nn.Conv2d(in_ch, out_ch, 1, bias=False), nn.BatchNorm2d(out_ch), nn.GELU())
def forward(self, x):
return self.net(x)
class CustomMSAFEffB0(nn.Module):
def __init__(self, num_classes=2, embed_dim=256, head_dim=256, dropout=0.35):
super().__init__()
self.backbone = timm.create_model("efficientnet_b0", pretrained=False, features_only=True, out_indices=(2, 3, 4))
channels = self.backbone.feature_info.channels()
self.proj = nn.ModuleList([Conv1x1BNAct(c, embed_dim) for c in channels])
self.se = nn.ModuleList([SEBlock(embed_dim) for _ in channels])
self.pool = nn.AdaptiveAvgPool2d(1)
self.scale_attn = nn.Sequential(nn.Linear(embed_dim, max(8, embed_dim // 4)), nn.GELU(), nn.Dropout(dropout * 0.25), nn.Linear(max(8, embed_dim // 4), 1))
self.head = nn.Sequential(nn.Linear(embed_dim * (len(channels) + 1), head_dim), nn.GELU(), nn.Dropout(dropout), nn.Linear(head_dim, num_classes))
def forward(self, x):
vecs = [self.pool(se(proj(f))).flatten(1) for f, proj, se in zip(self.backbone(x), self.proj, self.se)]
ms = torch.stack(vecs, dim=1)
weights = torch.softmax(self.scale_attn(ms).squeeze(-1), dim=1)
return self.head(torch.cat([(ms * weights.unsqueeze(-1)).sum(dim=1), ms.flatten(1)], dim=1))
def build_model(architecture, num_classes=2):
if architecture == "hybrid_dn121_effb0":
return HybridDN121EffB0(num_classes=num_classes)
if architecture == "hybrid_swint_effb0":
return HybridSwinTEffB0(num_classes=num_classes)
if architecture == "custom_msaf_effb0":
return CustomMSAFEffB0(num_classes=num_classes)
return timm.create_model(architecture, pretrained=False, num_classes=num_classes)