-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathGoT_app.py
More file actions
41 lines (33 loc) · 1.12 KB
/
Copy pathGoT_app.py
File metadata and controls
41 lines (33 loc) · 1.12 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
import argparse
from GoT_utils import *
from GoT_model import *
# Parsing command line arguments
arg_p = argparse.ArgumentParser()
arg_p.add_argument('-image_path', default='sample/')
arg_p.add_argument('-model_weights', default='saved_models/ver2.0_weights_final.hdf5')
args = vars(arg_p.parse_args())
IMAGE_PATH = args['image_path']
MODEL_WEIGHTS = args['model_weights']
sample = load_samples(IMAGE_PATH)
def GoT_algo(img_path):
"""Function that takes a image(s) via file path(s) and makes predictions using the final CNN model.
Outputs a plot of image(s) with the predicted label (+ for GoT, - for Not GoT)
Arguments:
img_path: File path of image(s) used to make a prediction.
"""
fig = plt.figure()
for i, image in enumerate(img_path):
test_img = path_to_tensor(image).astype('float32')/255
result = np.argmax(CNN(test_img, MODEL_WEIGHTS))
img = cv2.imread(image)
cv_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
ax = fig.add_subplot(4,5,i+1)
if result == 1:
plt.title('+')
else:
plt.title('-')
plt.axis('off')
plt.imshow(cv_rgb)
plt.suptitle('Sample Images')
plt.show()
GoT_algo(sample)