XenonFear128 commited on
Commit
aadef61
·
verified ·
1 Parent(s): fdcc140

Upload processing_agnes.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. processing_agnes.py +116 -0
processing_agnes.py ADDED
@@ -0,0 +1,116 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ """Processor for Agnes 3.0 Flash: tokenizer + image processor + video processor."""
15
+
16
+ import numpy as np
17
+
18
+ from transformers.processing_utils import MultiModalData, ProcessingKwargs, ProcessorMixin
19
+ from transformers.utils import auto_docstring, logging
20
+
21
+
22
+ logger = logging.get_logger(__name__)
23
+
24
+
25
+ class AgnesProcessorKwargs(ProcessingKwargs, total=False):
26
+ _defaults = {
27
+ "text_kwargs": {"padding": False, "return_token_type_ids": False, "return_mm_token_type_ids": True},
28
+ "videos_kwargs": {"return_metadata": True},
29
+ }
30
+
31
+
32
+ @auto_docstring
33
+ class AgnesProcessor(ProcessorMixin):
34
+ valid_processor_kwargs = AgnesProcessorKwargs
35
+
36
+ def __init__(self, image_processor=None, tokenizer=None, video_processor=None, chat_template=None, **kwargs):
37
+ self.image_token = getattr(tokenizer, "image_token", "<|image_pad|>")
38
+ self.video_token = getattr(tokenizer, "video_token", "<|video_pad|>")
39
+ self.image_token_id = getattr(tokenizer, "image_token_id", None) or tokenizer.convert_tokens_to_ids(self.image_token)
40
+ self.video_token_id = getattr(tokenizer, "video_token_id", None) or tokenizer.convert_tokens_to_ids(self.video_token)
41
+ super().__init__(image_processor, tokenizer, video_processor, chat_template=chat_template)
42
+ self.vision_start_token = getattr(tokenizer, "vision_start_token", "<|vision_start|>")
43
+ self.vision_end_token = getattr(tokenizer, "vision_end_token", "<|vision_end|>")
44
+ self.vision_start_token_id = getattr(tokenizer, "vision_start_token_id", None) or tokenizer.convert_tokens_to_ids(
45
+ self.vision_start_token
46
+ )
47
+ self.vision_end_token_id = getattr(tokenizer, "vision_end_token_id", None) or tokenizer.convert_tokens_to_ids(
48
+ self.vision_end_token
49
+ )
50
+
51
+ def replace_image_token(self, image_inputs: dict, image_idx: int) -> str:
52
+ per_token = self.image_processor.merge_size**2
53
+ n = image_inputs["image_grid_thw"][image_idx].prod() // per_token
54
+ return self.image_token * n
55
+
56
+ def replace_video_token(self, video_inputs: dict, video_idx: int) -> str:
57
+ per_token = self.video_processor.merge_size**2
58
+ thw = video_inputs["video_grid_thw"][video_idx]
59
+ n_frames = thw[0]
60
+ per_frame = thw[1:].prod() // per_token
61
+ meta = video_inputs["video_metadata"][video_idx]
62
+ if meta.fps is None:
63
+ logger.warning_once(
64
+ "Frame timestamps are needed to build the video prompt but the `fps` of the input video could not be "
65
+ "inferred (no `video_metadata`, pre-sampled frames?). Defaulting to `fps=24`."
66
+ )
67
+ meta.fps = 24 if meta.fps is None else meta.fps
68
+ stamps = self._frame_timestamps(meta.frames_indices, meta.fps, self.video_processor.temporal_patch_size)
69
+ text = ""
70
+ for f in range(n_frames):
71
+ text += f"<{stamps[f]:.1f} seconds>"
72
+ text += self.vision_start_token + self.video_token * per_frame + self.vision_end_token
73
+ return text
74
+
75
+ def _get_num_multimodal_tokens(self, image_sizes=None, video_sizes=None, **kwargs):
76
+ """Placeholder counts for inputs of the given sizes, without running the
77
+ processors on real pixels."""
78
+ data = {}
79
+ if image_sizes is not None:
80
+ ik = AgnesProcessorKwargs._defaults.get("images_kwargs", {})
81
+ ik.update(kwargs)
82
+ merge = ik.get("merge_size", None) or self.image_processor.merge_size
83
+ patches = [self.image_processor.get_number_of_image_patches(*s, ik) for s in image_sizes]
84
+ data.update({"num_image_tokens": [p // merge**2 for p in patches], "num_image_patches": patches})
85
+ if video_sizes is not None:
86
+ vk = AgnesProcessorKwargs._defaults.get("videos_kwargs", {})
87
+ vk.update(kwargs)
88
+ merge = vk.get("merge_size", None) or self.video_processor.merge_size
89
+ patches = [self.video_processor.get_number_of_video_patches(*s, vk) for s in video_sizes]
90
+ data["num_video_tokens"] = [p // merge**2 for p in patches]
91
+ return MultiModalData(**data)
92
+
93
+ def post_process_image_text_to_text(self, generated_outputs, skip_special_tokens=True, clean_up_tokenization_spaces=False, **kwargs):
94
+ """Decode generated ids to text."""
95
+ return self.tokenizer.batch_decode(
96
+ generated_outputs, skip_special_tokens=skip_special_tokens,
97
+ clean_up_tokenization_spaces=clean_up_tokenization_spaces, **kwargs,
98
+ )
99
+
100
+ @property
101
+ def model_input_names(self):
102
+ return super().model_input_names + ["mm_token_type_ids"]
103
+
104
+ @staticmethod
105
+ def _frame_timestamps(indices: list[int] | np.ndarray, video_fps: float, merge_size: int = 2):
106
+ """One timestamp per temporal patch: the mean of the first and last
107
+ frame time inside the patch."""
108
+ if not isinstance(indices, list):
109
+ indices = indices.tolist()
110
+ if len(indices) % merge_size != 0:
111
+ indices.extend(indices[-1] for _ in range(merge_size - len(indices) % merge_size))
112
+ times = [i / video_fps for i in indices]
113
+ return [(times[i] + times[i + merge_size - 1]) / 2 for i in range(0, len(times), merge_size)]
114
+
115
+
116
+ __all__ = ["AgnesProcessor"]