| import os |
| import logging |
| import numpy as np |
| from typing import List |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| import clip |
|
|
| from metrics.base_metrics_class import calculate_acc_for_train |
| from .base_detector import AbstractDetector |
| from detectors import DETECTOR |
| from loss import LOSSFUNC |
|
|
| logger = logging.getLogger(__name__) |
|
|
| |
| class LinearClassifier(nn.Module): |
| def __init__(self, input_size: int, hidden_size_list: List[int], num_classes: int): |
| super(LinearClassifier, self).__init__() |
| self.dropout = nn.Dropout(0.5) |
| self.fc1 = nn.Linear(input_size, hidden_size_list[0]) |
| self.fc2 = nn.Linear(hidden_size_list[0], hidden_size_list[1]) |
| self.fc3 = nn.Linear(hidden_size_list[1], num_classes) |
|
|
| def forward(self, x: torch.tensor) -> torch.tensor: |
| out = self.fc1(x) |
| out = F.relu(out) |
| out = self.dropout(out) |
| out = self.fc2(out) |
| out = F.relu(out) |
| out = self.fc3(out) |
| return out |
|
|
| @DETECTOR.register_module(module_name='clip_image') |
| class CLIPImageDetector(AbstractDetector): |
| def __init__(self, config): |
| super().__init__(config) |
| self.config = config |
| self.device = torch.device("cuda" if config['cuda'] else "cpu") |
| self.backbone = self.build_backbone(config) |
| self.classifier_module = self.build_classifier(config) |
| self.loss_func = self.build_loss(config) |
|
|
| def build_backbone(self, config) -> clip.model.CLIP: |
| |
| clip_model, _ = clip.load(config['clip_model_name'], device=self.device) |
| logger.info(f"Loaded CLIP image encoder: {config['clip_model_name']}") |
| return clip_model |
|
|
| def build_loss(self, config) -> nn.CrossEntropyLoss: |
| |
| loss_class = LOSSFUNC[config['loss_func']] |
| loss_func = loss_class() |
| return loss_func.to(self.device) |
|
|
| def features(self, data_dict: dict) -> torch.tensor: |
| |
| images = data_dict['image'].to(self.device) |
| |
| with torch.no_grad(): |
| image_emb = self.backbone.encode_image(images) |
| return image_emb.float() |
|
|
| def classifier(self, features: torch.tensor) -> torch.tensor: |
| |
| return self.classifier_module(features) |
|
|
| def get_losses(self, data_dict: dict, pred_dict: dict) -> dict: |
| |
| labels = data_dict['label'].to(self.device) |
| preds = pred_dict['cls'] |
| loss = self.loss_func(preds, labels) |
| return {'overall': loss} |
|
|
| def get_train_metrics(self, data_dict: dict, pred_dict: dict) -> dict: |
| |
| labels = data_dict['label'].detach().cpu() |
| preds = pred_dict['cls'].detach().cpu() |
| num_classes = self.config['classifier_config']['num_classes'] |
| acc, mAP = calculate_acc_for_train(labels, preds, num_classes) |
| return {'acc': acc, 'mAP': mAP} |
|
|
| def forward(self, data_dict: dict, inference=False) -> dict: |
| |
| features = self.features(data_dict) |
| cls_pred = self.classifier(features) |
| prob = torch.softmax(cls_pred, dim=1) |
| return {'cls': cls_pred, 'prob': prob, 'feat': features} |
|
|
| def build_classifier(self, config) -> LinearClassifier: |
| |
| input_size = 512 |
| hidden_size_list = config['classifier_config']['hidden_size_list'] |
| num_classes = config['classifier_config']['num_classes'] |
| classifier = LinearClassifier(input_size, hidden_size_list, num_classes) |
| return classifier.to(self.device) |
| |
|
|