Merge pull request #24 from johnlockejrr/unifying-training-models

Unifying training models
This commit is contained in:
Clemens Neudecker 2025-06-03 09:00:56 +02:00 committed by GitHub
commit d6ccb83bf5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 9 additions and 9 deletions

View file

@ -434,7 +434,7 @@ def generate_arrays_from_folder_reading_order(classes_file_dir, modal_dir, batch
batchcount = 0
while True:
for i in all_labels_files:
file_name = i.split('.')[0]
file_name = os.path.splitext(i)[0]
img = cv2.imread(os.path.join(modal_dir,file_name+'.png'))
label_class = int( np.load(os.path.join(classes_file_dir,i)) )
@ -479,7 +479,7 @@ def data_gen(img_folder, mask_folder, batch_size, input_height, input_width, n_c
for i in range(c, c + batch_size): # initially from 0 to 16, c = 0.
try:
filename = n[i].split('.')[0]
filename = os.path.splitext(n[i])[0]
train_img = cv2.imread(img_folder + '/' + n[i]) / 255.
train_img = cv2.resize(train_img, (input_width, input_height),
@ -745,7 +745,7 @@ def provide_patches(imgs_list_train, segs_list_train, dir_img, dir_seg, dir_flow
indexer = 0
for im, seg_i in tqdm(zip(imgs_list_train, segs_list_train)):
img_name = im.split('.')[0]
img_name = os.path.splitext(im)[0]
if task == "segmentation" or task == "binarization":
dir_of_label_file = os.path.join(dir_seg, img_name + '.png')
elif task=="enhancement":