-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathcxr_models.py
More file actions
48 lines (34 loc) · 1.49 KB
/
Copy pathcxr_models.py
File metadata and controls
48 lines (34 loc) · 1.49 KB
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
import torch.nn as nn
import torchvision
import torch
import numpy as np
from torch.nn.functional import kl_div, softmax, log_softmax
from .loss import RankingLoss, CosineLoss
import torch.nn.functional as F
class CXRModels(nn.Module):
def __init__(self, args, device='cpu'):
super(CXRModels, self).__init__()
self.args = args
self.device = device
self.vision_backbone = getattr(torchvision.models, self.args.vision_backbone)(pretrained=self.args.pretrained)
classifiers = [ 'classifier', 'fc']
for classifier in classifiers:
cls_layer = getattr(self.vision_backbone, classifier, None)
if cls_layer is None:
continue
d_visual = cls_layer.in_features
setattr(self.vision_backbone, classifier, nn.Identity(d_visual))
break
self.bce_loss = torch.nn.BCELoss(size_average=True)
self.classifier = nn.Sequential(nn.Linear(d_visual, self.args.vision_num_classes))
self.feats_dim = d_visual
def forward(self, x, labels=None, n_crops=0, bs=16):
lossvalue_bce = torch.zeros(1).to(self.device)
visual_feats = self.vision_backbone(x)
preds = self.classifier(visual_feats)
preds = torch.sigmoid(preds)
if n_crops > 0:
preds = preds.view(bs, n_crops, -1).mean(1)
if labels is not None:
lossvalue_bce = self.bce_loss(preds, labels)
return preds, lossvalue_bce, visual_feats