do_prediction*: refactor into separate module

This commit is contained in:
Robert Sachunsky 2026-07-30 15:28:01 +02:00
parent bad250fd72
commit b2777a2562
6 changed files with 560 additions and 410 deletions

View file

@ -1,5 +1,5 @@
""" """
extract images? extract image regions only
""" """
from concurrent.futures import ProcessPoolExecutor from concurrent.futures import ProcessPoolExecutor
@ -19,6 +19,7 @@ from .model_zoo.model_zoo import EynollahModelZoo
from .writer import EynollahXmlWriter from .writer import EynollahXmlWriter
from .eynollah import Eynollah from .eynollah import Eynollah
from .utils import box2rect, is_image_filename from .utils import box2rect, is_image_filename
from .utils.tiling import do_prediction_new_concept
from .plot import EynollahPlotter from .plot import EynollahPlotter
from .utils import Region from .utils import Region
@ -110,8 +111,9 @@ class EynollahImageExtractor(Eynollah):
img_h_new = img_w_new * img_height_h // img_width_h img_h_new = img_w_new * img_height_h // img_width_h
img_resized = resize_image(img, img_h_new, img_w_new) img_resized = resize_image(img, img_h_new, img_w_new)
prediction_regions, _ = self.do_prediction_new_concept( prediction_regions, _ = do_prediction_new_concept(
True, img_resized, self.model_zoo.get("extract_images")) 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) prediction_regions = resize_image(prediction_regions, img_height_h, img_width_h)
mask_texts_only = (prediction_regions == label_text).astype(np.uint8) mask_texts_only = (prediction_regions == label_text).astype(np.uint8)
@ -158,17 +160,11 @@ class EynollahImageExtractor(Eynollah):
**kwargs **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") self.logger.debug("enter run")
# Log enabled features directly # Log enabled features directly
enabled_modes = [] 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: if self.enable_plotting:
self.logger.info("Saving debug plots") self.logger.info("Saving debug plots")
if dir_of_cropped_images: if dir_of_cropped_images:

View file

