| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from torchvision import models |
|
|
| |
| class MetaNet(nn.Module): |
| """ |
| Implementação da abordagem MetaNet |
| Fusing Metadata and Dermoscopy Images for Skin Disease Diagnosis - https://ieeexplore.ieee.org/document/9098645 |
| """ |
| def __init__(self, in_channels, middle_channels, out_channels): |
| super(MetaNet, self).__init__() |
| self.metanet = nn.Sequential( |
| nn.Conv2d(in_channels, middle_channels, 1), |
| nn.ReLU(), |
| nn.Conv2d(middle_channels, out_channels, 1), |
| nn.Sigmoid() |
| ) |
|
|
| def forward(self, feat_maps, metadata): |
| |
| |
| metadata = metadata.unsqueeze(-1).unsqueeze(-1) |
| |
| x = self.metanet(metadata) |
| |
| x = x * feat_maps |
| return x |
|
|
| |
| class MetaBlock(nn.Module): |
| """ |
| Implementação do Metadata Processing Block (MetaBlock) |
| """ |
| def __init__(self, V, U): |
| """ |
| V: número de canais de features da imagem (ex.: 1664 da DenseNet-169) |
| U: dimensão dos metadados (ex.: 85) |
| """ |
| super(MetaBlock, self).__init__() |
| self.fb = nn.Sequential(nn.Linear(U, V), nn.LayerNorm(V)) |
| self.gb = nn.Sequential(nn.Linear(U, V), nn.LayerNorm(V)) |
|
|
| def forward(self, img_features, metadata): |
| |
| |
| t1 = self.fb(metadata) |
| t2 = self.gb(metadata) |
| |
| t1 = t1.unsqueeze(-1).unsqueeze(-1) |
| t2 = t2.unsqueeze(-1).unsqueeze(-1) |
| |
| out = torch.sigmoid(torch.tanh(img_features * t1) + t2) |
| return out |
|
|
| |
| class MDNet(nn.Module): |
| def __init__(self, meta_dim=85, num_classes=6, cnn_model_name="densenet169", text_model_name="one-hot-encode", hidden_dim=128, device="cpu", unfreeze_weights=False): |
| super(MDNet, self).__init__() |
| self.device = device |
| self.num_channels = 1664 |
| self.meta_dim = meta_dim |
| self.num_classes = num_classes |
| self.cnn_model_name=cnn_model_name |
| self.text_model_name = text_model_name |
| |
| densenet = models.densenet169(pretrained=True) |
| |
| for param in densenet.parameters(): |
| param.requires_grad = unfreeze_weights |
| self.feature_extractor = densenet.features |
| |
| |
| self.meta_net = MetaNet(in_channels=meta_dim, middle_channels=hidden_dim, out_channels=self.num_channels) |
| |
| self.meta_block = MetaBlock(V=self.num_channels, U=meta_dim) |
| |
| |
| self.avg_pool = nn.AdaptiveAvgPool2d((1, 1)) |
| self.classifier = nn.Linear(self.num_channels, self.num_classes) |
| |
| def forward(self, image, metadata): |
| |
| |
| |
| image_features = self.feature_extractor(image) |
| |
| |
| meta_net_features = self.meta_net(image_features, metadata) |
| |
| |
| meta_block_features = self.meta_block(image_features, metadata) |
| |
| |
| fused_features = meta_net_features + meta_block_features |
| |
| |
| pooled_features = self.avg_pool(fused_features) |
| pooled_features = pooled_features.view(pooled_features.size(0), -1) |
| out = self.classifier(pooled_features) |
| return out |
|
|