Views
No views yet
1import torch
2from PIL import Image
3import torch.nn.functional as F
4from networks import ErnieLayoutConfig, ErnieLayoutForQuestionAnswering, \
5 ErnieLayoutProcessor, ErnieLayoutTokenizerFast
6from transformers.models.layoutlmv3 import LayoutLMv3ImageProcessor
7
8pretrain_torch_model_or_path = "Norm/ERNIE-Layout-Pytorch"
9doc_imag_path = "./dummy_input.jpeg"
10
11context = ['This is an example sequence', 'All ocr boxes are inserted into this list']
12layout = [[381, 91, 505, 115], [738, 96, 804, 122]] # make sure all boxes are normalized between 0 - 1000
13pil_image = Image.open(doc_imag_path).convert("RGB")
14
15# initialize tokenizer
16tokenizer = ErnieLayoutTokenizerFast.from_pretrained(pretrained_model_name_or_path=pretrain_torch_model_or_path)
17
18# initialize feature extractor
19feature_extractor = LayoutLMv3ImageProcessor(apply_ocr=False)
20processor = ErnieLayoutProcessor(image_processor=feature_extractor, tokenizer=tokenizer)
21
22# Tokenize context & questions
23question = "what is it?"
24encoding = processor(pil_image, question, context, boxes=layout, return_tensors="pt")
25
26# dummy answer start && end index
27start_positions = torch.tensor([6])
28end_positions = torch.tensor([12])
29
30# initialize config
31config = ErnieLayoutConfig.from_pretrained(pretrained_model_name_or_path=pretrain_torch_model_or_path)
32config.num_classes = 2 # start and end
33
34# initialize ERNIE for VQA
35model = ErnieLayoutForQuestionAnswering.from_pretrained(
36 pretrained_model_name_or_path=pretrain_torch_model_or_path,
37 config=config,
38)
39
40output = model(**encoding, start_positions=start_positions, end_positions=end_positions)
41
42# decode output
43start_max = torch.argmax(F.softmax(output.start_logits, dim=-1))
44end_max = torch.argmax(F.softmax(output.end_logits, dim=-1)) + 1 # add one ##because of python list indexing
45answer = tokenizer.decode(encoding.input_ids[0][start_max: end_max])
46print(answer)
47
48