inferencerlabs commited on
Commit
56205d8
·
verified ·
1 Parent(s): a3560a5

Upload model file

Browse files
Files changed (1) hide show
  1. processing.py +54 -0
processing.py ADDED
@@ -0,0 +1,54 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from fractions import Fraction
2
+
3
+ from transformers import LlavaNextProcessor
4
+ from transformers.image_processing_utils import select_best_resolution
5
+
6
+
7
+
8
+ class Granite4VisionProcessor(LlavaNextProcessor):
9
+ model_type = "granite4_vision"
10
+
11
+ def __init__(
12
+ self,
13
+ image_processor=None,
14
+ tokenizer=None,
15
+ patch_size=None,
16
+ vision_feature_select_strategy=None,
17
+ chat_template=None,
18
+ image_token="<image>", # set the default and let users change if they have peculiar special tokens in rare cases
19
+ num_additional_image_tokens=0,
20
+ downsample_rate=None,
21
+ **kwargs,
22
+ ):
23
+ super().__init__(image_processor=image_processor,
24
+ tokenizer=tokenizer,
25
+ patch_size=patch_size,
26
+ vision_feature_select_strategy=vision_feature_select_strategy,
27
+ chat_template=chat_template,
28
+ image_token=image_token,
29
+ num_additional_image_tokens=num_additional_image_tokens,
30
+ )
31
+ self.downsample_rate = downsample_rate
32
+
33
+ def _get_number_of_features(self, orig_height: int, orig_width: int, height: int, width: int) -> int:
34
+ image_grid_pinpoints = self.image_processor.image_grid_pinpoints
35
+
36
+ height_best_resolution, width_best_resolution = select_best_resolution(
37
+ [orig_height, orig_width], image_grid_pinpoints
38
+ )
39
+ scale_height, scale_width = height_best_resolution // height, width_best_resolution // width
40
+
41
+ patches_height = height // self.patch_size
42
+ patches_width = width // self.patch_size
43
+ if self.downsample_rate is not None:
44
+ ds_rate = Fraction(self.downsample_rate)
45
+ patches_height = int(patches_height * ds_rate)
46
+ patches_width = int(patches_width * ds_rate)
47
+
48
+ unpadded_features, newline_features = self._get_unpadded_features(
49
+ orig_height, orig_width, patches_height, patches_width, scale_height, scale_width
50
+ )
51
+ # The base patch covers the entire image (+1 for the CLS)
52
+ base_features = patches_height * patches_width + self.num_additional_image_tokens
53
+ num_image_tokens = unpadded_features + newline_features + base_features
54
+ return num_image_tokens