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.
- specunet_pkg/__init__.py +0 -0
- specunet_pkg/__main__.py +5 -0
- specunet_pkg/config/SpecUNet.json +119 -0
- specunet_pkg/dataset.py +297 -0
- specunet_pkg/hyperparameter_search.py +34 -0
- specunet_pkg/logger.py +59 -0
- specunet_pkg/main.py +193 -0
- specunet_pkg/metrics.py +798 -0
- specunet_pkg/models.py +92 -0
- specunet_pkg/praser.py +265 -0
- specunet_pkg/spectra_crop.ipynb +1019 -0
- specunet_pkg/spectra_crop.py +722 -0
- specunet_pkg/spectra_crop_from_locs.py +769 -0
- specunet_pkg/test.py +549 -0
- specunet_pkg/train.py +248 -0
- specunet_pkg/utils.py +263 -0
- specunet_pkg-1.0.0.dist-info/METADATA +346 -0
- specunet_pkg-1.0.0.dist-info/RECORD +20 -0
- specunet_pkg-1.0.0.dist-info/WHEEL +5 -0
- specunet_pkg-1.0.0.dist-info/top_level.txt +1 -0
specunet_pkg/__init__.py
ADDED
|
File without changes
|
specunet_pkg/__main__.py
ADDED
|
@@ -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
|
+
}
|
specunet_pkg/dataset.py
ADDED
|
@@ -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
|