specunet-pkg 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
File without changes
@@ -0,0 +1,5 @@
1
+ # specunet_pkg/__main__.py
2
+ from .main import main
3
+
4
+ if __name__ == "__main__":
5
+ main()
@@ -0,0 +1,119 @@
1
+ {
2
+ "experiment_name": "Default",
3
+ "device_args": {
4
+ "distributed": false,
5
+ "device": "cuda",
6
+ "gpu_ids": [
7
+ 0
8
+ ]
9
+ },
10
+ "phase": {
11
+ "train": false,
12
+ "test_sim": false,
13
+ "test_exp": false
14
+ },
15
+ "exp_path": {
16
+ "base_dir": "test_models",
17
+ "save_tiff_images": false,
18
+ "save_model_onnx": false,
19
+ "sim_results": "sim_rep_samples",
20
+ "exp_results": "exp_rep_samples",
21
+ "configs": "configs",
22
+ "spectrum_metrics": "spectrum_metrics",
23
+ "peak_metrics": "peak_metrics"
24
+ },
25
+ "datasets": {
26
+ "data_type": {
27
+ "type": "spectral",
28
+ "norm": true,
29
+ "input_name": "sptimg4",
30
+ "target_name": "tbg4",
31
+ "GTspt": "GTspt",
32
+ "spt": "spt",
33
+ "seed": 42
34
+ },
35
+ "train": {
36
+ "args": {
37
+ "data_root": "data/Sample_TrainingData_10000.mat",
38
+ "data_len": 0,
39
+ "percent": false
40
+ },
41
+ "dataloader": {
42
+ "validation_split": 0,
43
+ "args": {
44
+ "batch_size": 16,
45
+ "num_workers": 0,
46
+ "shuffle": true,
47
+ "pin_memory": false
48
+ },
49
+ "val_args": {
50
+ "batch_size": 16,
51
+ "num_workers": 0,
52
+ "shuffle": false,
53
+ "pin_memory": false
54
+ }
55
+ }
56
+ },
57
+ "test_sim": {
58
+ "args": {
59
+ "data_root": "data/TestingData.mat",
60
+ "data_len": 0
61
+ },
62
+ "dataloader": {
63
+ "args": {
64
+ "batch_size": 32,
65
+ "num_workers": 0,
66
+ "pin_memory": false,
67
+ "shuffle": false
68
+ }
69
+ }
70
+ },
71
+ "test_exp": {
72
+ "args": {
73
+ "data_root": "data/ExpTestingData.mat",
74
+ "data_len": 0
75
+ },
76
+ "dataloader": {
77
+ "args": {
78
+ "batch_size": 16,
79
+ "num_workers": 0,
80
+ "pin_memory": false,
81
+ "shuffle": false
82
+ }
83
+ }
84
+ }
85
+ },
86
+ "model": {
87
+ "name": "all",
88
+ "hyperparameters": {
89
+ "optimizer": "adam",
90
+ "epochs": 250,
91
+ "lr": 0.001,
92
+ "weight_decay": 1e-05,
93
+ "lr_scheduler": "step",
94
+ "initializer": "none",
95
+ "scheduler_args": {
96
+ "step_size": 50,
97
+ "gamma": 0.2,
98
+ "T_max": 250,
99
+ "pct_start": 0.1
100
+ }
101
+ },
102
+ "input_size": [
103
+ 1,
104
+ 16,
105
+ 128
106
+ ],
107
+ "loss_fn": "mse",
108
+ "loss_fn_args": {
109
+ "tv_weight": 1e-05
110
+ },
111
+
112
+ "metrics": {
113
+ "image_wise": ["RMSE", "PSNR", "SSIM"],
114
+ "localization_wise": [],
115
+ "spectral_wise": ["Spearman rho", "Pearson r", "Centroid GT (nm)", "Centroid Raw (nm)", "Centroid % Error"],
116
+ "summary_output": ["mean_centroid_rawspt", "std_centroid_rawspt", "mean_centroid_spef", "std_centroid_spef"]
117
+ }
118
+ }
119
+ }
@@ -0,0 +1,297 @@
1
+ import torch
2
+ import h5py
3
+ import numpy as np
4
+ from torch.utils.data import Dataset, random_split, DataLoader
5
+ from scipy.io import loadmat
6
+ import sys
7
+
8
+ class SpecUNet_Dataset(Dataset):
9
+ def __init__(self, X, Y=None, GTspt=None, spt=None):
10
+ self.X = X
11
+ self.Y = Y
12
+ self.GTspt = GTspt
13
+ self.spt = spt
14
+
15
+ def __len__(self):
16
+ return len(self.X)
17
+
18
+ def __getitem__(self, idx):
19
+ x_tensor = torch.tensor(self.X[idx], dtype=torch.float64)
20
+
21
+ # For missing Y or GTspt, return zero tensors with expected shape
22
+ if self.Y is None:
23
+ y_tensor = torch.zeros_like(x_tensor) # or another shape as needed
24
+ else:
25
+ y_tensor = torch.tensor(self.Y[idx], dtype=torch.float64)
26
+
27
+ if self.GTspt is None:
28
+ gt_tensor = torch.zeros_like(x_tensor)
29
+ else:
30
+ gt_tensor = torch.tensor(self.GTspt[idx], dtype=torch.float64)
31
+
32
+ return x_tensor, y_tensor, gt_tensor
33
+
34
+ def get_spt(self):
35
+ return self.spt
36
+
37
+ def normalize_dataset(data, input_shape):
38
+ """
39
+ Normalizes dataset shape.
40
+ - If input is image data (3D/4D): returns (N, 1, H, W).
41
+ - If input is spectral data (2D): returns (N, Features).
42
+
43
+ Parameters
44
+ ----------
45
+ data : np.ndarray
46
+ Input dataset array (image or spectral curves).
47
+ input_shape : tuple
48
+ Expected single-sample image shape (e.g., (1, 128, 16)).
49
+ Used primarily for image dimension validation.
50
+
51
+ Returns
52
+ -------
53
+ np.ndarray
54
+ Normalized dataset.
55
+ """
56
+ if not isinstance(data, np.ndarray):
57
+ data = np.array(data)
58
+
59
+ # --- CASE A: Handle 2D Data (Spectral Vectors) ---
60
+ # Target: (N, 301)
61
+ if data.ndim == 2:
62
+ # Check if data is in (Features, Samples) format, e.g., (301, 5000)
63
+ # We assume if dim0 is 301 (your spectral channels), it needs rotation.
64
+ # OR if dim1 is significantly larger than dim0, it's likely (Feat, Samp).
65
+ if data.shape[0] == 301 or (data.shape[1] > data.shape[0]):
66
+ # print(f"2D Input detected {data.shape}: Transposing to (N, Features).")
67
+ return data.T
68
+
69
+ # Already (N, Features)
70
+ return data
71
+
72
+ # --- CASE B: Handle 3D/4D Data (Images) ---
73
+ expected_c, expected_h, expected_w = input_shape
74
+
75
+ # If array is 4D: (N, C, H, W) or mixed
76
+ if data.ndim == 4:
77
+ # Check if we need to swap H and W (e.g. 128 vs 16) inside the 4D array
78
+ # This can happen if data is (N, C, W, H) instead of (N, C, H, W)
79
+ _, _, h, w = data.shape
80
+ if h == expected_w and w == expected_h:
81
+ # print(f"4D Input detected {data.shape}: Swapping H and W axes.")
82
+ data = np.swapaxes(data, 2, 3)
83
+ return data
84
+
85
+ # If array is 3D: Missing channel dimension
86
+ if data.ndim == 3:
87
+ # Get current dimensions
88
+ dim0, dim1, dim2 = data.shape
89
+
90
+ # Case 1: (N, H, W) -> Correct
91
+ if dim1 == expected_h and dim2 == expected_w:
92
+ pass
93
+
94
+ # Case 2: (N, W, H) -> Transpose W and H
95
+ elif dim1 == expected_w and dim2 == expected_h:
96
+ # print(f"3D Input detected {data.shape}: Swapping spatial axes to (N, {expected_h}, {expected_w})")
97
+ data = np.transpose(data, (0, 2, 1))
98
+
99
+ # Case 3: (H, W, N) -> Move samples to front
100
+ elif dim0 == expected_h and dim1 == expected_w:
101
+ data = np.transpose(data, (2, 0, 1))
102
+
103
+ # Case 4: (W, H, N) -> Swap H/W, Move samples to front
104
+ elif dim0 == expected_w and dim1 == expected_h:
105
+ data = np.transpose(data, (2, 1, 0))
106
+
107
+ else:
108
+ # Fallback: If dimensions don't match input_shape, try to guess N
109
+ # This handles cases where input_shape might be (1,16,128) but data is slightly different
110
+ pass
111
+
112
+ # Add channel dimension (N, H, W) -> (N, 1, H, W)
113
+ data = np.expand_dims(data, axis=1)
114
+ return data
115
+
116
+ raise ValueError(f"[dataset] Unsupported array dimensions: {data.ndim}")
117
+
118
+
119
+ def load_matlab_data(file_path, input_shape, input_name, target_name=None, gt_name=None, spt_name=None):
120
+ """Loads MATLAB v5 or v7.3 data and expands dimensions."""
121
+
122
+ # Detect file type from first bytes
123
+ with open(file_path, 'rb') as f:
124
+ header = f.read(8)
125
+
126
+ sptimg4, tbg4, gt_spt, spt = None, None, None, None
127
+
128
+ if header.startswith(b'MATLAB 5'):
129
+ # MATLAB v5/v7.0 binary file
130
+ try:
131
+ mat_data = loadmat(file_path)
132
+ sptimg4 = mat_data[input_name]
133
+ if target_name is not None and target_name in mat_data:
134
+ tbg4 = mat_data[target_name]
135
+ if gt_name is not None and gt_name in mat_data:
136
+ gt_spt = mat_data[gt_name]
137
+ if spt_name is not None and spt_name in mat_data:
138
+ spt = mat_data[spt_name]
139
+ except FileNotFoundError:
140
+ print(f"[dataset] Error: The file {file_path} was not found.")
141
+ sys.exit(0)
142
+ except KeyError as e:
143
+ print(f"[dataset] Error: Required key {e} not found in the MATLAB file. Available dataset keys: {mat_data.keys()}")
144
+ sptimg4 = None
145
+ sys.exit(0)
146
+ except Exception as e:
147
+ print(f"[dataset] An unexpected error occurred: {e}")
148
+ sys.exit(0)
149
+
150
+ elif header.startswith(b'MATLAB 7'):
151
+ # MATLAB v7.3 (HDF5) file
152
+ with h5py.File(file_path, 'r') as f:
153
+ try:
154
+ sptimg4 = f[input_name][:]
155
+ if target_name is not None and target_name in f:
156
+ tbg4 = f[target_name][:]
157
+ if gt_name is not None and gt_name in f:
158
+ gt_spt = f[gt_name][:]
159
+ if spt_name is not None and spt_name in f:
160
+ spt = f[spt_name][:]
161
+ except FileNotFoundError:
162
+ print(f"[dataset] Error: The dataset {file_path} was not found.")
163
+ except KeyError as e:
164
+ print(f"[dataset] Error: Required key {e} not found in the MATLAB file. Available dataset keys: {f.keys()}")
165
+ except Exception as e:
166
+ print(f"[dataset] An unexpected error occurred: {e}")
167
+
168
+ else:
169
+ try:
170
+ with h5py.File(file_path, 'r') as f:
171
+ # print("Dataset keys: ", f.keys())
172
+ sptimg4 = f[input_name][:]
173
+ if target_name is not None and target_name in f:
174
+ tbg4 = f[target_name][:]
175
+ if gt_name is not None and gt_name in f:
176
+ gt_spt = f[gt_name][:]
177
+ if spt_name is not None and spt_name in f:
178
+ spt = f[spt_name][:]
179
+ except FileNotFoundError:
180
+ print(f"[dataset] Error: The dataset {file_path} was not found.")
181
+ except KeyError as e:
182
+ print(f"[dataset] Error: Required key {e} not found in the MATLAB file. Available dataset keys: {f.keys()}")
183
+ except Exception as e:
184
+ print(f"[dataset] An unexpected error occurred: {e}")
185
+
186
+ # Expand dims if loaded
187
+ if sptimg4 is not None:
188
+ # print(sptimg4.shape)
189
+ # print(input_shape)
190
+ sptimg4 = normalize_dataset(sptimg4, input_shape)
191
+ # print(sptimg4.shape)
192
+ if tbg4 is not None:
193
+ # print(tbg4.shape)
194
+ tbg4 = normalize_dataset(tbg4, tuple(input_shape))
195
+ # print(tbg4.shape)
196
+ if gt_spt is not None:
197
+ # print(gt_spt.shape)
198
+ gt_spt = normalize_dataset(gt_spt, tuple(input_shape))
199
+ # print(gt_spt.shape)
200
+ if spt is not None:
201
+ # print(spt.shape)
202
+ spt = normalize_dataset(spt, tuple(input_shape))
203
+ # print(spt.shape)
204
+
205
+ return sptimg4, tbg4, gt_spt, spt
206
+
207
+ def create_dataset(sptimg4, tbg4=None, gt_spt=None, spt=None):
208
+ """Creates a SpecUNet_Dataset."""
209
+ return SpecUNet_Dataset(sptimg4, tbg4, gt_spt, spt)
210
+
211
+ def create_dataloader(dataset, opt):
212
+ """Creates a DataLoader based on training or testing."""
213
+ return DataLoader(dataset, batch_size=opt['batch_size'], shuffle=opt['shuffle'],
214
+ num_workers=opt['num_workers'], pin_memory=opt['pin_memory'])
215
+
216
+
217
+ def get_train_datasets(opt, input_shape):
218
+ """Main data loading script."""
219
+
220
+ train_loader = None
221
+ val_loader = None
222
+
223
+ sptimg4_train, tbg4_train, GTspt_train, _ = load_matlab_data(
224
+ file_path=opt['train']['args']['data_root'],
225
+ input_name=opt['data_type']['input_name'],
226
+ target_name=opt['data_type']['target_name'],
227
+ gt_name=opt['data_type']['GTspt'],
228
+ input_shape=input_shape
229
+ )
230
+
231
+ train_dataset = create_dataset(sptimg4_train, tbg4_train, GTspt_train)
232
+
233
+ validation_split = opt['train']['dataloader']['validation_split']
234
+ if validation_split > 0:
235
+ val_size = int(validation_split * len(train_dataset))
236
+ train_size = len(train_dataset) - val_size
237
+
238
+ print(f"[dataset] "
239
+ f"Validation split of {validation_split} used. Splitting selected training dataset into training and "
240
+ f"validation datasets of sizes {train_size} and {val_size}, respectively.")
241
+
242
+ train_dataset, val_dataset = random_split(train_dataset, [train_size, val_size])
243
+
244
+ train_loader = create_dataloader(train_dataset, opt['train']['dataloader']['args'])
245
+ val_loader = create_dataloader(val_dataset, opt['train']['dataloader']['val_args'])
246
+
247
+ print("[dataset] Loaded training and validation datasets!")
248
+
249
+ return train_loader, val_loader
250
+
251
+ else:
252
+ sptimg4_test, tbg4_test, GTspt_test, _ = load_matlab_data(
253
+ file_path=opt['test_sim']['args']['data_root'],
254
+ input_name=opt['data_type']['input_name'],
255
+ target_name=opt['data_type']['target_name'],
256
+ gt_name=opt['data_type']['GTspt'],
257
+ input_shape=input_shape
258
+ )
259
+
260
+ test_dataset = create_dataset(sptimg4_test, tbg4_test, GTspt_test)
261
+
262
+ print(
263
+ f"[dataset] No validation split passed in. Using testing dataset as validation dataset "
264
+ f"for training and validation datasets of sizes {len(train_dataset)} and {len(test_dataset)}"
265
+ f", respectively.")
266
+
267
+ train_loader = create_dataloader(train_dataset, opt['train']['dataloader']['args'])
268
+ val_loader = create_dataloader(test_dataset, opt['train']['dataloader']['val_args'])
269
+
270
+ return train_loader, val_loader
271
+
272
+ def get_test_datasets(opt_dataset, input_shape, opt_phase):
273
+ if opt_phase['test_sim']:
274
+ sptimg4_test, tbg4_test, GTspt_test, spt_test = load_matlab_data(
275
+ file_path=opt_dataset['test_sim']['args']['data_root'],
276
+ input_name=opt_dataset['data_type']['input_name'],
277
+ target_name=opt_dataset['data_type']['target_name'],
278
+ gt_name=opt_dataset['data_type']['GTspt'],
279
+ spt_name=opt_dataset['data_type']['spt'],
280
+ input_shape=input_shape
281
+ )
282
+ test_dataset = create_dataset(sptimg4_test, tbg4_test, GTspt_test, spt_test)
283
+ print(f"[dataset] Using test_sim dataset of size {len(test_dataset)}")
284
+ return test_dataset
285
+
286
+ else:
287
+ sptimg4_test, tbg4_test, GTspt_test, _ = load_matlab_data(
288
+ file_path=opt_dataset['test_exp']['args']['data_root'],
289
+ input_name=opt_dataset['data_type']['input_name'],
290
+ input_shape=input_shape)
291
+ test_dataset = create_dataset(sptimg4_test, tbg4_test, GTspt_test)
292
+ print(f"Using test_exp dataset of size {len(test_dataset)}")
293
+ return test_dataset
294
+
295
+
296
+
297
+
@@ -0,0 +1,34 @@
1
+ import itertools
2
+ import subprocess
3
+ import datetime
4
+
5
+ # Define hyperparameter values
6
+ lr_values = [0.01, 0.0075, 0.005]
7
+ batch_sizes = [50]
8
+ weight_decay_values = [0.01, 0.0075, 0.005]
9
+ scheduler_gamma_values = [0.2]
10
+
11
+ train_path = "Training_Perlin40k_Pperlin50k_MatchNorm.mat"
12
+ train_size = 80000
13
+ train_test_split = 0.8
14
+ loss_function = 'mse'
15
+ epochs = 250
16
+
17
+ # Generate all combinations
18
+ hyperparameter_combinations = list(itertools.product(lr_values, batch_sizes,
19
+ weight_decay_values, scheduler_gamma_values))
20
+
21
+ for lr, batch_size, weight_decay, scheduler_gamma in hyperparameter_combinations:
22
+ exp_name = f"Test_{lr}lr_{batch_size}bs_{weight_decay}l2reg_{scheduler_gamma}gamma_{loss_function}"
23
+ command = (f"python -m specunet_pkg main.py --train --test_exp "
24
+ f"--exp_name {exp_name} "
25
+ f"--train_path {train_path} "
26
+ f"--train_size {train_size} "
27
+ f"--train_test_split {train_test_split}"
28
+ f"--lr {lr} --bs {batch_size} "
29
+ f"--weight_decay {weight_decay} --scheduler_gamma {scheduler_gamma} "
30
+ f"--loss_fn {loss_function} "
31
+ f"--epochs {epochs} ")
32
+ subprocess.run(command, shell=True)
33
+
34
+ print("Finished at time: ", datetime.datetime.now())
specunet_pkg/logger.py ADDED
@@ -0,0 +1,59 @@
1
+ import logging
2
+ import sys
3
+ import os
4
+
5
+ LOG_FILE = "log.log"
6
+
7
+ # Define a function to set up and return a logger
8
+ def get_logger(path, name="main"):
9
+ logger = logging.getLogger(name)
10
+
11
+ # Prevent adding multiple handlers in case of multiple imports
12
+ if not logger.hasHandlers():
13
+ logger.setLevel(logging.INFO)
14
+
15
+ # Explicitly set encoding='utf-8'
16
+ file_handler = logging.FileHandler(os.path.join(path, LOG_FILE), encoding='utf-8')
17
+
18
+ file_handler.setFormatter(logging.Formatter(
19
+ "%(asctime)s - %(levelname)s - %(message)s",
20
+ datefmt="%m/%d/%Y %I:%M:%S %p"
21
+ ))
22
+
23
+ # Console handler
24
+ console_handler = logging.StreamHandler(sys.stdout)
25
+ console_handler.setFormatter(logging.Formatter("%(levelname)s - %(message)s"))
26
+
27
+ # Add handlers
28
+ logger.addHandler(file_handler)
29
+ logger.addHandler(console_handler)
30
+
31
+ # Redirect print statements to logger
32
+ sys.stdout = StreamToLogger(logger, logging.INFO)
33
+ sys.stderr = StreamToLogger(logger, logging.ERROR) # Redirect stderr as well
34
+
35
+ return logger
36
+
37
+ def log_print(logger, *args, **kwargs):
38
+ """
39
+ Prints and logs the message using the provided logger.
40
+ """
41
+ message = " ".join(str(arg) for arg in args) # Convert args to string
42
+ logger.info(message) # Log the message
43
+
44
+ # Helper class to redirect stdout to the logger
45
+ class StreamToLogger:
46
+ def __init__(self, logger, log_level):
47
+ self.logger = logger
48
+ self.log_level = log_level
49
+ self.line_buffer = ""
50
+
51
+ def write(self, message):
52
+ if message.strip(): # Avoid logging empty lines
53
+ self.logger.log(self.log_level, message.strip())
54
+
55
+ def flush(self):
56
+ pass # No need to flush manually; logging handles it
57
+
58
+ def isatty(self):
59
+ return False