Skip to content

Commit cfbf79d

Browse files
committed
Upload New File
1 parent 9794312 commit cfbf79d

1 file changed

Lines changed: 130 additions & 0 deletions

File tree

utils.py

Lines changed: 130 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,130 @@
1+
import os
2+
import shutil
3+
import itertools
4+
import tifffile as tiff
5+
6+
import numpy as np
7+
8+
import torch
9+
10+
EPS = 1e-12
11+
12+
13+
def normalization(image0, image1, image2, cfg):
14+
if cfg.normalization == 'min_max':
15+
min_, max_ = cfg.min_max
16+
image0 = (image0 - min_)/(max_ - min_)
17+
image1 = (image1 - min_)/(max_ - min_)
18+
image2 = (image2 - min_)/(max_ - min_)
19+
20+
21+
elif cfg.normalization == 'zscore':
22+
image0 = (image0 - image0.float().mean()) / max(image0.float().std(), EPS)
23+
image1 = (image1 - image1.float().mean()) / max(image1.float().std(), EPS)
24+
image2 = (image2 - image2.float().mean()) / max(image2.float().std(), EPS)
25+
26+
return image0, image1, image2
27+
28+
29+
def create_zdataset(list_of_files:str, cfg:dict):
30+
"""
31+
Generates a 'datasets' folder containing one dataset per channel. Each dataset comprises PyTorch tensors,
32+
with each tensor representing three consecutive 2D slices extracted from the input 5D image array.
33+
34+
Args:
35+
list_of_files (str): File paths of the input images.
36+
cfg (dict): Configuration options including dataset size, model name, and validation split ratio.
37+
38+
Returns:
39+
None
40+
"""
41+
dataset_size = cfg.dataset_size
42+
model_name = cfg.model_name
43+
p_val = cfg.p_val
44+
45+
triplets = np.empty((0, 5), dtype=int)
46+
47+
root = '.'
48+
for i, file in enumerate(list_of_files):
49+
50+
path = f"{root}/data/train/{file}"
51+
image = tiff.imread(path)
52+
53+
format = 'y'
54+
len_shape = len(image.shape)
55+
if len_shape <= 2:
56+
raise ValueError('Error: images must be at least 3D.')
57+
elif len_shape == 3:
58+
image = image[None, :, None, :, :]
59+
elif len_shape == 4:
60+
if format == 'y':
61+
image = image[:, :, None, :, :]
62+
elif format == 'n':
63+
image = image[None, :, :, :, :]
64+
else:
65+
raise ValueError('Error: data file-format not recognized.')
66+
else:
67+
raise ValueError('Error: images must be at most 5D.')
68+
69+
t, z = image.shape[0:2]
70+
71+
if model_name == 'zaugnet+':
72+
all_triplets = np.array(list(itertools.combinations(range(z), 3)))
73+
selected_all_triplets = []
74+
for idx in range(len(all_triplets)):
75+
if (all_triplets[idx][2] - all_triplets[idx][0]) <= cfg.distance_triplets :
76+
selected_all_triplets.append(all_triplets[idx])
77+
all_triplets = np.array(selected_all_triplets)
78+
79+
for dt in range(t):
80+
rand_ = np.random.randint(0, len(all_triplets), dataset_size)
81+
82+
selected_triplets = np.concatenate((np.repeat([[i,dt]], dataset_size, axis=0), all_triplets[rand_]), axis=1)
83+
triplets = np.concatenate((triplets, selected_triplets))
84+
85+
86+
else :
87+
for dt, dz in itertools.product(range(t), range(z - 2)):
88+
triplets = np.concatenate((triplets, np.array([(i, dt, dz, dz+1, dz+2)])))
89+
90+
91+
np.random.shuffle(triplets)
92+
triplets_train = triplets[:int((1-p_val)*len(triplets))]
93+
triplets_val = triplets[int((1-p_val)*len(triplets)):]
94+
95+
if cfg.save_dataset :
96+
if os.path.exists('./dataset/'):
97+
shutil.rmtree('./dataset/')
98+
save_data(f'{root}/data/train/', f'{root}/dataset/train_{cfg.model_name}/', triplets_train)
99+
save_data(f'{root}/data/train/', f'{root}/dataset/val_{cfg.model_name}/', triplets_val)
100+
101+
102+
def save_data(data_path, dataset_path, triplets):
103+
os.makedirs(dataset_path, exist_ok=True)
104+
105+
files = os.listdir(data_path)
106+
for file_name in set(list(triplets[:,0])):
107+
selected_triplets = triplets[np.where(triplets[:,0] == file_name)]
108+
image = tiff.imread(f"{data_path}{files[file_name]}")
109+
110+
format = 'y'
111+
len_shape = len(image.shape)
112+
if len_shape <= 2:
113+
raise ValueError('Error: images must be at least 3D.')
114+
elif len_shape == 3:
115+
image = image[None, :, None, :, :]
116+
elif len_shape == 4:
117+
if format == 'y':
118+
image = image[:, :, None, :, :]
119+
elif format == 'n':
120+
image = image[None, :, :, :, :]
121+
else:
122+
raise ValueError('Error: data file-format not recognized.')
123+
else:
124+
raise ValueError('Error: images must be at most 5D.')
125+
126+
image = torch.from_numpy(image) #.to(torch.uint16)
127+
image = image.to(torch.int32) # for torch.uint16 problem nuclei
128+
for tri in selected_triplets:
129+
torch.save(image[tri[1], tri[2:]], f"{dataset_path}{tri[0]}_{tri[1]}_{tri[2]}_{tri[3]}_{tri[4]}.pt")
130+

0 commit comments

Comments
 (0)