XenonFear128 commited on
Commit
c75c717
·
verified ·
1 Parent(s): 33ee44e

Upload image_processing_agnes.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. image_processing_agnes.py +190 -0
image_processing_agnes.py ADDED
@@ -0,0 +1,190 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 Agnes AI. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ """Image processor for Agnes 3.0 Flash: dynamic-resolution patching."""
15
+
16
+ import math
17
+ from collections.abc import Iterable
18
+
19
+ import torch
20
+ from torchvision.transforms.v2 import functional as tvF
21
+
22
+ from transformers.image_processing_backends import TorchvisionBackend
23
+ from transformers.image_processing_utils import BatchFeature
24
+ from transformers.image_transforms import group_images_by_shape, reorder_images
25
+ from transformers.image_utils import ImageInput, PILImageResampling, SizeDict
26
+ from transformers.processing_utils import ImagesKwargs, Unpack
27
+ from transformers.utils import TensorType, auto_docstring
28
+
29
+
30
+ class AgnesImageProcessorKwargs(ImagesKwargs, total=False):
31
+ r"""
32
+ min_pixels (`int`, *optional*, defaults to `256 * 256`):
33
+ Lower bound on the pixel count after resizing.
34
+ max_pixels (`int`, *optional*, defaults to `4096 * 4096`):
35
+ Upper bound on the pixel count after resizing.
36
+ patch_size (`int`, *optional*, defaults to 16):
37
+ Spatial patch size of the vision tower.
38
+ temporal_patch_size (`int`, *optional*, defaults to 2):
39
+ Temporal patch size of the vision tower (images are duplicated to fill it).
40
+ merge_size (`int`, *optional*, defaults to 2):
41
+ Side of the patch square merged into one language-model token.
42
+ """
43
+
44
+ min_pixels: int
45
+ max_pixels: int
46
+ patch_size: int
47
+ temporal_patch_size: int
48
+ merge_size: int
49
+
50
+
51
+ def fit_to_grid(height: int, width: int, factor: int = 32, min_pixels: int = 256 * 256, max_pixels: int = 4096 * 4096):
52
+ """Pick a (height, width) that is a multiple of `factor` on both sides,
53
+ keeps the pixel count inside [min_pixels, max_pixels] and stays as close
54
+ as possible to the original aspect ratio."""
55
+ if max(height, width) / min(height, width) > 200:
56
+ raise ValueError(f"absolute aspect ratio must be smaller than 200, got {max(height, width) / min(height, width)}")
57
+ h = round(height / factor) * factor
58
+ w = round(width / factor) * factor
59
+ if h * w > max_pixels:
60
+ scale = math.sqrt((height * width) / max_pixels)
61
+ h = max(factor, math.floor(height / scale / factor) * factor)
62
+ w = max(factor, math.floor(width / scale / factor) * factor)
63
+ elif h * w < min_pixels:
64
+ scale = math.sqrt(min_pixels / (height * width))
65
+ h = math.ceil(height * scale / factor) * factor
66
+ w = math.ceil(width * scale / factor) * factor
67
+ return h, w
68
+
69
+
70
+ @auto_docstring
71
+ class AgnesImageProcessor(TorchvisionBackend):
72
+ do_resize = True
73
+ resample = PILImageResampling.BICUBIC
74
+ size = {"shortest_edge": 256 * 256, "longest_edge": 4096 * 4096}
75
+ default_to_square = False
76
+ do_rescale = True
77
+ do_normalize = True
78
+ image_mean = [0.5, 0.5, 0.5]
79
+ image_std = [0.5, 0.5, 0.5]
80
+ do_convert_rgb = True
81
+ patch_size = 16
82
+ temporal_patch_size = 2
83
+ merge_size = 2
84
+ valid_kwargs = AgnesImageProcessorKwargs
85
+ model_input_names = ["pixel_values", "image_grid_thw"]
86
+
87
+ def __init__(self, **kwargs: Unpack[AgnesImageProcessorKwargs]):
88
+ size = kwargs.pop("size", None)
89
+ min_pixels = kwargs.pop("min_pixels", None)
90
+ max_pixels = kwargs.pop("max_pixels", None)
91
+ size = self.size if size is None else size
92
+ # min_pixels / max_pixels are the older spelling of the two size keys
93
+ if min_pixels is not None:
94
+ size["shortest_edge"] = min_pixels
95
+ size.pop("min_pixels", None)
96
+ if max_pixels is not None:
97
+ size["longest_edge"] = max_pixels
98
+ size.pop("max_pixels", None)
99
+ if "shortest_edge" not in size or "longest_edge" not in size:
100
+ raise ValueError("size must contain 'shortest_edge' and 'longest_edge' keys.")
101
+ super().__init__(size=size, **kwargs)
102
+
103
+ def _standardize_kwargs(
104
+ self,
105
+ size: int | Iterable[int] | dict[str, int] | SizeDict | None = None,
106
+ min_pixels: int | None = None,
107
+ max_pixels: int | None = None,
108
+ **kwargs,
109
+ ) -> dict:
110
+ if min_pixels is not None and max_pixels is not None:
111
+ size = SizeDict(shortest_edge=min_pixels, longest_edge=max_pixels)
112
+ kwargs = super()._standardize_kwargs(size=size, **kwargs)
113
+ size = kwargs.get("size", self.size)
114
+ if not size.shortest_edge or not size.longest_edge:
115
+ raise ValueError("size must contain 'shortest_edge' and 'longest_edge' keys.")
116
+ return kwargs
117
+
118
+ @auto_docstring
119
+ def preprocess(self, images: ImageInput, **kwargs: Unpack[AgnesImageProcessorKwargs]) -> BatchFeature:
120
+ return super().preprocess(images, **kwargs)
121
+
122
+ def _preprocess(
123
+ self,
124
+ images: list["torch.Tensor"],
125
+ do_resize: bool,
126
+ size: SizeDict,
127
+ resample: "PILImageResampling | tvF.InterpolationMode | int | None",
128
+ do_rescale: bool,
129
+ rescale_factor: float,
130
+ do_normalize: bool,
131
+ image_mean: float | list[float] | None,
132
+ image_std: float | list[float] | None,
133
+ patch_size: int,
134
+ temporal_patch_size: int,
135
+ merge_size: int,
136
+ disable_grouping: bool | None,
137
+ return_tensors: str | TensorType | None,
138
+ **kwargs,
139
+ ) -> BatchFeature:
140
+ # 1. resize, batched per input shape
141
+ by_shape, order = group_images_by_shape(images, disable_grouping=disable_grouping)
142
+ resized = {}
143
+ for shape, batch in by_shape.items():
144
+ height, width = batch.shape[-2:]
145
+ if do_resize:
146
+ new_h, new_w = fit_to_grid(
147
+ height, width, factor=patch_size * merge_size,
148
+ min_pixels=size.shortest_edge, max_pixels=size.longest_edge,
149
+ )
150
+ batch = self.resize(image=batch, size=SizeDict(height=new_h, width=new_w), resample=resample)
151
+ resized[shape] = batch
152
+ images = reorder_images(resized, order)
153
+
154
+ # 2. normalise and cut into merge-ordered patches, batched per resized shape
155
+ by_shape, order = group_images_by_shape(images, disable_grouping=disable_grouping)
156
+ flat = {}
157
+ grids = {}
158
+ for shape, batch in by_shape.items():
159
+ new_h, new_w = batch.shape[-2:]
160
+ px = self.rescale_and_normalize(batch, do_rescale, rescale_factor, do_normalize, image_mean, image_std)
161
+ n, c = px.shape[:2]
162
+ gh, gw = new_h // patch_size, new_w // patch_size
163
+ px = px.reshape(n, c, gh // merge_size, merge_size, patch_size, gw // merge_size, merge_size, patch_size)
164
+ # -> [n, gh/merge, gw/merge, merge, merge, c, patch, patch]: patches of one
165
+ # merge square end up adjacent in the flattened sequence
166
+ px = px.permute(0, 2, 5, 3, 6, 1, 4, 7)
167
+ px = (
168
+ px.unsqueeze(6)
169
+ .expand(-1, -1, -1, -1, -1, -1, temporal_patch_size, -1, -1)
170
+ .reshape(n, gh * gw, c * temporal_patch_size * patch_size * patch_size)
171
+ )
172
+ flat[shape] = px
173
+ grids[shape] = [[1, gh, gw]] * n
174
+
175
+ pixel_values = torch.cat(reorder_images(flat, order), dim=0)
176
+ image_grid_thw = torch.tensor(reorder_images(grids, order), dtype=torch.long)
177
+ return BatchFeature(data={"pixel_values": pixel_values, "image_grid_thw": image_grid_thw}, tensor_type=return_tensors)
178
+
179
+ def get_number_of_image_patches(self, height: int, width: int, images_kwargs=None):
180
+ """Number of vision patches an image of this size produces (used by
181
+ serving engines to lay out placeholders without running the processor)."""
182
+ min_pixels = images_kwargs["min_pixels"] if "min_pixels" in images_kwargs else self.size["shortest_edge"]
183
+ max_pixels = images_kwargs["max_pixels"] if "max_pixels" in images_kwargs else self.size["longest_edge"]
184
+ patch_size = images_kwargs.get("patch_size", self.patch_size)
185
+ merge_size = images_kwargs.get("merge_size", self.merge_size)
186
+ new_h, new_w = fit_to_grid(height, width, patch_size * merge_size, min_pixels=min_pixels, max_pixels=max_pixels)
187
+ return (new_h // patch_size) * (new_w // patch_size)
188
+
189
+
190
+ __all__ = ["AgnesImageProcessor"]