diff --git a/.gitignore b/.gitignore index 49835a7..e3356ea 100644 --- a/.gitignore +++ b/.gitignore @@ -12,3 +12,5 @@ output.html *.sw? TAGS uv.lock +/ignore +*.log diff --git a/src/eynollah/Amiri-Regular.ttf b/src/eynollah/Amiri-Regular.ttf new file mode 100644 index 0000000..df5e1df Binary files /dev/null and b/src/eynollah/Amiri-Regular.ttf differ diff --git a/src/eynollah/Charis-Regular.ttf b/src/eynollah/Charis-Regular.ttf deleted file mode 100644 index a4e75a4..0000000 Binary files a/src/eynollah/Charis-Regular.ttf and /dev/null differ diff --git a/src/eynollah/eynollah_ocr.py b/src/eynollah/eynollah_ocr.py index 9cc2c9a..7fbfa5f 100644 --- a/src/eynollah/eynollah_ocr.py +++ b/src/eynollah/eynollah_ocr.py @@ -343,7 +343,7 @@ class Eynollah_ocr(Eynollah): if out_image_with_text: image_text = Image.new("RGB", (img.shape[1], img.shape[0]), "white") draw = ImageDraw.Draw(image_text) - font = get_font() + font = get_font(font_size=40) for indexer_text, bb_ind in enumerate(total_bb_coordinates): x_bb = bb_ind[0] diff --git a/src/eynollah/training/cli.py b/src/eynollah/training/cli.py index ccabb82..4b1b391 100644 --- a/src/eynollah/training/cli.py +++ b/src/eynollah/training/cli.py @@ -11,6 +11,7 @@ from .train import train_cli from .convert import convert_cli from .extract_line_gt import linegt_cli from .weights_ensembling import ensemble_cli +from .generate_or_update_cnn_rnn_ocr_character_list import main as update_ocr_characters_cli @click.group('training') def main(): @@ -23,3 +24,4 @@ main.add_command(train_cli, 'train') main.add_command(convert_cli, 'convert') main.add_command(linegt_cli, 'export_textline_images_and_text') main.add_command(ensemble_cli, 'ensembling') +main.add_command(update_ocr_characters_cli, 'generate_or_update_cnn_rnn_ocr_character_list') diff --git a/src/eynollah/training/extract_line_gt.py b/src/eynollah/training/extract_line_gt.py index 58fc253..fe9f60d 100644 --- a/src/eynollah/training/extract_line_gt.py +++ b/src/eynollah/training/extract_line_gt.py @@ -50,6 +50,12 @@ from ..utils import is_image_filename is_flag=True, help="if this parameter set to true, cropped textline images will not be masked with textline contour.", ) +@click.option( + "--exclude_vertical_lines", + "-exv", + is_flag=True, + help="if this parameter set to true, vertical textline images will be excluded.", +) def linegt_cli( image, dir_in, @@ -57,6 +63,7 @@ def linegt_cli( dir_out, pref_of_dataset, do_not_mask_with_textline_contour, + exclude_vertical_lines, ): assert bool(dir_in) ^ bool(image), "Set --dir-in or --image-filename, not both" if dir_in: @@ -100,6 +107,9 @@ def linegt_cli( x, y, w, h = cv2.boundingRect(textline_coords) + if exclude_vertical_lines and h > 2 * w: + continue + total_bb_coordinates.append([x, y, w, h]) img_poly_on_img = np.copy(img) diff --git a/src/eynollah/training/generate_gt_for_training.py b/src/eynollah/training/generate_gt_for_training.py index cc5a1b2..a848b65 100644 --- a/src/eynollah/training/generate_gt_for_training.py +++ b/src/eynollah/training/generate_gt_for_training.py @@ -6,6 +6,7 @@ from pathlib import Path from PIL import Image, ImageDraw, ImageFont import cv2 import numpy as np +from eynollah.utils.font import get_font from .gt_gen_utils import ( filter_contours_area_of_image, @@ -393,11 +394,15 @@ def visualize_reading_order(xml_file, dir_xml, dir_out, dir_imgs): layout = np.zeros( (y_len,x_len,3) ) layout = cv2.fillPoly(layout, pts =co_text_all, color=(1,1,1)) - img_file_name_with_format = find_format_of_given_filename_in_dir(dir_imgs, f_name) - img = cv2.imread(os.path.join(dir_imgs, img_file_name_with_format)) - - overlayed = overlay_layout_on_image(layout, img, cx_ordered, cy_ordered, color, thickness) - cv2.imwrite(os.path.join(dir_out, f_name+'.png'), overlayed) + try: + img_file_name_with_format = find_format_of_given_filename_in_dir(dir_imgs, f_name) + img = cv2.imread(os.path.join(dir_imgs, img_file_name_with_format)) + + overlayed = overlay_layout_on_image(layout, img, cx_ordered, cy_ordered, color, thickness) + cv2.imwrite(os.path.join(dir_out, f_name+'.png'), overlayed) + except: + pass + else: img = np.zeros( (y_len,x_len,3) ) @@ -452,14 +457,17 @@ def visualize_textline_segmentation(xml_file, dir_xml, dir_out, dir_imgs): xml_file = os.path.join(dir_xml,ind_xml ) f_name = Path(ind_xml).stem - img_file_name_with_format = find_format_of_given_filename_in_dir(dir_imgs, f_name) - img = cv2.imread(os.path.join(dir_imgs, img_file_name_with_format)) + try: + img_file_name_with_format = find_format_of_given_filename_in_dir(dir_imgs, f_name) + img = cv2.imread(os.path.join(dir_imgs, img_file_name_with_format)) + + co_tetxlines, y_len, x_len = get_textline_contours_for_visualization(xml_file) - co_tetxlines, y_len, x_len = get_textline_contours_for_visualization(xml_file) - - added_image = visualize_image_from_contours(co_tetxlines, img) - - cv2.imwrite(os.path.join(dir_out, f_name+'.png'), added_image) + added_image = visualize_image_from_contours(co_tetxlines, img) + + cv2.imwrite(os.path.join(dir_out, f_name+'.png'), added_image) + except: + pass @@ -509,15 +517,17 @@ def visualize_layout_segmentation(xml_file, dir_xml, dir_out, dir_imgs): f_name = Path(ind_xml).stem print(f_name, 'f_name') - img_file_name_with_format = find_format_of_given_filename_in_dir(dir_imgs, f_name) - img = cv2.imread(os.path.join(dir_imgs, img_file_name_with_format)) + try: + img_file_name_with_format = find_format_of_given_filename_in_dir(dir_imgs, f_name) + img = cv2.imread(os.path.join(dir_imgs, img_file_name_with_format)) + + co_text, co_graphic, co_sep, co_img, co_table, co_map, co_music, co_noise, y_len, x_len = get_layout_contours_for_visualization(xml_file) - co_text, co_graphic, co_sep, co_img, co_table, co_map, co_noise, y_len, x_len = get_layout_contours_for_visualization(xml_file) - - - added_image = visualize_image_from_contours_layout(co_text['paragraph'], co_text['header']+co_text['heading'], co_text['drop-capital'], co_sep, co_img, co_text['marginalia'], co_table, img) + added_image = visualize_image_from_contours_layout(co_text['paragraph'], co_text['header']+co_text['heading'], co_text['drop-capital'], co_sep, co_img, co_text['marginalia'], co_table, co_map, co_music, img) - cv2.imwrite(os.path.join(dir_out, f_name+'.png'), added_image) + cv2.imwrite(os.path.join(dir_out, f_name+'.png'), added_image) + except: + pass @@ -552,8 +562,8 @@ def visualize_ocr_text(xml_file, dir_xml, dir_out): else: xml_files_ind = [xml_file] - font_path = "Charis-7.000/Charis-Regular.ttf" # Make sure this file exists! - font = ImageFont.truetype(font_path, 40) + ###font_path = "Charis-7.000/Charis-Regular.ttf" # Make sure this file exists! + font = get_font(font_size=40)#ImageFont.truetype(font_path, 40) for ind_xml in tqdm(xml_files_ind): indexer = 0 @@ -590,11 +600,11 @@ def visualize_ocr_text(xml_file, dir_xml, dir_out): is_vertical = h > 2*w # Check orientation - font = fit_text_single_line(draw, ocr_texts[index], font_path, w, int(h*0.4) ) + font = fit_text_single_line(draw, ocr_texts[index], w, int(h*0.4) ) if is_vertical: - vertical_font = fit_text_single_line(draw, ocr_texts[index], font_path, h, int(w * 0.8)) + vertical_font = fit_text_single_line(draw, ocr_texts[index], h, int(w * 0.8)) text_img = Image.new("RGBA", (h, w), (255, 255, 255, 0)) # Note: dimensions are swapped text_draw = ImageDraw.Draw(text_img) diff --git a/src/eynollah/training/generate_or_update_cnn_rnn_ocr_character_list.py b/src/eynollah/training/generate_or_update_cnn_rnn_ocr_character_list.py new file mode 100644 index 0000000..8620515 --- /dev/null +++ b/src/eynollah/training/generate_or_update_cnn_rnn_ocr_character_list.py @@ -0,0 +1,59 @@ +import os +import numpy as np +import json +import click +import logging + + + +def run_character_list_update(dir_labels, out, current_character_list): + ls_labels = os.listdir(dir_labels) + ls_labels = [ind for ind in ls_labels if ind.endswith('.txt')] + + if current_character_list: + with open(current_character_list, 'r') as f_name: + characters = json.load(f_name) + + characters = set(characters) + else: + characters = set() + + + for ind in ls_labels: + label = open(os.path.join(dir_labels,ind),'r').read().split('\n')[0] + + for char in label: + characters.add(char) + + + characters = sorted(list(set(characters))) + + with open(out, 'w') as f_name: + json.dump(characters, f_name) + + +@click.command() +@click.option( + "--dir_labels", + "-dl", + help="directory of labels which are .txt files", + type=click.Path(exists=True, file_okay=False), + required=True, +) +@click.option( + "--current_character_list", + "-ccl", + help="existing character list in a .txt file that needs to be updated with a set of labels", + type=click.Path(exists=True, file_okay=True), + required=False, +) +@click.option( + "--out", + "-o", + help="An output .txt file where the generated or updated character list will be written", + type=click.Path(exists=False, file_okay=True), +) + +def main(dir_labels, out, current_character_list): + run_character_list_update(dir_labels, out, current_character_list) + diff --git a/src/eynollah/training/gt_gen_utils.py b/src/eynollah/training/gt_gen_utils.py index 796e896..e085fa0 100644 --- a/src/eynollah/training/gt_gen_utils.py +++ b/src/eynollah/training/gt_gen_utils.py @@ -8,7 +8,7 @@ from shapely import geometry from pathlib import Path from PIL import ImageFont from ocrd_utils import bbox_from_points - +from eynollah.utils.font import get_font KERNEL = np.ones((5, 5), np.uint8) NS = { 'pc': 'http://schema.primaresearch.org/PAGE/gts/pagecontent/2019-07-15' @@ -18,7 +18,7 @@ with warnings.catch_warnings(): warnings.simplefilter("ignore") -def visualize_image_from_contours_layout(co_par, co_header, co_drop, co_sep, co_image, co_marginal, co_table, co_map, img): +def visualize_image_from_contours_layout(co_par, co_header, co_drop, co_sep, co_image, co_marginal, co_table, co_map, co_music, img): alpha = 0.5 blank_image = np.ones( (img.shape[:]), dtype=np.uint8) * 255 @@ -32,6 +32,7 @@ def visualize_image_from_contours_layout(co_par, co_header, co_drop, co_sep, co_ col_marginal = (106, 90, 205) col_table = (0, 90, 205) col_map = (90, 90, 205) + col_music = (90, 90, 0) if len(co_image)>0: cv2.drawContours(blank_image, co_image, -1, col_image, thickness=cv2.FILLED) # Fill the contour @@ -59,6 +60,9 @@ def visualize_image_from_contours_layout(co_par, co_header, co_drop, co_sep, co_ if len(co_map)>0: cv2.drawContours(blank_image, co_map, -1, col_map, thickness=cv2.FILLED) # Fill the contour + + if len(co_music)>0: + cv2.drawContours(blank_image, co_music, -1, col_music, thickness=cv2.FILLED) # Fill the contour img_final =cv2.cvtColor(blank_image, cv2.COLOR_BGR2RGB) @@ -352,11 +356,11 @@ def get_textline_contours_and_ocr_text(xml_file): ocr_textlines.append(ocr_text_in[0]) return co_use_case, y_len, x_len, ocr_textlines -def fit_text_single_line(draw, text, font_path, max_width, max_height): +def fit_text_single_line(draw, text, max_width, max_height): initial_font_size = 50 font_size = initial_font_size while font_size > 10: # Minimum font size - font = ImageFont.truetype(font_path, font_size) + font = get_font(font_size=font_size)# ImageFont.truetype(font_path, font_size) text_bbox = draw.textbbox((0, 0), text, font=font) # Get text bounding box text_width = text_bbox[2] - text_bbox[0] text_height = text_bbox[3] - text_bbox[1] @@ -366,7 +370,7 @@ def fit_text_single_line(draw, text, font_path, max_width, max_height): font_size -= 2 # Reduce font size and retry - return ImageFont.truetype(font_path, 10) # Smallest font fallback + return get_font(font_size=10)#ImageFont.truetype(font_path, 10) # Smallest font fallback def get_layout_contours_for_visualization(xml_file): tree1 = ET.parse(xml_file, parser = ET.XMLParser(encoding='utf-8')) @@ -389,6 +393,7 @@ def get_layout_contours_for_visualization(xml_file): co_img=[] co_table=[] co_map=[] + co_music=[] co_noise=[] types_text = [] @@ -630,6 +635,31 @@ def get_layout_contours_for_visualization(xml_file): elif vv.tag!=link+'Point' and sumi>=1: break co_map.append(np.array(c_t_in)) + + if tag.endswith('}MusicRegion') or tag.endswith('}musicregion'): + #print('sth') + for nn in root1.iter(tag): + c_t_in=[] + sumi=0 + for vv in nn.iter(): + # check the format of coords + if vv.tag==link+'Coords': + coords=bool(vv.attrib) + if coords: + p_h=vv.attrib['points'].split(' ') + c_t_in.append( np.array( [ [ int(x.split(',')[0]) , int(x.split(',')[1]) ] for x in p_h] ) ) + break + else: + pass + + + if vv.tag==link+'Point': + c_t_in.append([ int(float(vv.attrib['x'])) , int(float(vv.attrib['y'])) ]) + sumi+=1 + #print(vv.tag,'in') + elif vv.tag!=link+'Point' and sumi>=1: + break + co_music.append(np.array(c_t_in)) if tag.endswith('}NoiseRegion') or tag.endswith('}noiseregion'): @@ -656,7 +686,7 @@ def get_layout_contours_for_visualization(xml_file): elif vv.tag!=link+'Point' and sumi>=1: break co_noise.append(np.array(c_t_in)) - return co_text, co_graphic, co_sep, co_img, co_table, co_map, co_noise, y_len, x_len + return co_text, co_graphic, co_sep, co_img, co_table, co_map, co_music, co_noise, y_len, x_len def get_images_of_ground_truth( gt_list, @@ -870,7 +900,7 @@ def get_images_of_ground_truth( types_graphic_label = list(types_graphic_dict.values()) - labels_rgb_color = [ (0,0,0), (255,0,0), (255,125,0), (255,0,125), (125,255,125), (125,125,0), (0,125,255), (0,125,0), (125,125,125), (255,0,255), (125,0,125), (0,255,0),(0,0,255), (0,255,255), (255,125,125), (0,125,125), (0,255,125), (255,125,255), (125,255,0), (125,255,255)] + labels_rgb_color = [ (0,0,0), (255,0,0), (255,125,0), (255,0,125), (125,255,125), (125,125,0), (0,125,255), (0,125,0), (125,125,125), (255,0,255), (125,0,125), (0,255,0),(0,0,255), (0,255,255), (255,125,125), (0,125,125), (0,255,125), (255,125,255), (125,255,0), (125,255,255), (125,125,255)] region_tags=np.unique([x for x in alltags if x.endswith('Region')]) @@ -882,6 +912,7 @@ def get_images_of_ground_truth( co_img=[] co_table=[] co_map=[] + co_music=[] co_noise=[] for tag in region_tags: @@ -966,19 +997,21 @@ def get_images_of_ground_truth( if "rest_as_decoration" in types_graphic: types_graphic_without_decoration = [element for element in types_graphic if element!='rest_as_decoration' and element!='decoration'] if len(types_graphic_without_decoration) == 0: - if "type" in nn.attrib: - c_t_in_graphic['decoration'].append( np.array( [ [ int(x.split(',')[0]) , int(x.split(',')[1]) ] for x in p_h] ) ) + #if "type" in nn.attrib: + c_t_in_graphic['decoration'].append( np.array( [ [ int(x.split(',')[0]) , int(x.split(',')[1]) ] for x in p_h] ) ) elif len(types_graphic_without_decoration) >= 1: if "type" in nn.attrib: if nn.attrib['type'] in types_graphic_without_decoration: c_t_in_graphic[nn.attrib['type']].append( np.array( [ [ int(x.split(',')[0]) , int(x.split(',')[1]) ] for x in p_h] ) ) else: c_t_in_graphic['decoration'].append( np.array( [ [ int(x.split(',')[0]) , int(x.split(',')[1]) ] for x in p_h] ) ) - + else: + c_t_in_graphic['decoration'].append( np.array( [ [ int(x.split(',')[0]) , int(x.split(',')[1]) ] for x in p_h] ) ) else: if "type" in nn.attrib: if nn.attrib['type'] in all_defined_graphic_types: - c_t_in_graphic[nn.attrib['type']].append( np.array( [ [ int(x.split(',')[0]) , int(x.split(',')[1]) ] for x in p_h] ) ) + c_t_in_graphic[nn.attrib['type']].append( np.array( [ [ int(x.split(',')[0]) , int(x.split(',')[1]) ] for x in p_h] ) ) + break else: @@ -989,9 +1022,9 @@ def get_images_of_ground_truth( if "rest_as_decoration" in types_graphic: types_graphic_without_decoration = [element for element in types_graphic if element!='rest_as_decoration' and element!='decoration'] if len(types_graphic_without_decoration) == 0: - if "type" in nn.attrib: - c_t_in_graphic['decoration'].append( [ int(float(vv.attrib['x'])) , int(float(vv.attrib['y'])) ] ) - sumi+=1 + #if "type" in nn.attrib: + c_t_in_graphic['decoration'].append( [ int(float(vv.attrib['x'])) , int(float(vv.attrib['y'])) ] ) + sumi+=1 elif len(types_graphic_without_decoration) >= 1: if "type" in nn.attrib: if nn.attrib['type'] in types_graphic_without_decoration: @@ -1000,6 +1033,9 @@ def get_images_of_ground_truth( else: c_t_in_graphic['decoration'].append( [ int(float(vv.attrib['x'])) , int(float(vv.attrib['y'])) ] ) sumi+=1 + else: + c_t_in_graphic['decoration'].append( [ int(float(vv.attrib['x'])) , int(float(vv.attrib['y'])) ] ) + sumi+=1 else: if "type" in nn.attrib: @@ -1118,6 +1154,32 @@ def get_images_of_ground_truth( elif vv.tag!=link+'Point' and sumi>=1: break co_map.append(np.array(c_t_in)) + + if 'musicregion' in keys: + if tag.endswith('}MusicRegion') or tag.endswith('}musicregion'): + #print('sth') + for nn in root1.iter(tag): + c_t_in=[] + sumi=0 + for vv in nn.iter(): + # check the format of coords + if vv.tag==link+'Coords': + coords=bool(vv.attrib) + if coords: + p_h=vv.attrib['points'].split(' ') + c_t_in.append( np.array( [ [ int(x.split(',')[0]) , int(x.split(',')[1]) ] for x in p_h] ) ) + break + else: + pass + + + if vv.tag==link+'Point': + c_t_in.append([ int(float(vv.attrib['x'])) , int(float(vv.attrib['y'])) ]) + sumi+=1 + #print(vv.tag,'in') + elif vv.tag!=link+'Point' and sumi>=1: + break + co_music.append(np.array(c_t_in)) if 'noiseregion' in keys: if tag.endswith('}NoiseRegion') or tag.endswith('}noiseregion'): @@ -1195,6 +1257,10 @@ def get_images_of_ground_truth( erosion_rate = 0#2 dilation_rate = 3#4 co_map, img_boundary = update_region_contours(co_map, img_boundary, erosion_rate, dilation_rate, y_len, x_len ) + if "musicregion" in elements_with_artificial_class: + erosion_rate = 0#2 + dilation_rate = 3#4 + co_music, img_boundary = update_region_contours(co_music, img_boundary, erosion_rate, dilation_rate, y_len, x_len ) @@ -1222,6 +1288,8 @@ def get_images_of_ground_truth( img_poly=cv2.fillPoly(img, pts =co_table, color=labels_rgb_color[ config_params['tableregion']]) if 'mapregion' in keys: img_poly=cv2.fillPoly(img, pts =co_map, color=labels_rgb_color[ config_params['mapregion']]) + if 'musicregion' in keys: + img_poly=cv2.fillPoly(img, pts =co_music, color=labels_rgb_color[ config_params['musicregion']]) if 'noiseregion' in keys: img_poly=cv2.fillPoly(img, pts =co_noise, color=labels_rgb_color[ config_params['noiseregion']]) @@ -1286,6 +1354,9 @@ def get_images_of_ground_truth( if 'mapregion' in keys: color_label = config_params['mapregion'] img_poly=cv2.fillPoly(img, pts =co_map, color=(color_label,color_label,color_label)) + if 'musicregion' in keys: + color_label = config_params['musicregion'] + img_poly=cv2.fillPoly(img, pts =co_music, color=(color_label,color_label,color_label)) if 'noiseregion' in keys: color_label = config_params['noiseregion'] img_poly=cv2.fillPoly(img, pts =co_noise, color=(color_label,color_label,color_label)) diff --git a/src/eynollah/training/inference.py b/src/eynollah/training/inference.py index 2be937d..58ccb22 100644 --- a/src/eynollah/training/inference.py +++ b/src/eynollah/training/inference.py @@ -132,15 +132,31 @@ class SBBPredict: self.model = Model( self.model.get_layer(name = "image").input, self.model.get_layer(name = "dense2").output) + assert isinstance(self.model, Model) + elif self.task == "transformer-ocr": + import torch + from transformers import VisionEncoderDecoderModel, TrOCRProcessor + + self.model = VisionEncoderDecoderModel.from_pretrained(self.model_dir) + self.processor = TrOCRProcessor.from_pretrained(self.model_dir) + + if self.cpu: + self.device = torch.device('cpu') + else: + self.device = torch.device('cuda:0') + + self.model.to(self.device) + + assert isinstance(self.model, torch.nn.Module) else: self.model = load_model(self.model_dir, compile=False, custom_objects={"PatchEncoder": PatchEncoder, "Patches": Patches}) + assert isinstance(self.model, Model) ##if self.weights_dir!=None: ##self.model.load_weights(self.weights_dir) - assert isinstance(self.model, Model) if self.task != 'classification' and self.task != 'reading_order': last = self.model.layers[-1] self.img_height = last.output_shape[1] @@ -230,6 +246,13 @@ class SBBPredict: pred_texts = decode_batch_predictions(preds, num_to_char) pred_texts = pred_texts[0].replace("[UNK]", "") return pred_texts + + elif self.task == "transformer-ocr": + from PIL import Image + image = Image.open(image_dir).convert("RGB") + pixel_values = self.processor(image, return_tensors="pt").pixel_values + generated_ids = self.model.generate(pixel_values.to(self.device)) + return self.processor.batch_decode(generated_ids, skip_special_tokens=True)[0] elif self.task == 'reading_order': @@ -566,6 +589,8 @@ class SBBPredict: cv2.imwrite(self.save,res) elif self.task == "cnn-rnn-ocr": print(f"Detected text: {res}") + elif self.task == "transformer-ocr": + print(f"Detected text: {res}") else: img_seg_overlayed, only_layout = self.visualize_model_output(res, self.img_org, self.task) if self.save: @@ -672,7 +697,7 @@ def main(image, dir_in, model, patches, save, save_layout, ground_truth, xml_fil with open(os.path.join(model,'config.json')) as f: config_params_model = json.load(f) task = config_params_model['task'] - if task not in ['classification', 'reading_order', "cnn-rnn-ocr"]: + if task not in ['classification', 'reading_order', "cnn-rnn-ocr", "transformer-ocr"]: assert not image or save, "For segmentation or binarization, an input single image -i also requires an output filename -s" assert not dir_in or out, "For segmentation or binarization, an input directory -di also requires an output directory -o" x = SBBPredict(image, dir_in, model, task, config_params_model, diff --git a/src/eynollah/training/train.py b/src/eynollah/training/train.py index 62d8e51..a6b5331 100644 --- a/src/eynollah/training/train.py +++ b/src/eynollah/training/train.py @@ -3,6 +3,7 @@ import sys import io import json import click +from typing import Optional from tqdm import tqdm import requests @@ -397,7 +398,7 @@ def run(_config, f1_threshold_classification=None, classification_classes_name=None, ## if task=cnn-rnn-ocr - characters_txt_file=None, + characters_txt_file: Optional[str]=None, color_padding_rotation=False, thetha_padd=None, bin_deg=False, @@ -698,6 +699,79 @@ def run(_config, callbacks=callbacks, initial_epoch=index_start) + elif task=="transformer-ocr": + import torch + from torch.utils.data import Dataset as TorchDataset + from transformers import TrOCRProcessor, VisionEncoderDecoderModel, Seq2SeqTrainer, Seq2SeqTrainingArguments, default_data_collator + dir_img, dir_lab = get_dirs_or_files(dir_train) + + if continue_training: + model = VisionEncoderDecoderModel.from_pretrained(dir_of_start_model) + else: + model = VisionEncoderDecoderModel.from_pretrained("microsoft/trocr-base-printed") + + processor = TrOCRProcessor.from_pretrained("microsoft/trocr-base-printed") + + + # Create a DataLoader + class TransformerOCRTorchDataset(TorchDataset): + """ + Wraps preprocess_imgs in a format consumable by torch + """ + def __init__(self, config, dir_img, dir_lab): + self.samples = preprocess_imgs( + config, + dir_img, + dir_lab, + processor=processor, + ) + + def __len__(self): + return len(self.samples) + + def __iter__(self): + yield from self.samples + + dataset = TransformerOCRTorchDataset(_config, dir_img, dir_lab) + data_loader = torch.utils.data.DataLoader(dataset, batch_size=1) + train_dataset = data_loader.dataset + + # set special tokens used for creating the decoder_input_ids from the labels + model.config.decoder_start_token_id = processor.tokenizer.cls_token_id + model.config.pad_token_id = processor.tokenizer.pad_token_id + # make sure vocab size is set correctly + model.config.vocab_size = model.config.decoder.vocab_size + + # set beam search parameters + model.config.eos_token_id = processor.tokenizer.sep_token_id + model.config.max_length = max_len + model.config.early_stopping = True + model.config.no_repeat_ngram_size = 3 + model.config.length_penalty = 2.0 + model.config.num_beams = 4 + + + training_args = Seq2SeqTrainingArguments( + predict_with_generate=True, + num_train_epochs=n_epochs, + learning_rate=learning_rate, + per_device_train_batch_size=n_batch, + fp16=True, + output_dir=dir_output, + logging_steps=2, + save_steps=save_interval, + ) + + # instantiate trainer + trainer = Seq2SeqTrainer( + model=model, + tokenizer=processor.feature_extractor, + args=training_args, + train_dataset=train_dataset, + data_collator=default_data_collator, + ) + trainer.train() + elif task=='classification': if continue_training: model = load_model(dir_of_start_model, compile=False) @@ -741,7 +815,7 @@ def run(_config, usable_checkpoints = [os.path.join(dir_output, 'model_{epoch:02d}'.format(epoch=epoch + 1)) for epoch in usable_checkpoints] ens_path = os.path.join(dir_output, 'model_ens_avg') - run_ensembling(usable_checkpoints, ens_path) + run_ensembling(usable_checkpoints, ens_path, framework='tensorflow') _log.info("ensemble model saved under '%s'", ens_path) elif task=='reading_order': diff --git a/src/eynollah/training/utils.py b/src/eynollah/training/utils.py index 33a1fd2..e84ac66 100644 --- a/src/eynollah/training/utils.py +++ b/src/eynollah/training/utils.py @@ -789,7 +789,7 @@ def preprocess_imgs(config, lab = cv2.imread(os.path.join(dir_lab, img_name + '.png')) elif config['task'] == "enhancement": lab = cv2.imread(os.path.join(dir_lab, img)) - elif config['task'] == "cnn-rnn-ocr": + elif config['task'] in ["cnn-rnn-ocr", "transformer-ocr"]: # assert lab == 'img_name + '.txt' with open(os.path.join(dir_lab, img_name + '.txt'), 'r') as f: lab = f.read().split('\n')[0] @@ -797,7 +797,7 @@ def preprocess_imgs(config, lab = None try: - if config['task'] == "cnn-rnn-ocr": + if config['task'] in ["cnn-rnn-ocr", "transformer-ocr"]: yield from preprocess_img_ocr(img, img_name, lab, **config) continue else: @@ -1116,14 +1116,25 @@ def preprocess_img_ocr( number_of_backgrounds_per_image=None, list_all_possible_background_images=None, list_all_possible_foreground_rgbs=None, + task=None, + processor=None, **kwargs ): def scale_image(img): return scale_padd_image_for_ocr(img, input_height, input_width).astype(np.float32) / 255. #lab = vectorize_label(lab, char_to_num, padding_token, max_len) # now padded at Dataset.padded_batch - lab = char_to_num(tf.strings.unicode_split(lab, input_encoding="UTF-8")) - yield scale_image(img), lab + if task == 'cnn-rnn-ocr': + assert char_to_num, 'task is cnn-rnn-ocr, so preprocess_imgs_ocr should be passed "char_to_num"' + lab = char_to_num(tf.strings.unicode_split(lab, input_encoding="UTF-8")) + yield_encoder = lambda x: x + elif task == 'transformer-ocr': + assert processor, 'task is transformer-ocr, so preprocess_imgs_ocr should be passed "processor"' + # TODO make max_length configurable again, if deemed sensible + lab = [l if l != self.processor.tokenizer.pad_token_id else -100 + for l in processor.tokenizer(lab, padding="max_length", max_length=128).input_ids] + yield_encoder = lambda img_, lab_: {"pixel_values": processor(Image.fromarray(img_), return_tensors="pt").pixel_values.squeeze(), "labels": torch.tensor(lab_)} + yield yield_encoder(scale_image(img), lab) #to_yield = {"image": ret_x, "label": ret_y} if dir_img_bin: @@ -1139,32 +1150,32 @@ def preprocess_img_ocr( for padd_col in padd_colors: img_pad = do_padding_for_ocr(img, 1.2, padd_col) img_rot = rotation_not_90_func_single_image(img_pad, thetha_ind) - yield scale_image(img_rot), lab + yield yield_encoder(scale_image(img_rot), lab) if rotation_not_90: for thetha_ind in thetha: img_rot = rotation_not_90_func_single_image(img, thetha_ind) - yield scale_image(img_rot), lab + yield yield_encoder(scale_image(img_rot), lab) if blur_aug: for blur_type in blur_k: img_blur = bluring(img, blur_type) - yield scale_image(img_blur), lab + yield yield_encoder(scale_image(img_blur), lab) if degrading: for deg_scale_ind in degrade_scales: img_deg = do_degrading(img, deg_scale_ind) - yield scale_image(img_deg), lab + yield yield_encoder(scale_image(img_deg), lab) if bin_deg: for deg_scale_ind in degrade_scales: img_deg = do_degrading(img_bin_corr, deg_scale_ind) - yield scale_image(img_deg), lab + yield yield_encoder(scale_image(img_deg), lab) if brightening: for bright_scale_ind in brightness: img_bright = do_brightening(img, bright_scale_ind) - yield scale_image(img_bright), lab + yield yield_encoder(scale_image(img_bright), lab) if padding_white: for padding_size in white_padds: for padd_col in padd_colors: img_pad = do_padding_for_ocr(img, padding_size, padd_col) - yield scale_image(img_pad), lab + yield yield_encoder(scale_image(img_pad), lab) if adding_rgb_foreground: for i_n in range(number_of_backgrounds_per_image): background_image_chosen_name = random.choice(list_all_possible_background_images) @@ -1178,7 +1189,7 @@ def preprocess_img_ocr( img_fg = \ return_binary_image_with_given_rgb_background_and_given_foreground_rgb( img_bin_corr, img_rgb_background_chosen, foreground_rgb_chosen) - yield scale_image(img_fg), lab + yield yield_encoder(scale_image(img_fg), lab) if adding_rgb_background: for i_n in range(number_of_backgrounds_per_image): background_image_chosen_name = random.choice(list_all_possible_background_images) @@ -1186,59 +1197,59 @@ def preprocess_img_ocr( cv2.imread(dir_rgb_backgrounds + '/' + background_image_chosen_name) img_bg = \ return_binary_image_with_given_rgb_background(img_bin_corr, img_rgb_background_chosen) - yield scale_image(img_bg), lab + yield yield_encoder(scale_image(img_bg), lab) if binarization: - yield scale_image(img_bin_corr), lab + yield yield_encoder(scale_image(img_bin_corr), lab) if image_inversion: img_inv = invert_image(img_bin_corr) - yield scale_image(img_inv), lab + yield yield_encoder(scale_image(img_inv), lab) if channels_shuffling: for shuffle_index in shuffle_indexes: img_shuf = return_shuffled_channels(img, shuffle_index) - yield scale_image(img_shuf), lab + yield yield_encoder(scale_image(img_shuf), lab) if add_red_textlines: img_red = return_image_with_red_elements(img, img_bin_corr) - yield scale_image(img_red), lab + yield yield_encoder(scale_image(img_red), lab) if white_noise_strap: img_noisy = return_image_with_strapped_white_noises(img) - yield scale_image(img_noisy), lab + yield yield_encoder(scale_image(img_noisy), lab) if textline_skewing: for des_scale_ind in skewing_amplitudes: img_rot = do_deskewing(img, des_scale_ind) - yield scale_image(img_rot), lab + yield yield_encoder(scale_image(img_rot), lab) if textline_skewing_bin: for des_scale_ind in skewing_amplitudes: img_rot = do_deskewing(img_bin_corr, des_scale_ind) - yield scale_image(img_rot), lab + yield yield_encoder(scale_image(img_rot), lab) if textline_left_in_depth: img_warp = do_direction_in_depth(img, 'left') - yield scale_image(img_warp), lab + yield yield_encoder(scale_image(img_warp), lab) if textline_left_in_depth_bin: img_warp = do_direction_in_depth(img_bin_corr, 'left') - yield scale_image(img_warp), lab + yield yield_encoder(scale_image(img_warp), lab) if textline_right_in_depth: img_warp = do_direction_in_depth(img, 'right') - yield scale_image(img_warp), lab + yield yield_encoder(scale_image(img_warp), lab) if textline_right_in_depth_bin: img_warp = do_direction_in_depth(img_bin_corr, 'right') - yield scale_image(img_warp), lab + yield yield_encoder(scale_image(img_warp), lab) if textline_up_in_depth: img_warp = do_direction_in_depth(img, 'up') - yield scale_image(img_warp), lab + yield yield_encoder(scale_image(img_warp), lab) if textline_up_in_depth_bin: img_warp = do_direction_in_depth(img_bin_corr, 'up') - yield scale_image(img_warp), lab + yield yield_encoder(scale_image(img_warp), lab) if textline_down_in_depth: img_warp = do_direction_in_depth(img, 'down') - yield scale_image(img_warp), lab + yield yield_encoder(scale_image(img_warp), lab) if textline_down_in_depth_bin: img_warp = do_direction_in_depth(img_bin_corr, 'down') - yield scale_image(img_warp), lab + yield yield_encoder(scale_image(img_warp), lab) if pepper_aug: for pepper_ind in pepper_indexes: img_noisy = add_salt_and_pepper_noise(img, pepper_ind, pepper_ind) - yield scale_image(img_noisy), lab + yield yield_encoder(scale_image(img_noisy), lab) if pepper_bin_aug: for pepper_ind in pepper_indexes: img_noisy = add_salt_and_pepper_noise(img_bin_corr, pepper_ind, pepper_ind) - yield scale_image(img_noisy), lab + yield yield_encoder(scale_image(img_noisy), lab) diff --git a/src/eynollah/training/weights_ensembling.py b/src/eynollah/training/weights_ensembling.py index f651c56..1c175c0 100644 --- a/src/eynollah/training/weights_ensembling.py +++ b/src/eynollah/training/weights_ensembling.py @@ -1,4 +1,5 @@ import os +from typing import Optional from warnings import catch_warnings, simplefilter import click @@ -11,33 +12,56 @@ from ocrd_utils import tf_disable_interactive_logs tf_disable_interactive_logs() import tensorflow as tf from tensorflow.keras.models import load_model +import torch +from transformers import VisionEncoderDecoderModel from ..patch_encoder import ( PatchEncoder, Patches, ) -def run_ensembling(model_dirs, out_dir): - all_weights = [] - - for model_dir in model_dirs: - assert os.path.isdir(model_dir), model_dir - model = load_model(model_dir, compile=False, - custom_objects=dict(PatchEncoder=PatchEncoder, - Patches=Patches)) - all_weights.append(model.get_weights()) +def run_ensembling(dir_models, out, framework): + ls_models = os.listdir(dir_models) + # model: Optional[VisionEncoderDecoderModel] = None + # model_name: Optional[str] = None + if framework=="torch": + models = [] + sd_models = [] - new_weights = [] - for layer_weights in zip(*all_weights): - layer_weights = np.array([np.array(weights).mean(axis=0) - for weights in zip(*layer_weights)]) - new_weights.append(layer_weights) + for model_name in ls_models: + model = VisionEncoderDecoderModel.from_pretrained(os.path.join(dir_models, model_name)) + models.append(model) + sd_models.append(model.state_dict()) + for key in sd_models[0]: + sd_models[0][key] = sum(sd[key] for sd in sd_models) / len(sd_models) + + model.load_state_dict(sd_models[0]) + os.system("mkdir "+out) + torch.save(model.state_dict(), os.path.join(out, "pytorch_model.bin")) + os.system('cp ' + os.path.join(os.path.join(dir_models, model_name), "config.json") + " " + out) + + else: + weights=[] - #model = tf.keras.models.clone_model(model) - model.set_weights(new_weights) + for model_name in ls_models: + model = load_model(os.path.join(dir_models, model_name), compile=False, custom_objects={'PatchEncoder':PatchEncoder, 'Patches': Patches}) + weights.append(model.get_weights()) + + new_weights = list() - model.save(out_dir) - os.system('cp ' + os.path.join(model_dirs[0], "config.json ") + out_dir + "/") + for weights_list_tuple in zip(*weights): + new_weights.append( + [np.array(weights_).mean(axis=0)\ + for weights_ in zip(*weights_list_tuple)]) + + + + new_weights = [np.array(x) for x in new_weights] + + model.set_weights(new_weights) + model.save(out) + os.system('cp '+os.path.join(os.path.join(dir_models, model_name), "config.json") + " " + out) + os.system('cp '+os.path.join(os.path.join(dir_models, model_name), "characters_org.txt") + " " + out) @click.command() @click.option( @@ -56,12 +80,19 @@ def run_ensembling(model_dirs, out_dir): required=True, type=click.Path(exists=False, file_okay=False), ) -def ensemble_cli(in_, out): +@click.option( + "--framework", + "-fw", + help="this parameter gets tensorflow or torch as model framework", + type=click.Choice(['torch', 'tensorflow']), + default="tensorflow" +) + +def ensemble_cli(in_, out, framework): """ mix multiple model weights Load a sequence of models and mix them into a single ensemble model by averaging their weights. Write the resulting model. """ - run_ensembling(in_, out) - + run_ensembling(in_, out, framework) diff --git a/src/eynollah/utils/font.py b/src/eynollah/utils/font.py index 939933e..3e9e588 100644 --- a/src/eynollah/utils/font.py +++ b/src/eynollah/utils/font.py @@ -9,8 +9,8 @@ else: import importlib.resources as importlib_resources -def get_font(): +def get_font(font_size): #font_path = "Charis-7.000/Charis-Regular.ttf" # Make sure this file exists! - font = importlib_resources.files(__package__) / "../Charis-Regular.ttf" + font = importlib_resources.files(__package__) / "../Amiri-Regular.ttf" with importlib_resources.as_file(font) as font: - return ImageFont.truetype(font=font, size=40) + return ImageFont.truetype(font=font, size=font_size) diff --git a/train/config_params_trocr.json b/train/config_params_trocr.json new file mode 100644 index 0000000..34c6376 --- /dev/null +++ b/train/config_params_trocr.json @@ -0,0 +1,82 @@ +{ + "backbone_type" : "transformer", + "task": "transformer-ocr", + "n_classes" : 2, + "max_len": 192, + "n_epochs" : 1, + "input_height" : 32, + "input_width" : 512, + "weight_decay" : 1e-6, + "n_batch" : 1, + "learning_rate": 1e-5, + "save_interval": 1500, + "patches" : false, + "pretraining" : false, + "augmentation" : true, + "flip_aug" : false, + "blur_aug" : true, + "scaling" : false, + "adding_rgb_background": true, + "adding_rgb_foreground": true, + "add_red_textlines": true, + "white_noise_strap": true, + "textline_right_in_depth": true, + "textline_left_in_depth": true, + "textline_up_in_depth": true, + "textline_down_in_depth": true, + "textline_right_in_depth_bin": true, + "textline_left_in_depth_bin": true, + "textline_up_in_depth_bin": true, + "textline_down_in_depth_bin": true, + "bin_deg": true, + "textline_skewing": true, + "textline_skewing_bin": true, + "channels_shuffling": true, + "degrading": true, + "brightening": true, + "binarization" : true, + "pepper_aug": true, + "pepper_bin_aug": true, + "image_inversion": true, + "scaling_bluring" : false, + "scaling_binarization" : false, + "scaling_flip" : false, + "rotation": false, + "color_padding_rotation": true, + "padding_white": true, + "rotation_not_90": true, + "transformer_num_patches_xy": [56, 56], + "transformer_patchsize_x": 4, + "transformer_patchsize_y": 4, + "transformer_projection_dim": 64, + "transformer_mlp_head_units": [128, 64], + "transformer_layers": 1, + "transformer_num_heads": 1, + "transformer_cnn_first": false, + "blur_k" : ["blur","gauss","median"], + "padd_colors" : ["white", "black"], + "scales" : [0.6, 0.7, 0.8, 0.9], + "brightness" : [1.3, 1.5, 1.7, 2], + "degrade_scales" : [0.2, 0.4], + "pepper_indexes": [0.01, 0.005], + "skewing_amplitudes" : [5, 8], + "flip_index" : [0, 1, -1], + "shuffle_indexes" : [ [0,2,1], [1,2,0], [1,0,2] , [2,1,0]], + "thetha" : [0.1, 0.2, -0.1, -0.2], + "thetha_padd": [-0.6, -1, -1.4, -1.8, 0.6, 1, 1.4, 1.8], + "white_padds" : [0.1, 0.3, 0.5, 0.7, 0.9], + "number_of_backgrounds_per_image": 2, + "continue_training": false, + "index_start" : 0, + "dir_of_start_model" : " ", + "weighted_loss": false, + "is_loss_soft_dice": false, + "data_is_provided": false, + "dir_train": "/home/vahid/extracted_lines/1919_bin/train", + "dir_eval": "/home/vahid/Documents/test/sbb_pixelwise_segmentation/test_label/pageextractor_test/eval_new", + "dir_output": "/home/vahid/extracted_lines/1919_bin/output", + "dir_rgb_backgrounds": "/home/vahid/Documents/1_2_test_eynollah/set_rgb_background", + "dir_rgb_foregrounds": "/home/vahid/Documents/1_2_test_eynollah/out_set_rgb_foreground", + "dir_img_bin": "/home/vahid/extracted_lines/1919_bin/images_bin" + +} diff --git a/train/requirements.txt b/train/requirements.txt index 1734b67..cd4d74b 100644 --- a/train/requirements.txt +++ b/train/requirements.txt @@ -6,6 +6,10 @@ imutils scipy tensorflow-addons # for connected_components, depublished and only compatible with tensorflow < 2.16 tensorflow < 2.16 # for tensorflow-addons, so only needed in training -tf-keras # avoid keras 3 (also needs TF_USE_LEGACY_KERAS=1) -tf-data < 2.16 # for tensorflow-addons, so only needed in training +tf-keras < 2.16 # avoid keras 3 (also needs TF_USE_LEGACY_KERAS=1) protobuf < 5 # for tensorflow-addons, so only needed in training +torch +transformers <= 4.30.2 ; python_version < '3.10' +transformers >= 5 ; python_version >= '3.10' +eynollah-fork-tf2onnx == 1.17.0.post2 +ml_dtypes >= 0.5