Upload processing_agnes.py with huggingface_hub
Browse files- 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"]
|