File size: 7,844 Bytes
a639402 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 | from torchvision import models
import torch.nn as nn
import torch
from efficientnet_pytorch import EfficientNet
from vit_pytorch.crossformer import CrossFormer
from vit_pytorch.efficient import ViT
from nystrom_attention import Nystromformer
import timm
def create_soat_cnn_models(num_classes, model, pretrained, drop_out):
if model == 'vgg19':
model_ft = models.vgg19(pretrained=pretrained)
number_features = model_ft.classifier[6].in_features
if drop_out>0:
model_ft.classifier[6] = nn.Sequential(nn.Dropout(drop_out), nn.Linear(number_features, num_classes))
else:
model_ft.classifier[6] = nn.Linear(number_features, num_classes)
elif model == 'vgg19_bn':
model_ft = models.vgg19_bn(pretrained=pretrained)
number_features = model_ft.classifier[6].in_features
if drop_out>0:
model_ft.classifier[6] = nn.Sequential(nn.Dropout(drop_out), nn.Linear(number_features, num_classes))
else:
model_ft.classifier[6] = nn.Linear(number_features, num_classes)
elif model == 'resnet18':
model_ft = models.resnet18(pretrained=pretrained)
number_features = model_ft.fc.in_features
if drop_out>0:
model_ft.fc = nn.Sequential(nn.Dropout(drop_out), nn.Linear(number_features, num_classes))
else:
model_ft.fc = nn.Linear(number_features, num_classes)
elif model == 'resnet50':
model_ft = models.resnet50(pretrained=pretrained)
number_features = model_ft.fc.in_features
if drop_out>0:
model_ft.fc = nn.Sequential(nn.Dropout(drop_out), nn.Linear(number_features, num_classes))
else:
model_ft.fc = nn.Sequential(nn.Linear(number_features, num_classes))
elif model == 'resnet152':
model_ft = models.resnet152(pretrained=pretrained)
number_features = model_ft.fc.in_features
if drop_out>0:
model_ft.fc = nn.Sequential(nn.Dropout(drop_out), nn.Linear(number_features, num_classes))
else:
model_ft.fc = nn.Linear(number_features, num_classes)
elif model == 'densenet121':
model_ft = models.densenet121(pretrained=pretrained)
#print(model_ft)
number_features = model_ft.classifier.in_features
if drop_out>0:
model_ft.classifier = nn.Sequential(nn.Dropout(drop_out), nn.Linear(number_features, num_classes))
else:
model_ft.classifier = nn.Linear(number_features, num_classes)
elif model == 'densenet161':
model_ft = models.densenet161(pretrained=pretrained)
number_features = model_ft.classifier.in_features
if drop_out>0:
model_ft.classifier = nn.Sequential(nn.Dropout(drop_out), nn.Linear(number_features, num_classes))
else:
model_ft.classifier = nn.Linear(number_features, num_classes)
elif model == 'skinfold_efficientnet':
## Dummy Implementation of a New Efficient Based Classifier; Half Image + Full Image + Half Image ;
## Average performance; Maybe needs location of the Skinfold.
model_ft = SkinFoldEfficientNet(num_classes=num_classes, drop_out=drop_out)
elif model == 'efficientnet_b0':
## Author : Github : https://github.com/lukemelas/EfficientNet-PyTorch
## Repo with ImageNet Pretrained Models.
## Some models in Torchvision are also ported from this repo; https://github.com/pytorch/vision/tree/main/references/classification#efficientnet-v1
model_ft = EfficientNet.from_pretrained('efficientnet-b0', num_classes=num_classes, dropout_rate = drop_out)
elif model == 'efficientnet_b3':
## Author : Github : https://github.com/lukemelas/EfficientNet-PyTorch
## Repo with ImageNet Pretrained Models.
## Some models in Torchvision are also ported from this repo; https://github.com/pytorch/vision/tree/main/references/classification#efficientnet-v1
model_ft = EfficientNet.from_pretrained('efficientnet-b3', num_classes=num_classes, dropout_rate = drop_out)
elif model == 'crossformer':
## Not Implemented
model_ft = CrossFormer(
num_classes = num_classes, # number of output classes
dim = (64, 128, 256, 512), # dimension at each stage
depth = (2, 2, 8, 2), # depth of transformer at each stage
global_window_size = (8, 4, 2, 1), # global window sizes at each stage
local_window_size = 7, # local window size (can be customized for each stage, but in paper, held constant at 7 for all stages)
dropout = drop_out,
)
elif model =='max_vit':
model_ft = timm.create_model('maxvit_base_tf_512.in21k_ft_in1k',
pretrained=pretrained,
pretrained_cfg_overlay=dict(file="/home/z003ve3e/MammographyAnalysis/external_model/maxvit_base_tf_512.in21k_ft_in1k/model.safetensors"),
num_classes=num_classes,
drop_rate=drop_out
)
elif model == 'efficient_transformer':
## Not Implemented
efficient_transformer = Nystromformer(
dim = 512,
depth = 12,
heads = 8,
num_landmarks = 256
)
model_ft = ViT(
dim = 512,
image_size = 1024,
patch_size = 32,
num_classes = num_classes,
transformer = efficient_transformer,
dropout = drop_out
)
else:
NotImplementedError("Model not implemented. Check src.models_all.models.py")
return model_ft
class SkinFoldEfficientNet(nn.Module):
## Dummy Architecture; No performance gains.
def __init__(self, num_classes, drop_out):
super(SkinFoldEfficientNet, self).__init__()
# Main EfficientNet B0 for full image processing
self.main_model = EfficientNet.from_pretrained('efficientnet-b0', dropout_rate = drop_out)
self.main_model._fc = nn.Linear(self.main_model._fc.in_features, 512)
# EfficientNet B0 for top half image processing
self.top_model = EfficientNet.from_pretrained('efficientnet-b0', dropout_rate = drop_out)
self.top_model._fc = nn.Linear(self.top_model._fc.in_features, 512)
# EfficientNet B0 for bottom half image processing
self.bottom_model = EfficientNet.from_pretrained('efficientnet-b0', dropout_rate = drop_out)
self.bottom_model._fc = nn.Linear(self.bottom_model._fc.in_features, 512)
self.fc = nn.Sequential(nn.Linear(1536, num_classes))
def forward(self, x):
# Split the input image into top and bottom halves
top_half = x[:, :, :512, :]
bottom_half = x[:, :, 512:, :]
main_output = self.main_model(x)
top_output = self.top_model(top_half)
bottom_output = self.bottom_model(bottom_half)
concatenated = torch.cat((main_output, top_output, bottom_output), dim=1)
# Forward pass through the fully connected layer
output = self.fc(concatenated)
return output
|