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