forked from baumgach/ralis
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathvisualise_regions.py
More file actions
211 lines (178 loc) · 11.1 KB
/
Copy pathvisualise_regions.py
File metadata and controls
211 lines (178 loc) · 11.1 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
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
"""
Code for highlighting regions selected by the RL agent in images for ACDC and BraTS2018
Author:
Carina Schmidt
"""
import os
import numpy as np
import matplotlib.pyplot as plt
import matplotlib.colors as colors
import matplotlib.patches.Rectangle as Rectangle
import utils.parser as parser
from data import acdc, acdc_al, msdHeart, brats18_2D, brats18_2D_al, brats18, brats18_al
def main(args):
dataset = args.dataset
####------ Create segmentation, query and target networks ------####
# create colormap for matplotlib
#plt.rcParams['axes.prop_cycle'] = plt.cycler(color=["#000000","#FFFF00","#FF0000","#10A5F5"])
plt.rcParams['font.size'] = '8'
myColors = colors.ListedColormap(['black', 'blue', 'gold', 'magenta'])
if args.dataset == 'acdc':
train_set = acdc_al.ACDC_al('fine', 'train',
data_path=args.data_path,
code_path=args.code_path,
joint_transform=None,
transform=None,
target_transform=None, num_each_iter=args.num_each_iter,
only_last_labeled=None,
split='train' if args.al_algorithm == 'ralis' and not args.test else 'test', #if --train and --test, still test=TRUE
region_size=args.region_size)
elif args.dataset == 'brats18':
myColors = colors.ListedColormap(['black', 'darkred', 'orange', 'antiquewhite'])
print("al algo: ", args.al_algorithm)
train_set = brats18_2D_al.BraTS18_2D_al('fine', 'train',
data_path=args.data_path,
code_path=args.code_path,
joint_transform=None,
transform=None,
target_transform=None, num_each_iter=args.num_each_iter,
only_last_labeled=args.only_last_labeled,
split='train' if args.al_algorithm == 'ralis' and not args.test else 'test',
region_size=args.region_size)
list_images = train_set.imgs
n_ep = 0
# dict of image indexes already labelled as key, coordinates as value
idx_region_dict = {}
#labeled set file
if args.dataset == 'acdc':
ralis_path = open(os.path.join(args.ckpt_path, args.exp_name_toload, 'labeled_set_' + str(n_ep) + '.txt'), 'r')
entropy_path = open(os.path.join('/mnt/qb/baumgartner/cschmidt77_data/exp2_acdc_baselines_preTrainDT/2021-11-03-acdc_ImageNetBackbone_baseline_entropy_budget_2384_seed_123/', 'labeled_set_0.txt'), 'r')
bald_path = open(os.path.join('/mnt/qb/baumgartner/cschmidt77_data/exp2_acdc_baselines_preTrainDT/2021-11-04-acdc_ImageNetBackbone_baseline_bald_budget_2384_seed_123/', 'labeled_set_0.txt'), 'r')
random_path = open(os.path.join('/mnt/qb/baumgartner/cschmidt77_data/exp2_acdc_baselines_preTrainDT/2021-11-03-acdc_ImageNetBackbone_baseline_random_budget_2384_seed_123/', 'labeled_set_' + str(n_ep) + '.txt'), 'r')
elif args.dataset == 'brats18': # 'brats18'
n_ep = 36
ralis_path = open(os.path.join(args.ckpt_path, args.exp_name_toload, 'labeled_set_' + str(n_ep) + '.txt'), 'r') #ckpt_path: /mnt/qb/baumgartner/cschmidt77_data/exp1b_brats_baselines/2021-11-04-brats18_ImageNetBackbone_baseline_random_budget_17792_seed_123
entropy_path = open(os.path.join('/mnt/qb/baumgartner/cschmidt77_data/exp1b_brats_baselines/2021-11-04-brats18_ImageNetBackbone_baseline_entropy_budget_17792_seed_123/', 'labeled_set_0.txt'), 'r')
bald_path = open(os.path.join('/mnt/qb/baumgartner/cschmidt77_data/exp1b_brats_baselines/2021-11-04-brats18_ImageNetBackbone_baseline_bald_budget_17792_seed_234', 'labeled_set_0.txt'), 'r')
random_path = open(os.path.join('/mnt/qb/baumgartner/cschmidt77_data/exp1b_brats_baselines/2021-11-04-brats18_ImageNetBackbone_baseline_random_budget_17792_seed_123/', 'labeled_set_0.txt'), 'r')
print("ralis_path: ", ralis_path)
file_paths = [ralis_path, entropy_path, bald_path, random_path]
for i, alalgo in enumerate(file_paths):
print("alalgo: ", alalgo)
file_path = alalgo
print("i: ", i)
if i == 0:
al = 'ralis'
elif i == 1:
al = 'entropy'
elif i == 2:
al = 'bald'
elif i == 3:
al = 'random'
else:
print("AL algo not recognised")
for line in file_path:
# get img indices and region coordinates from labelled set
img_idx, coord_x_left_upper, coord_y_left_upper = line.rstrip('\n').split(',') #removes \n and splits by ,
img_idx, coord_x_left_upper, coord_y_left_upper = int(img_idx), int(coord_x_left_upper), int(coord_y_left_upper)
# add selected img idx to dict
# if idx already added, append with new region coordinates, else create new key
if img_idx in idx_region_dict.keys():
idx_region_dict[img_idx].append((coord_x_left_upper, coord_y_left_upper))
else:
idx_region_dict[img_idx] = [(coord_x_left_upper, coord_y_left_upper)]
v_minimum = 0.5
v_maximum = 0.5
i = 0
# get intesity values
for key, values in idx_region_dict.items():
img_idx = key
img_path, mask_path, img_name = list_images[img_idx]
print("img_path: ", img_path)
img, mask = np.load(img_path), np.load(mask_path)
vmin, vmax = img.min(), img.max()
if vmin < v_minimum:
v_minimum = vmin
if vmax > v_maximum:
v_maximum = vmax
i +=1
if i == 80:
break
print("vminimum: ", v_minimum)
print("vmaximum: ", v_maximum)
# iterate over dictionary with key: idx image, values: list of coordinate pairs
i = 0
for key, values in idx_region_dict.items():
img_idx = key
img_path, mask_path, img_name = list_images[img_idx]
print("img_path: ", img_path)
img, mask = np.load(img_path), np.load(mask_path)
if args.dataset == 'brats18':
img = img[:,:,1]
print("shape of img: ", img.shape)
coordinate_pairs = values
masked = np.full(mask.shape, 0)
image_masked = img
fig, (ax1,ax2) = plt.subplots(1, 2, figsize=(15,7)) #,ax3,ax4)
#ax1.imshow(img, cmap=cm.Greys_r, vmin=v_minimum, vmax=v_maximum)
ax1.imshow(img, cmap=cm.Greys_r, vmin=0.9*v_minimum, vmax=0.9*v_maximum) #mask
ax1.set_title("MRI slice with selected regions")
ax1.axis('off')
ax2.imshow(mask, cmap=myColors) #mask
ax2.set_title("Segmentation with selected regions")
ax2.axis('off')
cropped_regions = []
cropped_masks = []
# for each region coordinates
names_list = []
for pair in coordinate_pairs:
coord_x_left_upper = pair[1]
coord_y_left_upper = pair[0]
coord_x_right_bottom = coord_x_left_upper + args.region_size[1] #region size is here 64, 48
coord_y_right_bottom = coord_y_left_upper + args.region_size[0] #region size 40 or 48
# crop regions
region_img = img[coord_y_left_upper: coord_y_right_bottom, coord_x_left_upper: coord_x_right_bottom]
region_mask = mask[coord_y_left_upper: coord_y_right_bottom, coord_x_left_upper: coord_x_right_bottom]
cropped_regions.append(region_img)
cropped_masks.append(region_mask)
# mask out region of image
masked[coord_y_left_upper: coord_y_right_bottom, coord_x_left_upper: coord_x_right_bottom] = region_mask
image_masked[coord_y_left_upper: coord_y_right_bottom, coord_x_left_upper: coord_x_right_bottom] = region_img
# Create a Rectangle patch
coord_x_left_lower =coord_x_left_upper
coord_y_left_lower =coord_y_left_upper + args.region_size[0]
rect = Rectangle((coord_x_left_lower, coord_y_left_lower),args.region_size[1],args.region_size[0],linewidth=0.8,edgecolor='lime',facecolor='none')
# Add the patch to the Axes
ax1.add_patch(rect)
rect = Rectangle((coord_x_left_lower,coord_y_left_lower),args.region_size[1],args.region_size[0],linewidth=0.8,edgecolor='lime',facecolor='none')
ax2.add_patch(rect)
name = os.path.join('/home/carina/Desktop/regions_visualisation_RALIS_DQN/brats18-dqn', al, str(i) + "_" + img_name) #str(i))#
print("name: ", name)
plt.savefig(f'{name}.png', bbox_inches='tight')
i += 1
alalgo.close()
def rc_params():
plt.rc('text', usetex=True)
plt.rc('font', **{'family': 'serif', 'sans-serif': ['lmodern'], 'size': 20})
plt.rc('axes', **{'titlesize': 22, 'labelsize': 22})
plt.rc('xtick', **{'labelsize': 18})
plt.rc('ytick', **{'labelsize': 18})
plt.rc('legend', **{'fontsize': 18})
plt.rc('figure', **{'figsize': (12,7)})
# Experiment to load for presentation:
# /mnt/qb/baumgartner/cschmidt77_data/exp4_acdc_train_DT_small/2021-10-26-train_acdc_ImageNetBackbone_budget_128_lr_0.05_2patients_seed123
# labeled_set_0.txt
# labeled_set_49.txt
# #/mnt/qb/baumgartner/cschmidt77_data/ACDC_regionsize_3232/2021-11-19-acdc_3232_train_2patients_ImageNetBackbone_budget_512_lr_0.05_seed_123
# singularity exec --nv --bind '/mnt/qb/baumgartner/cschmidt77_data/' '/home/carina/tue-slurm-helloworld/ralis.sif' python3 -u '/home/carina/ralis/visualise_regions.py'
# --exp-name-toload '2021-11-19-acdc_3232_train_2patients_ImageNetBackbone_budget_512_lr_0.05_seed_123' --checkpointer --ckpt-path '/mnt/qb/baumgartner/cschmidt77_data/ACDC_regionsize_3232/'
# --data-path '/mnt/qb/baumgartner/cschmidt77_data/' --input-size 128 128 --dataset 'acdc' --al-algorithm 'ralis' --region-size 32 32
# --train-batch-size 2 --val-batch-size 1 --exp-name-toload-rl '2021-07-17-train_acdc_ImageNetBackbone_budget_608_lr_0.01_seed_123' --num-each-iter 1 --rl-pool 30 --test
# trained DQN:
#/mnt/qb/baumgartner/cschmidt77_data/exp1b_brats_baselines/2021-10-31-brats18_ImageNetBackbone_stdAug_budget_1536_lr_0.01_seed_55
# brats
# singularity exec --nv --bind '/mnt/qb/baumgartner/cschmidt77_data/' '/home/carina/tue-slurm-helloworld/ralis.sif' python3 -u '/home/carina/ralis/visualise_regions.py' --exp-name-toload '2021-10-31-brats18_ImageNetBackbone_stdAug_budget_1536_lr_0.01_seed_55' --checkpointer --ckpt-path '/mnt/qb/baumgartner/cschmidt77_data/exp1b_brats_baselines/' --data-path '/mnt/qb/baumgartner/cschmidt77_data/' --input-size 128 128 --dataset 'brats18' --al-algorithm 'ralis' --region-size 40 48 --train --test --final-test
if __name__ == '__main__':
####------ Parse arguments from console ------####
args = parser.get_arguments()
main(args)