Donut consists of a vision encoder (Swin Transformer) and a text decoder (BART).
Given an image, the encoder first encodes the image into a tensor of embeddings (of shape batch_size, seq_len, hidden_size),
after which the decoder autoregressively generates text, conditioned on the encoding of the encoder.
Donut_architecture
1 import torch
2 from PIL import Image
3 from transformers import DonutProcessor , VisionEncoderDecoderConfig , VisionEncoderDecoderModel
4
5 model_id = "hf-tuner/donut-base-finetuned-sroie"
6
7 config = VisionEncoderDecoderConfig . from_pretrained ( model_id )
8 processor = DonutProcessor . from_pretrained ( model_id )
9 device = "cuda" if torch . cuda . is_available ( ) else "cpu"
10 config . dtype = torch . float16 if torch . cuda . is_available ( ) else torch . float32
11
12 model = VisionEncoderDecoderModel . from_pretrained ( ckpt , config = config )
13 model . to ( device )
14
15 # inference
16 task_start_token = "<parsing>"
17
18 image = Image . open ( "test-receipt.png" )
19
20 pixel_values = processor ( image , return_tensors = "pt" ) . pixel_values . to ( device )
21 decoder_input_ids = processor . tokenizer ( task_start_token , add_special_tokens = False , return_tensors = "pt" ) . input_ids . to ( device )
22
23 generated_ids = model . generate ( pixel_values ,
24 decoder_input_ids = decoder_input_ids ,
25 max_length = 128 ,
26 bad_words_ids = [ [ processor . tokenizer . unk_token_id ] ]
27 )
28
29 generated_text = processor . batch_decode ( generated_ids , skip_special_tokens = False ) [ 0 ]
30 processor . token2json ( generated_text )
31
32 ### Output
33 # {'total': '62.60',
34 # 'date': '30/10/2017',
35 # 'company': 'gardenia bakeries (kl) sdn bhd',
36 # 'address': 'lot 3, jalan pelabur 23/1, 40300 shah alam, selangor.'}
37