acverma/documentAI
0
1# -*- coding: utf-8 -*-2"""DocAI_DeploymentGradio.ipynb3 4Automatically generated by Colaboratory.5 6Original file is located at7 https://colab.research.google.com/drive/1USSEj7nHh2n2hUhTJTC0Iwhj6mSR7-mD8"""9 10import os11os.system('pip3 install torch torchvision torchaudio --extra-index-url https://download.pytorch.org/whl/cpu')12 13os.system('pip install pyyaml==5.1')14 15os.system('pip install -q git+https://github.com/huggingface/transformers.git')16 17os.system('pip install -q datasets seqeval')18 19os.system('pip install torch==1.8.0+cu101 torchvision==0.9.0+cu101 -f https://download.pytorch.org/whl/torch_stable.html')20os.system('pip install -q detectron2 -f https://dl.fbaipublicfiles.com/detectron2/wheels/cu101/torch1.8/index.html')21os.system('pip install -q pytesseract')22 23#!pip install gradio24 25#!pip install -q git+https://github.com/huggingface/transformers.git26 27#!pip install h5py28 29#!pip install -q datasets seqeval30 31import gradio as gr32 33import numpy as np34import tensorflow as tf35 36import torch37import json38 39from datasets.features import ClassLabel40from transformers import AutoProcessor41 42from datasets import Features, Sequence, ClassLabel, Value, Array2D, Array3D43from datasets import load_dataset # this dataset uses the new Image feature :) 44 45from transformers import LayoutLMv3ForTokenClassification46from transformers import AutoModelForTokenClassification47 48#import cv249from PIL import Image, ImageDraw, ImageFont50 51dataset = load_dataset("nielsr/funsd-layoutlmv3")52 53example = dataset["test"][0]54 55#image_path = "/root/.cache/huggingface/datasets/nielsr___funsd-layoutlmv3/funsd/1.0.0/0e3f4efdfd59aa1c3b4952c517894f7b1fc4d75c12ef01bcc8626a69e41c1bb9/funsd-layoutlmv3-test.arrow"56 57image_path = '/root/.cache/huggingface/datasets/nielsr___funsd-layoutlmv3/funsd/1.0.0/0e3f4efdfd59aa1c3b4952c517894f7b1fc4d75c12ef01bcc8626a69e41c1bb9'58 59example = dataset["test"][0]60example["image"].save("example1.png")61 62example1 = dataset["test"][1]63example1["image"].save("example2.png")64 65example2 = dataset["test"][2]66example2["image"].save("example3.png")67 68example2["image"]69 70#Image.open(dataset[2][image_path]).convert("RGB").save("example1.png")71#Image.open(dataset[1]["image_path"]).convert("RGB").save("example2.png")72#Image.open(dataset[0]["image_path"]).convert("RGB").save("example3.png")73 74words, boxes, ner_tags = example["tokens"], example["bboxes"], example["ner_tags"]75 76features = dataset["test"].features77 78column_names = dataset["test"].column_names79image_column_name = "image"80text_column_name = "tokens"81boxes_column_name = "bboxes"82label_column_name = "ner_tags"83 84def get_label_list(labels):85 unique_labels = set()86 for label in labels:87 unique_labels = unique_labels | set(label)88 label_list = list(unique_labels)89 label_list.sort()90 return label_list91 92if isinstance(features[label_column_name].feature, ClassLabel):93 label_list = features[label_column_name].feature.names94 # No need to convert the labels since they are already ints.95 id2label = {k: v for k,v in enumerate(label_list)}96 label2id = {v: k for k,v in enumerate(label_list)}97else:98 label_list = get_label_list(dataset["train"][label_column_name])99 id2label = {k: v for k,v in enumerate(label_list)}100 label2id = {v: k for k,v in enumerate(label_list)}101num_labels = len(label_list)102 103label2color = {'question':'blue', 'answer':'green', 'header':'orange', 'other':'violet'}104 105def prepare_examples(examples):106 images = examples[image_column_name]107 words = examples[text_column_name]108 boxes = examples[boxes_column_name]109 word_labels = examples[label_column_name]110 111 encoding = processor(images, words, boxes=boxes, word_labels=word_labels,112 truncation=True, padding="max_length")113 114 return encoding115 116processor = AutoProcessor.from_pretrained("microsoft/layoutlmv3-base", apply_ocr=False)117 118model = LayoutLMv3ForTokenClassification.from_pretrained("microsoft/layoutlmv3-base",119 id2label=id2label,120 label2id=label2id)121 122# we need to define custom features for `set_format` (used later on) to work properly123features = Features({124 'pixel_values': Array3D(dtype="float32", shape=(3, 224, 224)),125 'input_ids': Sequence(feature=Value(dtype='int64')),126 'attention_mask': Sequence(Value(dtype='int64')),127 'bbox': Array2D(dtype="int64", shape=(512, 4)),128 'labels': Sequence(feature=Value(dtype='int64')),129})130 131eval_dataset = dataset["test"].map(132 prepare_examples,133 batched=True,134 remove_columns=column_names,135 features=features,136)137 138def unnormalize_box(bbox, width, height):139 return [140 width * (bbox[0] / 1000),141 height * (bbox[1] / 1000),142 width * (bbox[2] / 1000),143 height * (bbox[3] / 1000),144 ]145 146def process_image(image):147 148 print(type(image))149 width, height = image.size150 151 image = example["image"]152 words = example["tokens"]153 boxes = example["bboxes"]154 word_labels = example["ner_tags"]155 156 for k,v in encoding.items():157 print(k,v.shape)158 159 # encode160 #encoding = processor(image, truncation=True, return_offsets_mapping=True, return_tensors="pt")161 #offset_mapping = encoding.pop('offset_mapping')162 163 #encoding = processor(image, words, truncation=True,boxes=boxes, word_labels=word_labels,return_offsets_mapping=True, return_tensors="pt")164 #offset_mapping = encoding.pop('offset_mapping')165 166 encoding = processor(image, truncation=True,boxes=boxes, word_labels=word_labels,return_offsets_mapping=True, return_tensors="pt")167 offset_mapping = encoding.pop('offset_mapping')168 169 170 171 # forward pass172 with torch.no_grad():173 outputs = model(**encoding)174 175 # get predictions176 177 # We take the highest score for each token, using argmax. 178 # This serves as the predicted label for each token.179 logits = outputs.logits180 #logits.shape181 predictions = logits.argmax(-1).squeeze().tolist()182 183 labels = encoding.labels.squeeze().tolist()184 185 token_boxes = encoding.bbox.squeeze().tolist()186 width, height = image.size187 188 #true_predictions = [model.config.id2label[pred] for pred, label in zip(predictions, labels) if label != - 100]189 #true_labels = [model.config.id2label[label] for prediction, label in zip(predictions, labels) if label != -100]190 #true_boxes = [unnormalize_box(box, width, height) for box, label in zip(token_boxes, labels) if label != -100]191 192 193 # only keep non-subword predictions194 is_subword = np.array(offset_mapping.squeeze().tolist())[:,0] != 0195 true_predictions = [id2label[pred] for idx, pred in enumerate(predictions) if not is_subword[idx]]196 true_boxes = [unnormalize_box(box, width, height) for idx, box in enumerate(token_boxes) if not is_subword[idx]]197 198 # draw predictions over the image199 draw = ImageDraw.Draw(image)200 font = ImageFont.load_default()201 for prediction, box in zip(true_predictions, true_boxes):202 predicted_label = id2label(prediction)203 draw.rectangle(box, outline=label2color[predicted_label])204 draw.text((box[0]+10, box[1]-10), text=predicted_label, fill=label2color[predicted_label], font=font)205 206 return image207 208title = "DocumentAI - Extraction of Key Information using LayoutLMv3 model"209description = "Extraction of Form or Invoice Extraction - We use Microsoft's LayoutLMv3 trained on Invoice Dataset to predict the Biller Name, Biller Address, Biller post_code, Due_date, GST, Invoice_date, Invoice_number, Subtotal and Total. To use it, simply upload an image or use the example image below. Results will show up in a few seconds."210 211article="<b>References</b><br>[1] Y. Xu et al., “LayoutLMv3: Pre-training for Document AI with Unified Text and Image Masking.” 2022. <a href='https://arxiv.org/abs/2204.08387'>Paper Link</a><br>[2] <a href='https://github.com/NielsRogge/Transformers-Tutorials/tree/master/LayoutLMv3'>LayoutLMv3 training and inference</a>" 212 213examples =[['example1.png'],['example2.png'],['example3.png']]214 215css = """.output_image, .input_image {height: 600px !important}"""216 217iface = gr.Interface(fn=process_image, 218 inputs=gr.inputs.Image(type="pil"), 219 outputs=gr.outputs.Image(type="pil", label="annotated predict image"),220 title=title,221 description=description,222 article=article,223 examples=examples,224 css=css,225 analytics_enabled = True, enable_queue=True226 )227 228iface.launch(inline=False, share=False, debug=False)