-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathattribute_extractors.py
More file actions
executable file
·79 lines (65 loc) · 2.4 KB
/
Copy pathattribute_extractors.py
File metadata and controls
executable file
·79 lines (65 loc) · 2.4 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
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
from third_party.mgn import MGN
import torch
from torchvision import transforms
import numpy as np
import cv2
from PIL import Image
import os
from constants import INPUT_RESOLUTION, PER_CHANNEL_MEAN, PER_CHANNEL_STD
def ndarraytopil(img):
"""Return a PIL image of an ndarray"""
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
return Image.fromarray(img)
class MgnWrapper:
"""
This class is a wrapper class for the attribute extractor MGN
Attributes:
model (MGN): MGN model architecture
transform (Compose): transformations done to an image such as reshaping and normalizing
"""
def __init__(self, weights_path):
"""
The constructor for MgnWrapper class
Parameters:
weights_path (str): MGN model weights path (MGN.pt)
"""
if not os.path.exists(weights_path):
raise ValueError(
"Weights path given {} doesn't exist".format(weights_path))
self.model = MGN()
self.model.load_state_dict(torch.load(weights_path))
self.model.cuda()
self.model.eval()
self.transform = transforms.Compose([
transforms.Resize(INPUT_RESOLUTION, interpolation=Image.BILINEAR),
transforms.ToTensor(),
transforms.Normalize(mean=PER_CHANNEL_MEAN, std=PER_CHANNEL_STD)
])
def compute_feat_vector(self, inputs):
"""
Uses model to compute the feature vector given an image
Parameters:
inputs (Image or ndarray): PIL Image or ndarray of an image
Returns:
ndarray: The features extracted from MGN of the given image
"""
if isinstance(inputs, np.ndarray):
inputs = ndarraytopil(inputs)
inputs = self.transform(inputs).float()
ff = torch.FloatTensor(inputs.size(0), 2048).zero_()
inputs = inputs.unsqueeze(0)
for i in range(2):
if i == 1:
inputs = inputs.index_select(
3,
torch.arange(inputs.size(3) - 1, -1, -1).long())
input_img = inputs.to('cuda')
outputs = self.model(input_img)
f = outputs[0].data.cpu()
ff = ff + f
fnorm = torch.norm(ff, p=2, dim=1, keepdim=True)
ff = ff.div(fnorm.expand_as(ff))
return ff
def __call__(self, x):
""" Return the feature vector of an image x """
return self.compute_feat_vector(x)