From b2777a2562efb0fc73d9f5b92f0ed589dcf5fdc1 Mon Sep 17 00:00:00 2001 From: Robert Sachunsky Date: Thu, 30 Jul 2026 15:28:01 +0200 Subject: [PATCH] do_prediction*: refactor into separate module --- src/eynollah/extract_images.py | 16 +- src/eynollah/eynollah.py | 438 ++++----------------------------- src/eynollah/eynollah_ocr.py | 7 +- src/eynollah/sbb_binarize.py | 7 +- src/eynollah/utils/tiling.py | 394 +++++++++++++++++++++++++++++ tests/test_tiling.py | 108 ++++++++ 6 files changed, 560 insertions(+), 410 deletions(-) create mode 100644 src/eynollah/utils/tiling.py create mode 100644 tests/test_tiling.py diff --git a/src/eynollah/extract_images.py b/src/eynollah/extract_images.py index ed49368..499fa12 100644 --- a/src/eynollah/extract_images.py +++ b/src/eynollah/extract_images.py @@ -1,5 +1,5 @@ """ -extract images? +extract image regions only """ from concurrent.futures import ProcessPoolExecutor @@ -19,6 +19,7 @@ from .model_zoo.model_zoo import EynollahModelZoo from .writer import EynollahXmlWriter from .eynollah import Eynollah from .utils import box2rect, is_image_filename +from .utils.tiling import do_prediction_new_concept from .plot import EynollahPlotter from .utils import Region @@ -110,8 +111,9 @@ class EynollahImageExtractor(Eynollah): img_h_new = img_w_new * img_height_h // img_width_h img_resized = resize_image(img, img_h_new, img_w_new) - prediction_regions, _ = self.do_prediction_new_concept( - True, img_resized, self.model_zoo.get("extract_images")) + prediction_regions, _ = do_prediction_new_concept( + img_resized, self.model_zoo.get("extract_images"), + patches=True, logger=self.logger) prediction_regions = resize_image(prediction_regions, img_height_h, img_width_h) mask_texts_only = (prediction_regions == label_text).astype(np.uint8) @@ -158,17 +160,11 @@ class EynollahImageExtractor(Eynollah): **kwargs ): """ - Get image and scales, then extract the page of scanned image + Get scanned image and scales, then crop, and detect image regions """ self.logger.debug("enter run") # Log enabled features directly enabled_modes = [] - if self.full_layout: - enabled_modes.append("Full layout analysis") - if self.tables: - enabled_modes.append("Table detection") - if enabled_modes: - self.logger.info("Enabled modes: " + ", ".join(enabled_modes)) if self.enable_plotting: self.logger.info("Saving debug plots") if dir_of_cropped_images: diff --git a/src/eynollah/eynollah.py b/src/eynollah/eynollah.py index eea3b6e..be62210 100644 --- a/src/eynollah/eynollah.py +++ b/src/eynollah/eynollah.py @@ -21,7 +21,6 @@ import logging.handlers import sys from difflib import SequenceMatcher as sq -import math import os import time from typing import Optional, List, Tuple @@ -30,7 +29,6 @@ from functools import partial from pathlib import Path import multiprocessing as mp from concurrent.futures import ProcessPoolExecutor, as_completed -import gc import cv2 import numpy as np @@ -65,6 +63,7 @@ from .utils.separate_lines import ( from .utils.marginals import get_marginals from .utils.resize import resize_image from .utils.shm import share_ndarray +from .utils.tiling import do_prediction, do_prediction_new_concept from .utils import ( Region, TextRegion, @@ -77,7 +76,6 @@ from .utils import ( box2slice, find_num_col, otsu_copy_binary, - seg_mask_label, fill_bb_of_drop_capitals, split_textregion_main_vs_head, small_textlines_to_parent_adherence2, @@ -354,10 +352,12 @@ class Eynollah: img_new, _ = fun(img, num_col, conf_col, width_early) if img_new.shape[1] > img.shape[1]: - img_new = self.do_prediction(True, img_new, self.model_zoo.get("enhancement"), - marginal_of_patch_percent=0, - n_batch_inference=3, - is_enhancement=True) + img_new = do_prediction(img_new, self.model_zoo.get("enhancement"), + patches=True, + logger=self.logger, + marginal_of_patch_percent=0, + n_batch_inference=3, + is_enhancement=True) self.logger.info("Enhancement applied") image['img_res'] = img_new @@ -372,7 +372,10 @@ class Eynollah: img = self.imread(image) self.logger.info("Detected %s DPI", dpi) if self.input_binary: - prediction_bin = self.do_prediction(True, img, self.model_zoo.get("binarization"), n_batch_inference=5) + prediction_bin = do_prediction(img, self.model_zoo.get("binarization"), + patches=True, + logger=self.logger, + n_batch_inference=5) prediction_bin = 255 * (prediction_bin == 0) prediction_bin = np.repeat(prediction_bin[:, :, np.newaxis], 3, axis=2).astype(np.uint8) image['img_bin_uint8'] = prediction_bin @@ -436,375 +439,6 @@ class Eynollah: image['scale_x'] = 1.0 * img_res.shape[1] / img.shape[1] return is_image_enhanced, num_col, is_image_resized - def do_prediction( - self, patches, img, model, - n_batch_inference=1, - marginal_of_patch_percent=0.1, - thresholding_for_some_classes=False, - thresholding_for_heading=False, - heading_class=2, - thresholding_for_artificial_class=False, - threshold_art_class=0.1, - artificial_class=2, - is_enhancement=False, - ): - - self.logger.debug("enter do_prediction (patches=%d)", patches) - _, img_height_model, img_width_model, _ = model.input_shape - img_h_page = img.shape[0] - img_w_page = img.shape[1] - - img = img / 255. - img = img.astype(np.float16) - - if not patches: - img = resize_image(img, img_height_model, img_width_model) - - label_p_pred = model.predict(img[np.newaxis], verbose=0)[0] - if is_enhancement: - seg = (label_p_pred * 255).astype(np.uint8) - else: - seg = np.argmax(label_p_pred, axis=2).astype(np.uint8) - - if thresholding_for_artificial_class: - seg_mask_label( - seg, label_p_pred[:, :, artificial_class] >= threshold_art_class, - label=artificial_class, - skeletonize=True) - - if thresholding_for_heading: - seg_mask_label( - seg, label_p_pred[:, :, heading_class] >= 0.2, - label=heading_class) - - return resize_image(seg, img_h_page, img_w_page) - - if img_h_page < img_height_model: - img = resize_image(img, img_height_model, img.shape[1]) - if img_w_page < img_width_model: - img = resize_image(img, img.shape[0], img_width_model) - - self.logger.debug("Patch size: %sx%s", img_height_model, img_width_model) - margin = int(marginal_of_patch_percent * img_height_model) - window = 1 / (1 + np.exp(5.0 - 5 * np.arange(2 * margin) / margin)) - width_mid = img_width_model - 2 * margin - height_mid = img_height_model - 2 * margin - img_h = img.shape[0] - img_w = img.shape[1] - prediction: np.ndarray = None # type: ignore - nxf = math.ceil((img_w - 2.0 * margin) / width_mid) - nyf = math.ceil((img_h - 2.0 * margin) / height_mid) - - batch_i = [] - batch_j = [] - batch_x_u = [] - batch_x_d = [] - batch_x_s = [] - batch_y_u = [] - batch_y_d = [] - batch_y_s = [] - - batch = 0 - img_patch = np.zeros((n_batch_inference, - img_height_model, - img_width_model, - 3), dtype=np.float16) - for i in range(nxf): - for j in range(nyf): - index_x_d = i * width_mid - index_x_u = index_x_d + img_width_model - if index_x_u > img_w: - index_x_s = index_x_u - img_w - index_x_u = img_w - index_x_d = img_w - img_width_model - else: - index_x_s = 0 - index_y_d = j * height_mid - index_y_u = index_y_d + img_height_model - if index_y_u > img_h: - index_y_s = index_y_u - img_h - index_y_u = img_h - index_y_d = img_h - img_height_model - else: - index_y_s = 0 - - batch_i.append(i) - batch_j.append(j) - batch_x_u.append(index_x_u) - batch_x_d.append(index_x_d) - batch_x_s.append(index_x_s) - batch_y_d.append(index_y_d) - batch_y_u.append(index_y_u) - batch_y_s.append(index_y_s) - - img_patch[batch] = img[index_y_d: index_y_u, - index_x_d: index_x_u] - batch += 1 - if (batch == n_batch_inference or - # last batch - i == nxf - 1 and j == nyf - 1): - self.logger.debug("predicting patches on %s", str(img_patch.shape)) - label_p_pred = model.predict(img_patch, verbose=0) - if prediction is None: - # now we know the number of classes - prediction = np.zeros((img_h, img_w, label_p_pred.shape[-1]), dtype=float) - - for batch in range(batch): - where = np.index_exp[batch_y_d[batch]: batch_y_u[batch], - batch_x_d[batch]: batch_x_u[batch]] - # shorter window on last tile - part = np.index_exp[batch_y_s[batch]:, - batch_x_s[batch]:] - # normalize probability (where windows overlap) - attenuation_y = np.ones(img_height_model - batch_y_s[batch]) - attenuation_x = np.ones(img_width_model - batch_x_s[batch]) - if margin and batch_j[batch] > 0: - attenuation_y[:2 * margin] = window - if margin and batch_j[batch] < nyf - 1: - attenuation_y[-2 * margin:] = 1 - window - if margin and batch_i[batch] > 0: - attenuation_x[:2 * margin] = window - if margin and batch_i[batch] < nxf - 1: - attenuation_x[-2 * margin:] = 1 - window - label_p_pred[batch][part] *= attenuation_y[:, np.newaxis, np.newaxis] - label_p_pred[batch][part] *= attenuation_x[np.newaxis, :, np.newaxis] - prediction[where][part] += label_p_pred[batch][part] - - batch_i = [] - batch_j = [] - batch_x_u = [] - batch_x_d = [] - batch_x_s = [] - batch_y_u = [] - batch_y_d = [] - batch_y_s = [] - batch = 0 - img_patch[:] = 0 - - if is_enhancement: - seg = (prediction * 255).astype(np.uint8) - else: - seg = np.argmax(prediction, axis=2).astype(np.uint8) - if thresholding_for_some_classes: - seg_mask_label( - seg, prediction[:, :, 4] > 0.03, - label=4) # - seg_mask_label( - seg, prediction[:, :, 0] > 0.25, - label=0) # bg - seg_mask_label( - seg, prediction[:, :, 3] > 0.10 & seg == 0, - label=3) # line - if thresholding_for_artificial_class: - seg_art = prediction[:, :, artificial_class] >= threshold_art_class - seg_mask_label(seg, seg_art, - label=artificial_class, - only=True, - skeletonize=True, - dilate=3) - - if img_h != img_h_page or img_w != img_w_page: - seg = resize_image(seg, img_h_page, img_w_page) - - gc.collect() - return seg - - def do_prediction_new_concept( - self, patches, img, model, - n_batch_inference=1, - marginal_of_patch_percent=0.1, - thresholding_for_heading=False, - heading_class=2, - thresholding_for_artificial_class=False, - threshold_art_class=0.1, - artificial_class=4, - separator_class=0, - ): - - self.logger.debug("enter do_prediction_new_concept (patches=%d)", patches) - _, img_height_model, img_width_model, _ = model.input_shape - - img = img / 255.0 - img = img.astype(np.float16) - - if not patches: - img_h_page = img.shape[0] - img_w_page = img.shape[1] - img = resize_image(img, img_height_model, img_width_model) - - label_p_pred = model.predict(img[np.newaxis], verbose=0)[0] - seg = np.argmax(label_p_pred, axis=2).astype(np.uint8) - - prediction = resize_image(seg, img_h_page, img_w_page) - - if thresholding_for_artificial_class: - mask = resize_image(label_p_pred[:, :, artificial_class], - img_h_page, img_w_page) >= threshold_art_class - seg_mask_label(prediction, mask, - label=artificial_class, - only=True, - skeletonize=True, - dilate=3, - keep=separator_class) - if thresholding_for_heading: - mask = resize_image(label_p_pred[:, :, heading_class], - img_h_page, img_w_page) >= 0.2 - seg_mask_label(prediction, mask, - label=heading_class) - - conf = label_p_pred[tuple(np.indices(seg.shape)) + (seg,)] - conf = resize_image(conf, img_h_page, img_w_page) - return prediction, conf - - if img.shape[0] < img_height_model: - img = resize_image(img, img_height_model, img.shape[1]) - if img.shape[1] < img_width_model: - img = resize_image(img, img.shape[0], img_width_model) - - self.logger.debug("Patch size: %sx%s", img_height_model, img_width_model) - margin = int(marginal_of_patch_percent * img_height_model) - window = 1 / (1 + np.exp(5.0 - 5 * np.arange(2 * margin) / margin)) - width_mid = img_width_model - 2 * margin - height_mid = img_height_model - 2 * margin - img_h = img.shape[0] - img_w = img.shape[1] - prediction = None - nxf = math.ceil((img_w - 2.0 * margin) / width_mid) - nyf = math.ceil((img_h - 2.0 * margin) / height_mid) - - batch_i = [] - batch_j = [] - batch_x_u = [] - batch_x_d = [] - batch_x_s = [] - batch_y_u = [] - batch_y_d = [] - batch_y_s = [] - batch = 0 - img_patch = np.zeros((n_batch_inference, - img_height_model, - img_width_model, - 3), dtype=np.float16) - for i in range(nxf): - for j in range(nyf): - index_x_d = i * width_mid - index_x_u = index_x_d + img_width_model - if index_x_u > img_w: - index_x_s = index_x_u - img_w - index_x_u = img_w - index_x_d = img_w - img_width_model - else: - index_x_s = 0 - index_y_d = j * height_mid - index_y_u = index_y_d + img_height_model - if index_y_u > img_h: - index_y_s = index_y_u - img_h - index_y_u = img_h - index_y_d = img_h - img_height_model - else: - index_y_s = 0 - - batch_i.append(i) - batch_j.append(j) - batch_x_u.append(index_x_u) - batch_x_d.append(index_x_d) - batch_x_s.append(index_x_s) - batch_y_d.append(index_y_d) - batch_y_u.append(index_y_u) - batch_y_s.append(index_y_s) - - img_patch[batch] = img[index_y_d: index_y_u, - index_x_d: index_x_u] - batch += 1 - if (batch == n_batch_inference or - # last batch - i == nxf - 1 and j == nyf - 1): - self.logger.debug("predicting patches on %s", str(img_patch.shape)) - label_p_pred = model.predict(img_patch, verbose=0) - if prediction is None: - # now we know the number of classes - prediction = np.zeros((img_h, img_w, label_p_pred.shape[-1]), dtype=float) - - for batch in range(batch): - where = np.index_exp[batch_y_d[batch]: batch_y_u[batch], - batch_x_d[batch]: batch_x_u[batch]] - # shorter window on last tile - part = np.index_exp[batch_y_s[batch]:, - batch_x_s[batch]:] - # normalize probability (where windows overlap) - attenuation_y = np.ones(img_height_model - batch_y_s[batch]) - attenuation_x = np.ones(img_width_model - batch_x_s[batch]) - if margin and batch_j[batch] > 0: - attenuation_y[:2 * margin] = window - if margin and batch_j[batch] < nyf - 1: - attenuation_y[-2 * margin:] = 1 - window - if margin and batch_i[batch] > 0: - attenuation_x[:2 * margin] = window - if margin and batch_i[batch] < nxf - 1: - attenuation_x[-2 * margin:] = 1 - window - label_p_pred[batch][part] *= attenuation_y[:, np.newaxis, np.newaxis] - label_p_pred[batch][part] *= attenuation_x[np.newaxis, :, np.newaxis] - prediction[where][part] += label_p_pred[batch][part] - - batch_i = [] - batch_j = [] - batch_x_u = [] - batch_x_d = [] - batch_x_s = [] - batch_y_u = [] - batch_y_d = [] - batch_y_s = [] - batch = 0 - img_patch[:] = 0 - - # decode - seg = np.argmax(prediction, axis=2).astype(np.uint8) - conf = prediction[tuple(np.indices(seg.shape)) + (seg,)] - if thresholding_for_artificial_class: - seg_art = prediction[:, :, artificial_class] >= threshold_art_class - seg_mask_label(seg, seg_art, - label=artificial_class, - only=True, - skeletonize=True, - dilate=3, - keep=separator_class) - gc.collect() - return seg, conf - - # variant of do_prediction_new_concept with no need - # for resizing or tiling into patches - done on model - # (Tensorflow/CUDA) side - # (after loading wrapped resized or patched model) - def do_prediction_new_concept_autosize( - self, img, model, - n_batch_inference=None, - thresholding_for_heading=False, - thresholding_for_artificial_class=False, - threshold_art_class=0.1, - artificial_class=4, - ): - self.logger.debug("enter do_prediction_new_concept (%s)", model.name) - img = img / 255.0 - img = img.astype(np.float16) - - prediction = model.predict(img[np.newaxis])[0] - confidence = prediction[:, :, 1] - segmentation = np.argmax(prediction, axis=2).astype(np.uint8) - - if thresholding_for_artificial_class: - seg_mask_label(segmentation, - prediction[:, :, artificial_class] >= threshold_art_class, - label=artificial_class, - only=True, - skeletonize=True, - dilate=3) - if thresholding_for_heading: - seg_mask_label(segmentation, - prediction[:, :, 2] >= 0.2, - label=2) - gc.collect() - return segmentation, confidence - def extract_page(self, image): page_cropped = img = image['img_res'] h, w = img.shape[:2] @@ -816,7 +450,9 @@ class Eynollah: if not self.ignore_page_extraction: self.logger.debug("enter extract_page") #cv2.GaussianBlur(img, (5, 5), 0) - prediction = self.do_prediction(False, img, self.model_zoo.get("page")) + prediction = do_prediction(img, self.model_zoo.get("page"), + patches=False, + logger=self.logger) contours, _ = cv2.findContours(prediction, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) if len(contours): areas = np.array(list(map(cv2.contourArea, contours))) @@ -832,7 +468,9 @@ class Eynollah: if not self.ignore_page_extraction: self.logger.debug("enter early_page_for_num_of_column_classification") img2 = cv2.GaussianBlur(img, (5, 5), 0) - prediction = self.do_prediction(False, img2, self.model_zoo.get("page")) + prediction = do_prediction(img2, self.model_zoo.get("page"), + patches=False, + logger=self.logger) prediction = cv2.dilate(prediction, KERNEL, iterations=3) contours, _ = cv2.findContours(prediction, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) if len(contours): @@ -852,8 +490,10 @@ class Eynollah: img_height_h = img.shape[0] img_width_h = img.shape[1] - prediction_regions, confidence_regions = self.do_prediction_new_concept( - patches, img, self.model_zoo.get("region_fl" if patches else "region_fl_np"), + prediction_regions, confidence_regions = do_prediction_new_concept( + img, self.model_zoo.get("region_fl" if patches else "region_fl_np"), + patches=patches, + logger=self.logger, n_batch_inference=1, thresholding_for_heading=not patches) @@ -866,8 +506,10 @@ class Eynollah: img_width_h = img.shape[1] model_region = self.model_zoo.get("region_fl" if patches else "region_fl_np") - prediction_regions = self.do_prediction(patches, img, model_region, - marginal_of_patch_percent=0.1) + prediction_regions = do_prediction(img, model_region, + patches=patches, + logger=self.logger, + marginal_of_patch_percent=0.1) prediction_regions = resize_image(prediction_regions, img_height_h, img_width_h) self.logger.debug("exit extract_text_regions") return prediction_regions @@ -1016,14 +658,16 @@ class Eynollah: n_batch = 1 else: n_batch = 3 - prediction_textline, conf_textline = self.do_prediction_new_concept( - use_patches, img, self.model_zoo.get("textline"), + prediction_textline, conf_textline = do_prediction_new_concept( + img, self.model_zoo.get("textline"), + patches=use_patches, + logger=self.logger, artificial_class=2, n_batch_inference=n_batch, thresholding_for_artificial_class=True, threshold_art_class=self.threshold_art_class_textline) - #prediction_textline_longshot = self.do_prediction(False, img, self.model_zoo.get("textline")) + #prediction_textline_longshot = do_prediction(img, self.model_zoo.get("textline"), patches=False) self.logger.debug('exit textline_contours') # suppress artificial boundary label @@ -1086,13 +730,14 @@ class Eynollah: new_w, new_h, num_col_classifier) patches = True - prediction_regions, confidence_regions = \ - self.do_prediction_new_concept( - patches, img_resized, self.model_zoo.get("region_1_2"), - n_batch_inference=1, - thresholding_for_artificial_class=True, - threshold_art_class=self.threshold_art_class_layout, - separator_class=label_seps) + prediction_regions, confidence_regions = do_prediction_new_concept( + img_resized, self.model_zoo.get("region_1_2"), + patches=patches, + logger=self.logger, + n_batch_inference=1, + thresholding_for_artificial_class=True, + threshold_art_class=self.threshold_art_class_layout, + separator_class=label_seps) prediction_regions = resize_image(prediction_regions, img_height_h, img_width_h) confidence_regions = resize_image(confidence_regions, img_height_h, img_width_h) @@ -1463,9 +1108,10 @@ class Eynollah: return image_revised_last def get_tables_from_model(self, img): - table_prediction, table_confidence = self.do_prediction_new_concept( - False, img, - self.model_zoo.get("table"), + table_prediction, table_confidence = do_prediction_new_concept( + img, self.model_zoo.get("table"), + patches=False, + logger=self.logger, thresholding_for_artificial_class=True, threshold_art_class=0.05, artificial_class=2) diff --git a/src/eynollah/eynollah_ocr.py b/src/eynollah/eynollah_ocr.py index b5bf6d8..b5b1cf8 100644 --- a/src/eynollah/eynollah_ocr.py +++ b/src/eynollah/eynollah_ocr.py @@ -32,6 +32,7 @@ from .utils import ( from .utils.font import get_font from .utils.xml import etree_namespace_for_element_tag from .utils.resize import resize_image +from .utils.tiling import do_prediction from .utils.utils_ocr import ( break_curved_line_into_small_pieces_and_then_merge, fit_text_single_line, @@ -200,8 +201,10 @@ class Eynollah_ocr(Eynollah): if img_bin is None: # run ad-hoc binarization self.logger.info("running binarization for ensemble input") - img_bin = self.do_prediction(True, img, self.model_zoo.get("binarization"), - n_batch_inference=5) + img_bin = do_prediction(img, self.model_zoo.get("binarization"), + patches=True, + logger=self.logger, + n_batch_inference=5) img_bin = np.repeat(img_bin[:, :, np.newaxis], 3, axis=2) img_bin = 255 * (img_bin == 0).astype(np.uint8) diff --git a/src/eynollah/sbb_binarize.py b/src/eynollah/sbb_binarize.py index 9b154a8..7541ca4 100644 --- a/src/eynollah/sbb_binarize.py +++ b/src/eynollah/sbb_binarize.py @@ -18,6 +18,7 @@ import cv2 from .eynollah import Eynollah from .model_zoo import EynollahModelZoo from .utils.resize import resize_image +from .utils.tiling import do_prediction from .utils import is_image_filename class SbbBinarizer(Eynollah): @@ -84,8 +85,10 @@ class SbbBinarizer(Eynollah): ): image = self.cache_images(image_filename=img_filename, image_pil=img_pil) img = self.imread(image) - img_bin = self.do_prediction(use_patches, img, self.model_zoo.get("binarization"), - n_batch_inference=5) + img_bin = do_prediction(img, self.model_zoo.get("binarization"), + patches=use_patches, + logger=self.logger, + n_batch_inference=5) img_bin = 255 * (img_bin == 0).astype(np.uint8) #img_bin = np.repeat(img_bin[:, :, np.newaxis], 3, axis=2).astype(np.uint8) return img_bin diff --git a/src/eynollah/utils/tiling.py b/src/eynollah/utils/tiling.py new file mode 100644 index 0000000..a0a8822 --- /dev/null +++ b/src/eynollah/utils/tiling.py @@ -0,0 +1,394 @@ +from logging import getLogger +import math +import gc + +import numpy as np + +from . import seg_mask_label +from .resize import resize_image +from ..predictor import Predictor + +def do_prediction( + img: np.ndarray, + model: Predictor, + logger=None, + patches=False, + n_batch_inference=1, + marginal_of_patch_percent=0.1, + thresholding_for_some_classes=False, + thresholding_for_heading=False, + heading_class=2, + thresholding_for_artificial_class=False, + threshold_art_class=0.1, + artificial_class=2, + is_enhancement=False, +) -> np.ndarray: + if logger is None: + logger = getLogger('eynollah') + + logger.debug("enter do_prediction (patches=%d)", patches) + _, img_height_model, img_width_model, _ = model.input_shape + img_h_page = img.shape[0] + img_w_page = img.shape[1] + + img = img / 255. + img = img.astype(np.float16) + + if not patches: + img = resize_image(img, img_height_model, img_width_model) + + label_p_pred = model.predict(img[np.newaxis], verbose=0)[0] + if is_enhancement: + seg = np.round(label_p_pred * 255).astype(np.uint8) + else: + seg = np.argmax(label_p_pred, axis=2).astype(np.uint8) + + if thresholding_for_artificial_class: + seg_mask_label( + seg, label_p_pred[:, :, artificial_class] >= threshold_art_class, + label=artificial_class, + skeletonize=True) + + if thresholding_for_heading: + seg_mask_label( + seg, label_p_pred[:, :, heading_class] >= 0.2, + label=heading_class) + + return resize_image(seg, img_h_page, img_w_page) + + if img_h_page < img_height_model: + img = resize_image(img, img_height_model, img.shape[1]) + if img_w_page < img_width_model: + img = resize_image(img, img.shape[0], img_width_model) + + logger.debug("Patch size: %sx%s", img_height_model, img_width_model) + margin = int(marginal_of_patch_percent * img_height_model) + window = 1 / (1 + np.exp(5.0 - 5 * np.arange(2 * margin) / margin)) + width_mid = img_width_model - 2 * margin + height_mid = img_height_model - 2 * margin + img_h = img.shape[0] + img_w = img.shape[1] + prediction: np.ndarray = None # type: ignore + nxf = math.ceil((img_w - 2.0 * margin) / width_mid) + nyf = math.ceil((img_h - 2.0 * margin) / height_mid) + + batch_i = [] + batch_j = [] + batch_x_u = [] + batch_x_d = [] + batch_x_s = [] + batch_y_u = [] + batch_y_d = [] + batch_y_s = [] + + batch = 0 + img_patch = np.zeros((n_batch_inference, + img_height_model, + img_width_model, + 3), dtype=np.float16) + for i in range(nxf): + for j in range(nyf): + index_x_d = i * width_mid + index_x_u = index_x_d + img_width_model + if index_x_u > img_w: + index_x_s = index_x_u - img_w + index_x_u = img_w + index_x_d = img_w - img_width_model + else: + index_x_s = 0 + index_y_d = j * height_mid + index_y_u = index_y_d + img_height_model + if index_y_u > img_h: + index_y_s = index_y_u - img_h + index_y_u = img_h + index_y_d = img_h - img_height_model + else: + index_y_s = 0 + + batch_i.append(i) + batch_j.append(j) + batch_x_u.append(index_x_u) + batch_x_d.append(index_x_d) + batch_x_s.append(index_x_s) + batch_y_d.append(index_y_d) + batch_y_u.append(index_y_u) + batch_y_s.append(index_y_s) + + img_patch[batch] = img[index_y_d: index_y_u, + index_x_d: index_x_u] + batch += 1 + if (batch == n_batch_inference or + # last batch + i == nxf - 1 and j == nyf - 1): + logger.debug("predicting patches on %s", str(img_patch.shape)) + label_p_pred = model.predict(img_patch, verbose=0) + if prediction is None: + # now we know the number of classes + prediction = np.zeros((img_h, img_w, label_p_pred.shape[-1]), dtype=float) + + for batch in range(batch): + where = np.index_exp[batch_y_d[batch]: batch_y_u[batch], + batch_x_d[batch]: batch_x_u[batch]] + # shorter window on last tile + part = np.index_exp[batch_y_s[batch]:, + batch_x_s[batch]:] + # normalize probability (where windows overlap) + attenuation_y = np.ones(img_height_model - batch_y_s[batch]) + attenuation_x = np.ones(img_width_model - batch_x_s[batch]) + if margin and batch_j[batch] > 0: + attenuation_y[:2 * margin] = window + if margin and batch_j[batch] < nyf - 1: + attenuation_y[-2 * margin:] = 1 - window + if margin and batch_i[batch] > 0: + attenuation_x[:2 * margin] = window + if margin and batch_i[batch] < nxf - 1: + attenuation_x[-2 * margin:] = 1 - window + label_p_pred[batch][part] *= attenuation_y[:, np.newaxis, np.newaxis] + label_p_pred[batch][part] *= attenuation_x[np.newaxis, :, np.newaxis] + prediction[where][part] += label_p_pred[batch][part] + + batch_i = [] + batch_j = [] + batch_x_u = [] + batch_x_d = [] + batch_x_s = [] + batch_y_u = [] + batch_y_d = [] + batch_y_s = [] + batch = 0 + img_patch[:] = 0 + + if is_enhancement: + seg = (prediction * 255).astype(np.uint8) + else: + seg = np.argmax(prediction, axis=2).astype(np.uint8) + if thresholding_for_some_classes: + seg_mask_label( + seg, prediction[:, :, 4] > 0.03, + label=4) # + seg_mask_label( + seg, prediction[:, :, 0] > 0.25, + label=0) # bg + seg_mask_label( + seg, prediction[:, :, 3] > 0.10 & seg == 0, + label=3) # line + if thresholding_for_artificial_class: + seg_art = prediction[:, :, artificial_class] >= threshold_art_class + seg_mask_label(seg, seg_art, + label=artificial_class, + only=True, + skeletonize=True, + dilate=3) + + if img_h != img_h_page or img_w != img_w_page: + seg = resize_image(seg, img_h_page, img_w_page) + + gc.collect() + return seg + +def do_prediction_new_concept( + img: np.ndarray, + model: Predictor, + logger=None, + patches=False, + n_batch_inference=1, + marginal_of_patch_percent=0.1, + thresholding_for_heading=False, + heading_class=2, + thresholding_for_artificial_class=False, + threshold_art_class=0.1, + artificial_class=4, + separator_class=0, +) -> np.ndarray: + if logger is None: + logger = getLogger('eynollah') + + logger.debug("enter do_prediction_new_concept (patches=%d)", patches) + _, img_height_model, img_width_model, _ = model.input_shape + + img = img / 255.0 + img = img.astype(np.float16) + + if not patches: + img_h_page = img.shape[0] + img_w_page = img.shape[1] + img = resize_image(img, img_height_model, img_width_model) + + label_p_pred = model.predict(img[np.newaxis], verbose=0)[0] + seg = np.argmax(label_p_pred, axis=2).astype(np.uint8) + + prediction = resize_image(seg, img_h_page, img_w_page) + + if thresholding_for_artificial_class: + mask = resize_image(label_p_pred[:, :, artificial_class], + img_h_page, img_w_page) >= threshold_art_class + seg_mask_label(prediction, mask, + label=artificial_class, + only=True, + skeletonize=True, + dilate=3, + keep=separator_class) + if thresholding_for_heading: + mask = resize_image(label_p_pred[:, :, heading_class], + img_h_page, img_w_page) >= 0.2 + seg_mask_label(prediction, mask, + label=heading_class) + + conf = label_p_pred[tuple(np.indices(seg.shape)) + (seg,)] + conf = resize_image(conf, img_h_page, img_w_page) + return prediction, conf + + if img.shape[0] < img_height_model: + img = resize_image(img, img_height_model, img.shape[1]) + if img.shape[1] < img_width_model: + img = resize_image(img, img.shape[0], img_width_model) + + logger.debug("Patch size: %sx%s", img_height_model, img_width_model) + margin = int(marginal_of_patch_percent * img_height_model) + window = 1 / (1 + np.exp(5.0 - 5 * np.arange(2 * margin) / margin)) + width_mid = img_width_model - 2 * margin + height_mid = img_height_model - 2 * margin + img_h = img.shape[0] + img_w = img.shape[1] + prediction = None + nxf = math.ceil((img_w - 2.0 * margin) / width_mid) + nyf = math.ceil((img_h - 2.0 * margin) / height_mid) + + batch_i = [] + batch_j = [] + batch_x_u = [] + batch_x_d = [] + batch_x_s = [] + batch_y_u = [] + batch_y_d = [] + batch_y_s = [] + batch = 0 + img_patch = np.zeros((n_batch_inference, + img_height_model, + img_width_model, + 3), dtype=np.float16) + for i in range(nxf): + for j in range(nyf): + index_x_d = i * width_mid + index_x_u = index_x_d + img_width_model + if index_x_u > img_w: + index_x_s = index_x_u - img_w + index_x_u = img_w + index_x_d = img_w - img_width_model + else: + index_x_s = 0 + index_y_d = j * height_mid + index_y_u = index_y_d + img_height_model + if index_y_u > img_h: + index_y_s = index_y_u - img_h + index_y_u = img_h + index_y_d = img_h - img_height_model + else: + index_y_s = 0 + + batch_i.append(i) + batch_j.append(j) + batch_x_u.append(index_x_u) + batch_x_d.append(index_x_d) + batch_x_s.append(index_x_s) + batch_y_d.append(index_y_d) + batch_y_u.append(index_y_u) + batch_y_s.append(index_y_s) + + img_patch[batch] = img[index_y_d: index_y_u, + index_x_d: index_x_u] + batch += 1 + if (batch == n_batch_inference or + # last batch + i == nxf - 1 and j == nyf - 1): + logger.debug("predicting patches on %s", str(img_patch.shape)) + label_p_pred = model.predict(img_patch, verbose=0) + if prediction is None: + # now we know the number of classes + prediction = np.zeros((img_h, img_w, label_p_pred.shape[-1]), dtype=float) + + for batch in range(batch): + where = np.index_exp[batch_y_d[batch]: batch_y_u[batch], + batch_x_d[batch]: batch_x_u[batch]] + # shorter window on last tile + part = np.index_exp[batch_y_s[batch]:, + batch_x_s[batch]:] + # normalize probability (where windows overlap) + attenuation_y = np.ones(img_height_model - batch_y_s[batch]) + attenuation_x = np.ones(img_width_model - batch_x_s[batch]) + if margin and batch_j[batch] > 0: + attenuation_y[:2 * margin] = window + if margin and batch_j[batch] < nyf - 1: + attenuation_y[-2 * margin:] = 1 - window + if margin and batch_i[batch] > 0: + attenuation_x[:2 * margin] = window + if margin and batch_i[batch] < nxf - 1: + attenuation_x[-2 * margin:] = 1 - window + label_p_pred[batch][part] *= attenuation_y[:, np.newaxis, np.newaxis] + label_p_pred[batch][part] *= attenuation_x[np.newaxis, :, np.newaxis] + prediction[where][part] += label_p_pred[batch][part] + + batch_i = [] + batch_j = [] + batch_x_u = [] + batch_x_d = [] + batch_x_s = [] + batch_y_u = [] + batch_y_d = [] + batch_y_s = [] + batch = 0 + img_patch[:] = 0 + + # decode + seg = np.argmax(prediction, axis=2).astype(np.uint8) + conf = prediction[tuple(np.indices(seg.shape)) + (seg,)] + if thresholding_for_artificial_class: + seg_art = prediction[:, :, artificial_class] >= threshold_art_class + seg_mask_label(seg, seg_art, + label=artificial_class, + only=True, + skeletonize=True, + dilate=3, + keep=separator_class) + gc.collect() + return seg, conf + +# variant of do_prediction_new_concept with no need +# for resizing or tiling into patches - done on model +# (Tensorflow/CUDA) side +# (after loading wrapped resized or patched model) +def do_prediction_new_concept_autosize( + img: np.ndarray, + model: Predictor, + logger=None, + n_batch_inference=None, + thresholding_for_heading=False, + thresholding_for_artificial_class=False, + threshold_art_class=0.1, + artificial_class=4, +) -> np.ndarray: + if logger is None: + logger = getLogger('eynollah') + + logger.debug("enter do_prediction_new_concept (%s)", model.name) + img = img / 255.0 + img = img.astype(np.float16) + + prediction = model.predict(img[np.newaxis])[0] + confidence = prediction[:, :, 1] + segmentation = np.argmax(prediction, axis=2).astype(np.uint8) + + if thresholding_for_artificial_class: + seg_mask_label(segmentation, + prediction[:, :, artificial_class] >= threshold_art_class, + label=artificial_class, + only=True, + skeletonize=True, + dilate=3) + if thresholding_for_heading: + seg_mask_label(segmentation, + prediction[:, :, 2] >= 0.2, + label=2) + gc.collect() + return segmentation, confidence + diff --git a/tests/test_tiling.py b/tests/test_tiling.py new file mode 100644 index 0000000..5412b51 --- /dev/null +++ b/tests/test_tiling.py @@ -0,0 +1,108 @@ +import pytest +import cv2 +import numpy as np +from matplotlib import pyplot as plt + +from eynollah.utils.tiling import do_prediction, do_prediction_new_concept + +@pytest.mark.parametrize( + "height,width", + [ + (448, 448), + (672, 672), + (1088, 832), + (1152, 896), + ]) +def test_tiling_idem(image_resources, height, width): + infile = image_resources[0] + class PseudoModel: + def predict(self, images, **kwargs): + return images + @property + def input_shape(self): + return None, height, width, None + model = PseudoModel() + in_img = cv2.imread(infile) + outimg = do_prediction(in_img, model, patches=True, is_enhancement=True, + marginal_of_patch_percent=0) + assert in_img.shape == outimg.shape + assert np.all(in_img == outimg) + outimg = do_prediction(in_img, model, patches=True, is_enhancement=True, + marginal_of_patch_percent=0.1) + assert in_img.shape == outimg.shape + assert np.all(in_img == outimg) + outimg = do_prediction(in_img, model, patches=True, is_enhancement=True, + marginal_of_patch_percent=0.2) + assert in_img.shape == outimg.shape + assert in_img.dtype == outimg.dtype + assert np.all(in_img == outimg) + +@pytest.mark.parametrize( + "height,width", + [ + (448, 448), + (672, 672), + (1088, 832), + (1152, 896), + ]) +def test_tiling_min(image_resources, height, width): + infile = image_resources[0] + class PseudoModel: + def predict(self, images, **kwargs): + M = images.min(axis=(1, 2, 3)) + return 1. * (images == M) + @property + def input_shape(self): + return None, height, width, None + model = PseudoModel() + in_img = cv2.imread(infile) + outimg = do_prediction(in_img, model, patches=True, + marginal_of_patch_percent=0) + assert in_img.shape[:2] == outimg.shape + assert np.any(outimg) + outimg = do_prediction(in_img, model, patches=True, + marginal_of_patch_percent=0.1) + assert in_img.shape[:2] == outimg.shape + assert np.any(outimg) + outimg = do_prediction(in_img, model, patches=True, + marginal_of_patch_percent=0.2) + assert in_img.shape[:2] == outimg.shape + assert np.any(outimg) + +@pytest.mark.parametrize( + "height,width", + [ + (448, 448), + (672, 672), + (1088, 832), + (1152, 896), + ]) +def test_tiling_min_conf(image_resources, height, width): + infile = image_resources[0] + class PseudoModel: + def predict(self, images, **kwargs): + M = images.min(axis=(1, 2, 3)) + return 1. * (images == M) + @property + def input_shape(self): + return None, height, width, None + model = PseudoModel() + in_img = cv2.imread(infile) + outimg, conf = do_prediction_new_concept( + in_img, model, patches=True, + marginal_of_patch_percent=0) + assert in_img.shape[:2] == outimg.shape + assert np.any(outimg) + assert np.sum(conf) < conf.size + outimg, conf = do_prediction_new_concept( + in_img, model, patches=True, + marginal_of_patch_percent=0.1) + assert in_img.shape[:2] == outimg.shape + assert np.any(outimg) + assert np.sum(conf) < conf.size + outimg, conf = do_prediction_new_concept( + in_img, model, patches=True, + marginal_of_patch_percent=0.2) + assert in_img.shape[:2] == outimg.shape + assert np.any(outimg) + assert np.sum(conf) < conf.size