@ -21,7 +21,6 @@ import logging.handlers
import sys import sys
from difflib import SequenceMatcher as sq from difflib import SequenceMatcher as sq
import math
import os import os
import time import time
from typing import Optional, List, Tuple from typing import Optional, List, Tuple
@ -30,7 +29,6 @@ from functools import partial
from pathlib import Path from pathlib import Path
import multiprocessing as mp import multiprocessing as mp
from concurrent.futures import ProcessPoolExecutor, as_completed from concurrent.futures import ProcessPoolExecutor, as_completed
import gc
import cv2 import cv2
import numpy as np import numpy as np
@ -65,6 +63,7 @@ from .utils.separate_lines import (
from .utils.marginals import get_marginals from .utils.marginals import get_marginals
from .utils.resize import resize_image from .utils.resize import resize_image
from .utils.shm import share_ndarray from .utils.shm import share_ndarray
from .utils.tiling import do_prediction, do_prediction_new_concept
from .utils import ( from .utils import (
Region, Region,
TextRegion, TextRegion,
@ -77,7 +76,6 @@ from .utils import (
box2slice, box2slice,
find_num_col, find_num_col,
otsu_copy_binary, otsu_copy_binary,
seg_mask_label,
fill_bb_of_drop_capitals, fill_bb_of_drop_capitals,
split_textregion_main_vs_head, split_textregion_main_vs_head,
small_textlines_to_parent_adherence2, small_textlines_to_parent_adherence2,
@ -354,7 +352,9 @@ class Eynollah:
img_new, _ = fun(img, num_col, conf_col, width_early) img_new, _ = fun(img, num_col, conf_col, width_early)
if img_new.shape[1] > img.shape[1]: if img_new.shape[1] > img.shape[1]:
img_new = self.do_prediction(True, img_new, self.model_zoo.get("enhancement"), img_new = do_prediction(img_new, self.model_zoo.get("enhancement"),
patches=True,
logger=self.logger,
marginal_of_patch_percent=0, marginal_of_patch_percent=0,
n_batch_inference=3, n_batch_inference=3,
is_enhancement=True) is_enhancement=True)
@ -372,7 +372,10 @@ class Eynollah:
img = self.imread(image) img = self.imread(image)
self.logger.info("Detected %s DPI", dpi) self.logger.info("Detected %s DPI", dpi)
if self.input_binary: 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 = 255 * (prediction_bin == 0)
prediction_bin = np.repeat(prediction_bin[:, :, np.newaxis], 3, axis=2).astype(np.uint8) prediction_bin = np.repeat(prediction_bin[:, :, np.newaxis], 3, axis=2).astype(np.uint8)
image['img_bin_uint8'] = prediction_bin image['img_bin_uint8'] = prediction_bin
@ -436,375 +439,6 @@ class Eynollah:
image['scale_x'] = 1.0 * img_res.shape[1] / img.shape[1] image['scale_x'] = 1.0 * img_res.shape[1] / img.shape[1]
return is_image_enhanced, num_col, is_image_resized 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): def extract_page(self, image):
page_cropped = img = image['img_res'] page_cropped = img = image['img_res']
h, w = img.shape[:2] h, w = img.shape[:2]
@ -816,7 +450,9 @@ class Eynollah:
if not self.ignore_page_extraction: if not self.ignore_page_extraction:
self.logger.debug("enter extract_page") self.logger.debug("enter extract_page")
#cv2.GaussianBlur(img, (5, 5), 0) #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) contours, _ = cv2.findContours(prediction, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
if len(contours): if len(contours):
areas = np.array(list(map(cv2.contourArea, contours))) areas = np.array(list(map(cv2.contourArea, contours)))
@ -832,7 +468,9 @@ class Eynollah:
if not self.ignore_page_extraction: if not self.ignore_page_extraction:
self.logger.debug("enter early_page_for_num_of_column_classification") self.logger.debug("enter early_page_for_num_of_column_classification")
img2 = cv2.GaussianBlur(img, (5, 5), 0) 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) prediction = cv2.dilate(prediction, KERNEL, iterations=3)
contours, _ = cv2.findContours(prediction, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) contours, _ = cv2.findContours(prediction, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
if len(contours): if len(contours):
@ -852,8 +490,10 @@ class Eynollah:
img_height_h = img.shape[0] img_height_h = img.shape[0]
img_width_h = img.shape[1] img_width_h = img.shape[1]
prediction_regions, confidence_regions = self.do_prediction_new_concept( prediction_regions, confidence_regions = do_prediction_new_concept(
patches, img, self.model_zoo.get("region_fl" if patches else "region_fl_np"), img, self.model_zoo.get("region_fl" if patches else "region_fl_np"),
patches=patches,
logger=self.logger,
n_batch_inference=1, n_batch_inference=1,
thresholding_for_heading=not patches) thresholding_for_heading=not patches)
@ -866,7 +506,9 @@ class Eynollah:
img_width_h = img.shape[1] img_width_h = img.shape[1]
model_region = self.model_zoo.get("region_fl" if patches else "region_fl_np") model_region = self.model_zoo.get("region_fl" if patches else "region_fl_np")
prediction_regions = self.do_prediction(patches, img, model_region, prediction_regions = do_prediction(img, model_region,
patches=patches,
logger=self.logger,
marginal_of_patch_percent=0.1) marginal_of_patch_percent=0.1)
prediction_regions = resize_image(prediction_regions, img_height_h, img_width_h) prediction_regions = resize_image(prediction_regions, img_height_h, img_width_h)
self.logger.debug("exit extract_text_regions") self.logger.debug("exit extract_text_regions")
@ -1016,14 +658,16 @@ class Eynollah:
n_batch = 1 n_batch = 1
else: else:
n_batch = 3 n_batch = 3
prediction_textline, conf_textline = self.do_prediction_new_concept( prediction_textline, conf_textline = do_prediction_new_concept(
use_patches, img, self.model_zoo.get("textline"), img, self.model_zoo.get("textline"),
patches=use_patches,
logger=self.logger,
artificial_class=2, artificial_class=2,
n_batch_inference=n_batch, n_batch_inference=n_batch,
thresholding_for_artificial_class=True, thresholding_for_artificial_class=True,
threshold_art_class=self.threshold_art_class_textline) 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') self.logger.debug('exit textline_contours')
# suppress artificial boundary label # suppress artificial boundary label
@ -1086,9 +730,10 @@ class Eynollah:
new_w, new_h, num_col_classifier) new_w, new_h, num_col_classifier)
patches = True patches = True
prediction_regions, confidence_regions = \ prediction_regions, confidence_regions = do_prediction_new_concept(
self.do_prediction_new_concept( img_resized, self.model_zoo.get("region_1_2"),
patches, img_resized, self.model_zoo.get("region_1_2"), patches=patches,
logger=self.logger,
n_batch_inference=1, n_batch_inference=1,
thresholding_for_artificial_class=True, thresholding_for_artificial_class=True,
threshold_art_class=self.threshold_art_class_layout, threshold_art_class=self.threshold_art_class_layout,
@ -1463,9 +1108,10 @@ class Eynollah:
return image_revised_last return image_revised_last
def get_tables_from_model(self, img): def get_tables_from_model(self, img):
table_prediction, table_confidence = self.do_prediction_new_concept( table_prediction, table_confidence = do_prediction_new_concept(
False, img, img, self.model_zoo.get("table"),
self.model_zoo.get("table"), patches=False,
logger=self.logger,
thresholding_for_artificial_class=True, thresholding_for_artificial_class=True,
threshold_art_class=0.05, threshold_art_class=0.05,
artificial_class=2) artificial_class=2)

View file

@ -32,6 +32,7 @@ from .utils import (
from .utils.font import get_font from .utils.font import get_font
from .utils.xml import etree_namespace_for_element_tag from .utils.xml import etree_namespace_for_element_tag
from .utils.resize import resize_image from .utils.resize import resize_image
from .utils.tiling import do_prediction
from .utils.utils_ocr import ( from .utils.utils_ocr import (
break_curved_line_into_small_pieces_and_then_merge, break_curved_line_into_small_pieces_and_then_merge,
fit_text_single_line, fit_text_single_line,
@ -200,7 +201,9 @@ class Eynollah_ocr(Eynollah):
if img_bin is None: if img_bin is None:
# run ad-hoc binarization # run ad-hoc binarization
self.logger.info("running binarization for ensemble input") self.logger.info("running binarization for ensemble input")
img_bin = self.do_prediction(True, img, self.model_zoo.get("binarization"), img_bin = do_prediction(img, self.model_zoo.get("binarization"),
patches=True,
logger=self.logger,
n_batch_inference=5) n_batch_inference=5)
img_bin = np.repeat(img_bin[:, :, np.newaxis], 3, axis=2) img_bin = np.repeat(img_bin[:, :, np.newaxis], 3, axis=2)
img_bin = 255 * (img_bin == 0).astype(np.uint8) img_bin = 255 * (img_bin == 0).astype(np.uint8)

View file

@ -18,6 +18,7 @@ import cv2
from .eynollah import Eynollah from .eynollah import Eynollah
from .model_zoo import EynollahModelZoo from .model_zoo import EynollahModelZoo
from .utils.resize import resize_image from .utils.resize import resize_image
from .utils.tiling import do_prediction
from .utils import is_image_filename from .utils import is_image_filename
class SbbBinarizer(Eynollah): class SbbBinarizer(Eynollah):
@ -84,7 +85,9 @@ class SbbBinarizer(Eynollah):
): ):
image = self.cache_images(image_filename=img_filename, image_pil=img_pil) image = self.cache_images(image_filename=img_filename, image_pil=img_pil)
img = self.imread(image) img = self.imread(image)
img_bin = self.do_prediction(use_patches, img, self.model_zoo.get("binarization"), img_bin = do_prediction(img, self.model_zoo.get("binarization"),
patches=use_patches,
logger=self.logger,
n_batch_inference=5) n_batch_inference=5)
img_bin = 255 * (img_bin == 0).astype(np.uint8) img_bin = 255 * (img_bin == 0).astype(np.uint8)
#img_bin = np.repeat(img_bin[:, :, np.newaxis], 3, axis=2).astype(np.uint8) #img_bin = np.repeat(img_bin[:, :, np.newaxis], 3, axis=2).astype(np.uint8)

View file

@ -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

108
tests/test_tiling.py Normal file
View file

@ -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