# coding=utf-8 # Copyright 2024 The HuggingFace Inc. team. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. """Processor class for BailingMM2.""" import sys from typing import List, Union, Dict, Optional import torch import PIL from PIL import Image if sys.version_info >= (3, 11): from typing import Unpack else: from typing_extensions import Unpack from transformers.feature_extraction_utils import BatchFeature from transformers.image_utils import ImageInput from transformers.processing_utils import ( ProcessingKwargs, ProcessorMixin, ) from transformers.tokenization_utils_base import PreTokenizedInput, TextInput from bailingmm_utils import process_vision_info, VideoInput, process_ratio, process_reference_vision_info, get_default_image_gen_hw import torchvision import math DEFAULT_IMAGE_PATCH_TOKEN = "" DEFAULT_IM_START_TOKEN = "" DEFAULT_IM_END_TOKEN = "" DEFAULT_VID_START_TOKEN = "" DEFAULT_GEN_IMAGE_PATCH_TOKEN = "" DEFAULT_GEN_IM_START_TOKEN = "" DEFAULT_GEN_IM_END_TOKEN = "" PLACEHOLDER_IMAGE_TOKEN_IN_TEXT = "" DEFAULT_END_OF_CHUNK_TOKEN = "" DEFAULT_FRAME_PATCH_TOKEN = "" DEFAULT_TEXT_TOKEN = '' DEFAULT_ASR_TOKEN = '' DEFAULT_TTS_TOKEN = '' USER_PREFIX = "HUMAN" ASSISTANT_PREFIX = "ASSISTANT" SYSTEM_PROMPT_LINGV2_FLASH_NOTHINK = "SYSTEM你是一个友好的AI助手。\n\ndetailed thinking off" SYSTEM_PROMPT_LINGV2_FLASH_THINK = "SYSTEM你是一个友好的AI助手。\n\ndetailed thinking on" def check_single_quotes(s): count = s.count("'") if count % 2 != 0: return False positions = [i for i, char in enumerate(s) if char == "'"] for i in range(0, len(positions), 2): start = positions[i] end = positions[i+1] substr = s[start+1:end] chinese_count = 0 for char in substr: if '\u4e00' <= char <= '\u9fff': chinese_count += 1 other_count = len(substr) - chinese_count total = 3 * chinese_count + other_count if total >= 20: return False return True def get_text_from_prompt(prompt): if "'" in prompt and check_single_quotes(prompt): prompt = prompt.replace("'", '"') patterns = [r'\"(.*?)\"', r'‘(.*?)’', r'“(.*?)”'] import re texts = [] patterns = [r'\"(.*?)\"', r'‘(.*?)’', r'“(.*?)”'] for pattern in patterns: texts.extend(re.findall(pattern, prompt)) if len(texts) == 1: assert texts[0] in prompt is_remove = False remove_keywords = ["remove", "delete", "erase"] text_start = min([j for j in [prompt.find(i) for i in ['"', '‘', '“']] if j >= 0]) for kw in remove_keywords: if kw in prompt.lower(): if prompt.lower().find(kw) < text_start: is_remove = True break if is_remove: texts = [] text = " ".join(texts[-1:]) if len(text) > 0: text = f'Text "{text}"' text += ". " return text def crop_to_aspect_max(img: Image.Image, target_ratio: float) -> Image.Image: """ Center-crop a PIL.Image to the largest area fitting the target aspect ratio (width/height) without resizing. Uses torchvision CenterCrop. Args: img: PIL.Image.Image input image target_ratio: float target aspect ratio (width/height), must be positive Returns: the center-cropped PIL.Image.Image """ if not isinstance(img, Image.Image): raise TypeError("img must be a PIL.Image.Image") if not math.isfinite(target_ratio) or target_ratio <= 0: raise ValueError("target_ratio must be a positive, finite number") W, H = img.size if W <= 0 or H <= 0: raise ValueError("image size is invalid") orig_ratio = W / H if orig_ratio >= target_ratio: # image is wider than the target: use full height, crop left/right new_h = H new_w = int(math.floor(target_ratio * H)) new_w = max(1, min(new_w, W)) # guard against extreme ratios producing invalid sizes else: # image is narrower than the target: use full width, crop top/bottom new_w = W new_h = int(math.floor(W / target_ratio)) new_h = max(1, min(new_h, H)) crop = torchvision.transforms.CenterCrop((new_h, new_w)) # size is (h, w) return crop(img) def transform_reference_images( images, image_gen_aspect_ratio=None, image_gen_resolution=512, image_gen_input_channels=None, ): if image_gen_input_channels not in (3, 4): raise ValueError( "image_gen_input_channels must be explicitly set to 3 or 4 " "from the checkpoint capability contract" ) image_mode = "RGB" if image_gen_input_channels == 3 else "RGBA" images = [image.convert(image_mode) for image in images] ref_pil = images[0] if image_gen_aspect_ratio is not None: ref_pil = crop_to_aspect_max(ref_pil, image_gen_aspect_ratio) ori_h = ref_pil.size[1] ori_w = ref_pil.size[0] closest_size, _ = process_ratio(ori_h=ori_h, ori_w=ori_w, highres=image_gen_resolution) ref_pils = [torchvision.transforms.functional.resize(i, closest_size, interpolation=torchvision.transforms.InterpolationMode.BILINEAR) for i in images] ref_tensor = torch.cat([ ((torchvision.transforms.functional.to_tensor(i) - 0.5) * 2.0).unsqueeze(0) for i in ref_pils ], dim=0) return ref_tensor, ref_pil.size[1], ref_pil.size[0] class BailingMM2ProcessorKwargs(ProcessingKwargs, total=False): # see processing_utils.ProcessingKwargs documentation for usage. _defaults = { "text_kwargs": {"padding": True, "padding_side": "right"}, "image_kwargs": {}, "video_kwargs": {}, } class BailingMM2Processor(ProcessorMixin): r""" Constructs a BailingMM2 processor which wraps a bailingmm2 image processor, bailing audio processor and a LLaMa tokenizer into a single processor. Args: image_processor ([`BailingMM2ImageProcessor`], *optional*): The image processor is a required input. tokenizer ([`LlamaTokenizerFast`], *optional*): The tokenizer is a required input. chat_template (`str`, *optional*): A Jinja template which will be used to convert lists of messages in a chat into a tokenizable string. image_token (`str`, *optional*, defaults to `""`): Special token used to denote image location. video_token (`str`, *optional*, defaults to `"