| from typing import Any |
|
|
| import numpy as np |
| import torch |
| import torch.nn as nn |
| from numpy.linalg import norm, pinv |
| from scipy.special import logsumexp |
| from sklearn.covariance import EmpiricalCovariance |
| from tqdm import tqdm |
|
|
| from .base_postprocessor import BasePostprocessor |
|
|
|
|
| class VIMPostprocessor(BasePostprocessor): |
| def __init__(self, config): |
| super().__init__(config) |
| self.args = self.config.postprocessor.postprocessor_args |
| self.args_dict = self.config.postprocessor.postprocessor_sweep |
| self.dim = self.args.dim |
| self.setup_flag = False |
|
|
| def setup(self, net: nn.Module, id_loader_dict, ood_loader_dict): |
| if not self.setup_flag: |
| net.eval() |
|
|
| with torch.no_grad(): |
| self.w, self.b = net.get_fc() |
| print('Extracting id training feature') |
| feature_id_train = [] |
| for batch in tqdm(id_loader_dict['train'], |
| desc='Setup: ', |
| position=0, |
| leave=True): |
| data = batch['data'].cuda() |
| data = data.float() |
| _, feature = net(data, return_feature=True) |
| feature_id_train.append(feature.cpu().numpy()) |
| feature_id_train = np.concatenate(feature_id_train, axis=0) |
| logit_id_train = feature_id_train @ self.w.T + self.b |
|
|
| self.u = -np.matmul(pinv(self.w), self.b) |
| ec = EmpiricalCovariance(assume_centered=True) |
| ec.fit(feature_id_train - self.u) |
| eig_vals, eigen_vectors = np.linalg.eig(ec.covariance_) |
| self.Vim = np.ascontiguousarray( |
| (eigen_vectors.T[np.argsort(eig_vals * -1)[self.dim:]]).T) |
|
|
| vlogit_id_train = norm(np.matmul(feature_id_train - self.u, |
| self.NS), |
| axis=-1) |
| self.alpha = logit_id_train.max( |
| axis=-1).mean() / vlogit_id_train.mean() |
| print(f'self.alpha = {self.alpha:.4f}') |
|
|
| self.setup_flag = True |
| else: |
| pass |
|
|
| @torch.no_grad() |
| def postprocess(self, net: nn.Module, data: Any): |
| _, feature_ood = net.forward(data, return_feature=True) |
| feature_ood = feature_ood.cpu() |
| logit_ood = feature_ood @ self.w.T + self.b |
| _, pred = torch.max(logit_ood, dim=1) |
| energy_ood = logsumexp(logit_ood.numpy(), axis=-1) |
| vlogit_ood = norm(np.matmul(feature_ood.numpy() - self.u, self.NS), |
| axis=-1) * self.alpha |
| score_ood = -vlogit_ood + energy_ood |
| return pred, torch.from_numpy(score_ood) |
|
|
| def set_hyperparam(self, hyperparam: list): |
| self.dim = hyperparam[0] |
|
|
| def get_hyperparam(self): |
| return self.dim |
|
|