Views
No views yet
1from transformers.models.deit.modeling_deit import DeiTPatchEmbeddings, DeiTEmbeddings
2from transformers import TrOCRProcessor, VisionEncoderDecoderModel
3import torch.nn.functional as F
4from PIL import Image
5import torch
6
7def apply_deit_custom_size_patches():
8 """Apply patches to DeiT model to support custom image sizes"""
9
10 def deit_patch_forward(self, pixel_values, interpolate_pos_encoding=None):
11 embeddings = self.projection(pixel_values).flatten(2).transpose(1, 2)
12 return embeddings
13
14 def deit_embeddings_forward(self, pixel_values, bool_masked_pos=None, interpolate_pos_encoding=None):
15 batch_size, num_channels, height, width = pixel_values.shape
16 embeddings = self.patch_embeddings(pixel_values, interpolate_pos_encoding)
17
18 cls_tokens = self.cls_token.expand(batch_size, -1, -1)
19 distillation_tokens = self.distillation_token.expand(batch_size, -1, -1)
20 embeddings = torch.cat((cls_tokens, distillation_tokens, embeddings), dim=1)
21
22 patch_size = self.patch_embeddings.patch_size[0]
23 num_patches_h = height // patch_size
24 num_patches_w = width // patch_size
25 num_patches = num_patches_h * num_patches_w
26
27 pos_embed = self.position_embeddings
28
29 if num_patches + 2 != pos_embed.shape[1]:
30 special_pos_embed = pos_embed[:, :2, :]
31 patch_pos_embed = pos_embed[:, 2:, :]
32
33 orig_size = int(patch_pos_embed.shape[1] ** 0.5)
34 embed_dim = patch_pos_embed.shape[2]
35
36 patch_pos_embed = patch_pos_embed.reshape(1, orig_size, orig_size, embed_dim)
37 patch_pos_embed = patch_pos_embed.permute(0, 3, 1, 2)
38 patch_pos_embed = F.interpolate(patch_pos_embed,
39 size=(num_patches_h, num_patches_w),
40 mode='bicubic',
41 align_corners=False)
42 patch_pos_embed = patch_pos_embed.permute(0, 2, 3, 1).reshape(1, num_patches, embed_dim)
43
44 pos_embed = torch.cat([special_pos_embed, patch_pos_embed], dim=1)
45
46 embeddings = embeddings + pos_embed
47 embeddings = self.dropout(embeddings)
48
49 return embeddings
50
51 DeiTPatchEmbeddings.forward = deit_patch_forward
52 DeiTEmbeddings.forward = deit_embeddings_forward
53
54# Use it at the start of your inference script
55apply_deit_custom_size_patches()
56
57# Load model and processor
58processor = TrOCRProcessor.from_pretrained("Kansallisarkisto/multicentury-htr-model-small",
59 use_fast=True,
60 do_resize=True,
61 size={'height': 192,'width': 1024})
62
63model = VisionEncoderDecoderModel.from_pretrained("Kansallisarkisto/multicentury-htr-model-small")
64
65# Open an image of handwritten text
66image = Image.open("path_to_image.jpg")
67
68# Preprocess and predict
69pixel_values = processor(image, return_tensors="pt").pixel_values
70generated_ids = model.generate(pixel_values)
71generated_text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
72
73print(generated_text)