Team Ai
Modelpublic

optimum-intel-internal-testing/tiny-random-phi-4-multimodal

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes16kdownloads
processing_phi4mm.py734 linesDownload Raw Back to root
1# Copyright 2024 Microsoft and the HuggingFace Inc. team. 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 at6#7#     http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# 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 and13# limitations under the License.14 15"""16Processor class for Phi4MM17"""18import re19from typing import List, Optional, Tuple, Union20import math21from enum import Enum22 23import numpy as np24import scipy25import torch26import torchvision27 28from transformers import AutoFeatureExtractor, AutoImageProcessor29from transformers.feature_extraction_sequence_utils import SequenceFeatureExtractor30from transformers.image_processing_utils import BaseImageProcessor, BatchFeature31from transformers.image_utils import (32    ImageInput,33    make_list_of_images,34    valid_images,35)36from transformers.processing_utils import ProcessorMixin37from transformers.tokenization_utils_base import PaddingStrategy, TextInput, TruncationStrategy38from transformers.utils import TensorType, logging39from torch.nn.utils.rnn import pad_sequence40 41 42logger = logging.get_logger(__name__)43 44# Special tokens45_COMPATIBLE_IMAGE_SPECIAL_TOKEN_PATTERN = r'<\|image_\d+\|>'  # For backward compatibility46_COMPATIBLE_AUDIO_SPECIAL_TOKEN_PATTERN = r'<\|audio_\d+\|>'  # For backward compatibility47_IMAGE_SPECIAL_TOKEN = '<|endoftext10|>'48_AUDIO_SPECIAL_TOKEN = '<|endoftext11|>'49_IMAGE_SPECIAL_TOKEN_ID = 200010  # '<|endoftext10|>', or we can better name it (in `tokenizer_config.json`)50_AUDIO_SPECIAL_TOKEN_ID = 200011  # '<|endoftext11|>'51 52 53class InputMode(Enum):54    LANGUAGE = 055    VISION = 156    SPEECH = 257    VISION_SPEECH = 358 59 60class Phi4MMImageProcessor(BaseImageProcessor):61    r"""62    Constructs a Phi4MM image processor.63    """64    model_input_names = ["input_image_embeds", "image_sizes", "image_attention_mask"]65 66    def __init__(67        self,68        dynamic_hd,69        **kwargs,70    ) -> None:71        super().__init__(**kwargs)72        self.dynamic_hd = dynamic_hd73 74    def find_closest_aspect_ratio(self, aspect_ratio, target_ratios, width, height, image_size):75        best_ratio_diff = float('inf')76        best_ratio = (1, 1)77        area = width * height78        for ratio in target_ratios:79            target_aspect_ratio = ratio[0] / ratio[1]80            ratio_diff = abs(aspect_ratio - target_aspect_ratio)81            if ratio_diff < best_ratio_diff:82                best_ratio_diff = ratio_diff83                best_ratio = ratio84            elif ratio_diff == best_ratio_diff:85                if area > 0.5 * image_size * image_size * ratio[0] * ratio[1]:86                    best_ratio = ratio87        return best_ratio88 89    def dynamic_preprocess(self, image, min_num=1, max_num=12, image_size=384, mask_size=27, use_thumbnail=True):90        orig_width, orig_height = image.size91 92        w_crop_num = math.ceil(orig_width/float(image_size))93        h_crop_num = math.ceil(orig_height/float(image_size))94        if w_crop_num * h_crop_num > max_num:95 96            aspect_ratio = orig_width / orig_height97 98            # calculate the existing image aspect ratio99            target_ratios = set(100                (i, j) for n in range(min_num, max_num + 1) for i in range(1, n + 1) for j in range(1, n + 1) if101                i * j <= max_num and i * j >= min_num)102            target_ratios = sorted(target_ratios, key=lambda x: x[0] * x[1])103 104            # find the closest aspect ratio to the target105            target_aspect_ratio = self.find_closest_aspect_ratio(106                aspect_ratio, target_ratios, orig_width, orig_height, image_size)107 108            # calculate the target width and height109            target_width = image_size * target_aspect_ratio[0]110            target_height = image_size * target_aspect_ratio[1]111        else:112            target_width = image_size * w_crop_num113            target_height = image_size * h_crop_num114            target_aspect_ratio = (w_crop_num, h_crop_num)115 116        # Calculate the ratio117        ratio_width = target_width / orig_width118        ratio_height = target_height / orig_height119        if ratio_width < ratio_height:120            new_size = (target_width, int(orig_height * ratio_width))121            padding_width = 0122            padding_height = target_height - int(orig_height * ratio_width)123        else:124            new_size = (int(orig_width * ratio_height), target_height)125            padding_width = target_width - int(orig_width * ratio_height)126            padding_height = 0127 128        attention_mask = torch.ones((int(mask_size*target_aspect_ratio[1]), int(mask_size*target_aspect_ratio[0])))129        if padding_width >= 14:130            attention_mask[:, -math.floor(padding_width/14):] = 0131        if padding_height >= 14:132            attention_mask[-math.floor(padding_height/14):,:] = 0133        assert attention_mask.sum() > 0134 135        if min(new_size[1], target_height) < 10 or min(new_size[0], target_width) < 10:136            raise ValueError(f'the aspect ratio is very extreme {new_size}')137 138        image = torchvision.transforms.functional.resize(image, [new_size[1], new_size[0]],)139 140        resized_img = torchvision.transforms.functional.pad(image, [0, 0, padding_width, padding_height], fill=[255,255,255])141 142        return resized_img, attention_mask143 144    def pad_to_max_num_crops(self, images, max_crops=5):145        """146        images: B x 3 x H x W, B<=max_crops147        """148        B, _, H, W = images.shape149        if B < max_crops:150            pad = torch.zeros(max_crops - B, 3, H, W, dtype=images.dtype, device=images.device)151            images = torch.cat([images, pad], dim=0)152        return images153 154    def pad_mask_to_max_num_crops(self, masks, max_crops=5):155        B, H, W = masks.shape156        if B < max_crops:157            pad = torch.ones(max_crops - B, H, W, dtype=masks.dtype, device=masks.device)158            masks = torch.cat([masks, pad], dim=0)159        return masks160 161    def preprocess(162        self,163        images: ImageInput,164        return_tensors: Optional[Union[str, TensorType]] = None,165    ):166        """167        Args:168            images (`ImageInput`):169                Image to preprocess. Expects a single or batch of images with pixel values ranging from 0 to 255. If170                passing in images with pixel values between 0 and 1, set `do_rescale=False`.171            return_tensors (`str` or `TensorType`, *optional*):172                The type of tensors to return. Can be one of:173                - Unset: Return a list of `np.ndarray`.174                - `TensorType.TENSORFLOW` or `'tf'`: Return a batch of type `tf.Tensor`.175                - `TensorType.PYTORCH` or `'pt'`: Return a batch of type `torch.Tensor`.176                - `TensorType.NUMPY` or `'np'`: Return a batch of type `np.ndarray`.177                - `TensorType.JAX` or `'jax'`: Return a batch of type `jax.numpy.ndarray`.178        """179        images = make_list_of_images(images)180 181        if not valid_images(images):182            raise ValueError(183                "Invalid image type. Must be of type PIL.Image.Image, numpy.ndarray, "184                "torch.Tensor, tf.Tensor or jax.ndarray."185            )186 187        # Basic settings.188        img_processor = torchvision.transforms.Compose([189            torchvision.transforms.ToTensor(),190            torchvision.transforms.Normalize(191                (0.5, 0.5, 0.5),192                (0.5, 0.5, 0.5)193            ),194        ])195        dyhd_base_resolution = 448196 197        # Dynamic HD198        base_resolution = dyhd_base_resolution199        images = [image.convert('RGB') for image in images]200        # cover 384 and 448 resolution201        mask_resolution = base_resolution // 14202        elems, image_attention_masks = [], []203        for im in images:204            elem, attention_mask = self.dynamic_preprocess(im, max_num=self.dynamic_hd, image_size=base_resolution, mask_size=mask_resolution)205            elems.append(elem)206            image_attention_masks.append(attention_mask)207        hd_images = [img_processor(im) for im in elems]208        global_image = [torch.nn.functional.interpolate(im.unsqueeze(0).float(), size=(base_resolution, base_resolution), mode='bicubic',).to(im.dtype) for im in hd_images]209        shapes = [[im.size(1), im.size(2)] for im in hd_images]210        mask_shapes = [[mask.size(0), mask.size(1)] for mask in image_attention_masks]211        global_attention_mask = [torch.ones((1, mask_resolution, mask_resolution)) for _ in hd_images]212        hd_images_reshape = [im.reshape(1, 3,213                                            h//base_resolution,214                                            base_resolution,215                                            w//base_resolution,216                                            base_resolution217                                            ).permute(0,2,4,1,3,5).reshape(-1, 3, base_resolution, base_resolution).contiguous() for im, (h, w) in zip(hd_images, shapes)]218        attention_masks_reshape = [mask.reshape(1,219                                            h//mask_resolution,220                                            mask_resolution,221                                            w//mask_resolution,222                                            mask_resolution223                                            ).permute(0,1,3,2,4).reshape(-1, mask_resolution, mask_resolution).contiguous() for mask, (h, w) in zip(image_attention_masks, mask_shapes)]224        downsample_attention_masks = [mask[:,0::2,0::2].reshape(1,225                                            h//mask_resolution,226                                            w//mask_resolution,227                                            mask_resolution//2+mask_resolution%2,228                                            mask_resolution//2+mask_resolution%2229                                            ).permute(0,1,3,2,4) for mask, (h,w) in zip(attention_masks_reshape, mask_shapes)]230        downsample_attention_masks = [mask.reshape(mask.size(1)*mask.size(2), mask.size(3)*mask.size(4))for mask in downsample_attention_masks]231        num_img_tokens = [256 + 1 + int(mask.sum().item()) + int(mask[:,0].sum().item()) + 16 for mask in downsample_attention_masks]232 233        hd_images_reshape = [torch.cat([_global_image] + [_im], dim=0) for _global_image, _im in zip(global_image, hd_images_reshape)]234        hd_masks_reshape = [torch.cat([_global_mask] + [_mask], dim=0) for _global_mask, _mask in zip(global_attention_mask, attention_masks_reshape)]235        max_crops = max([img.size(0) for img in hd_images_reshape])236        image_transformed = [self.pad_to_max_num_crops(im, max_crops) for im in hd_images_reshape]237        image_transformed = torch.stack(image_transformed, dim=0)238        mask_transformed = [self.pad_mask_to_max_num_crops(mask, max_crops) for mask in hd_masks_reshape]239        mask_transformed = torch.stack(mask_transformed, dim=0)240 241        returned_input_image_embeds = image_transformed242        returned_image_sizes = torch.tensor(shapes, dtype=torch.long)243        returned_image_attention_mask = mask_transformed244        returned_num_img_tokens = num_img_tokens245 246        data = {247            "input_image_embeds": returned_input_image_embeds,248            "image_sizes": returned_image_sizes,249            "image_attention_mask": returned_image_attention_mask,250            "num_img_tokens": returned_num_img_tokens,251        }252 253        return BatchFeature(data=data, tensor_type=return_tensors)254 255 256AudioInput = Tuple[Union[np.ndarray, torch.Tensor], int]257AudioInputs = List[AudioInput]258 259 260def speechlib_mel(sample_rate, n_fft, n_mels, fmin=None, fmax=None):261    """Create a Mel filter-bank the same as SpeechLib FbankFC.262 263    Args:264        sample_rate (int): Sample rate in Hz. number > 0 [scalar]265        n_fft (int): FFT size. int > 0 [scalar]266        n_mel (int): Mel filter size. int > 0 [scalar]267        fmin (float): lowest frequency (in Hz). If None use 0.0.268            float >= 0 [scalar]269        fmax: highest frequency (in Hz). If None use sample_rate / 2.270            float >= 0 [scalar]271 272    Returns273        out (numpy.ndarray): Mel transform matrix274            [shape=(n_mels, 1 + n_fft/2)]275    """276 277    bank_width = int(n_fft // 2 + 1)278    if fmax is None:279        fmax = sample_rate / 2280    if fmin is None:281        fmin = 0282    assert fmin >= 0, "fmin cannot be negtive"283    assert fmin < fmax <= sample_rate / 2, "fmax must be between (fmin, samplerate / 2]"284 285    def mel(f):286        return 1127.0 * np.log(1.0 + f / 700.0)287 288    def bin2mel(fft_bin):289        return 1127.0 * np.log(1.0 + fft_bin * sample_rate / (n_fft * 700.0))290 291    def f2bin(f):292        return int((f * n_fft / sample_rate) + 0.5)293 294    # Spec 1: FFT bin range [f2bin(fmin) + 1, f2bin(fmax) - 1]295    klo = f2bin(fmin) + 1296    khi = f2bin(fmax)297 298    khi = max(khi, klo)299 300    # Spec 2: SpeechLib uses trianges in Mel space301    mlo = mel(fmin)302    mhi = mel(fmax)303    m_centers = np.linspace(mlo, mhi, n_mels + 2)304    ms = (mhi - mlo) / (n_mels + 1)305 306    matrix = np.zeros((n_mels, bank_width), dtype=np.float32)307    for m in range(0, n_mels):308        left = m_centers[m]309        center = m_centers[m + 1]310        right = m_centers[m + 2]311        for fft_bin in range(klo, khi):312            mbin = bin2mel(fft_bin)313            if left < mbin < right:314                matrix[m, fft_bin] = 1.0 - abs(center - mbin) / ms315 316    return matrix317 318 319class Phi4MMAudioFeatureExtractor(SequenceFeatureExtractor):320    model_input_names = ["input_audio_embeds", "audio_embed_sizes", "audio_attention_mask"]321 322    def __init__(self, audio_compression_rate, audio_downsample_rate, audio_feat_stride, **kwargs):323        feature_size = 80324        sampling_rate = 16000325        padding_value = 0.0326        super().__init__(feature_size, sampling_rate, padding_value, **kwargs)327 328        self.compression_rate = audio_compression_rate329        self.qformer_compression_rate = audio_downsample_rate330        self.feat_stride = audio_feat_stride331 332        self._eightk_method = "fillzero"333        self._mel = speechlib_mel(16000, 512, 80, fmin=None, fmax=7690).T334 335        self._hamming400 = np.hamming(400)  # for 16k audio336        self._hamming200 = np.hamming(200)  # for 8k audio337 338    def duration_to_frames(self, duration):339        """duration in s, estimated frames"""340        frame_rate = 10341 342        num_frames = duration * 1000 // frame_rate343        return num_frames344 345    def __call__(346        self,347        audios: List[AudioInput],348        return_tensors: Optional[Union[str, TensorType]] = None,349    ):350        # Ref: https://github.com/huggingface/transformers/blob/v4.47.0/src/transformers/models/audio_spectrogram_transformer/feature_extraction_audio_spectrogram_transformer.py#L161351        returned_input_audio_embeds = []352        returned_audio_embed_sizes = []353        audio_frames_list = []354 355        for audio_data, sample_rate in audios:356            audio_embeds = self._extract_features(audio_data, sample_rate)357            audio_frames = len(audio_embeds) * self.feat_stride358            audio_embed_size = self._compute_audio_embed_size(audio_frames)359 360            returned_input_audio_embeds.append(torch.tensor(audio_embeds))361            returned_audio_embed_sizes.append(torch.tensor(audio_embed_size).long())362            audio_frames_list.append(audio_frames)363 364        returned_input_audio_embeds = pad_sequence(365            returned_input_audio_embeds, batch_first=True366        )367        returned_audio_embed_sizes = torch.stack(returned_audio_embed_sizes, dim=0)368        audio_frames = torch.tensor(audio_frames_list)369        returned_audio_attention_mask = torch.arange(0, audio_frames.max()).unsqueeze(0) < audio_frames.unsqueeze(1) if len(audios) > 1 else None370 371        data = {372            "input_audio_embeds": returned_input_audio_embeds,373            "audio_embed_sizes": returned_audio_embed_sizes,374        }375        if returned_audio_attention_mask is not None:376            data["audio_attention_mask"] = returned_audio_attention_mask377 378        return BatchFeature(data=data, tensor_type=return_tensors)379 380    def _extract_spectrogram(self, wav, fs):381        """Extract spectrogram features from waveform.382        Args:383            wav (1D array): waveform of the input384            fs (int): sampling rate of the waveform, 16000 or 8000.385                If fs=8000, the waveform will be resampled to 16000Hz.386        Output:387            log_fbank (2D array): a TxD matrix of log Mel filterbank features.388                D=80, and T is the number of frames.389        """390        if wav.ndim > 1:391            wav = np.squeeze(wav)392 393        # by default, we extract the mean if stereo394        if len(wav.shape) == 2:395            wav = wav.mean(1)396 397        # Resample to 16000 or 8000 if needed398        if fs > 16000:399            wav = scipy.signal.resample_poly(wav, 1, fs // 16000)400            fs = 16000401        elif 8000 < fs < 16000:402            wav = scipy.signal.resample_poly(wav, 1, fs // 8000)403            fs = 8000404        elif fs < 8000:405            raise RuntimeError(f"Unsupported sample rate {fs}")406 407        if fs == 8000:408            if self._eightk_method == "resample":409                # Input audio is 8 kHz. Convert to 16 kHz before feature410                # extraction411                wav = scipy.signal.resample_poly(wav, 2, 1)412                fs = 16000413            # Do nothing here for fillzero method414        elif fs != 16000:415            # Input audio is not a supported sample rate.416            raise RuntimeError(f"Input data using an unsupported sample rate: {fs}")417 418        preemphasis = 0.97419 420        if fs == 8000:421            n_fft = 256422            win_length = 200423            hop_length = 80424            fft_window = self._hamming200425        elif fs == 16000:426            n_fft = 512427            win_length = 400428            hop_length = 160429            fft_window = self._hamming400430 431        # Spec 1: SpeechLib cut remaining sample insufficient for a hop432        n_batch = (wav.shape[0] - win_length) // hop_length + 1433        # Here we don't use stride_tricks since the input array may not satisfy434        # memory layout requirement and we need writeable output435        # Here we only use list of views before copy to desination436        # so it is more efficient than broadcasting437        y_frames = np.array(438            [wav[_stride : _stride + win_length] for _stride in range(0, hop_length * n_batch, hop_length)],439            dtype=np.float32,440        )441 442        # Spec 2: SpeechLib applies preemphasis within each batch443        y_frames_prev = np.roll(y_frames, 1, axis=1)444        y_frames_prev[:, 0] = y_frames_prev[:, 1]445        y_frames = (y_frames - preemphasis * y_frames_prev) * 32768446 447        S = np.fft.rfft(fft_window * y_frames, n=n_fft, axis=1).astype(np.complex64)448 449        if fs == 8000:450            # Need to pad the output to look like 16 kHz data but with zeros in451            # the 4 to 8 kHz bins.452            frames, bins = S.shape453            padarray = np.zeros((frames, bins))454            S = np.concatenate((S[:, 0:-1], padarray), axis=1)  # Nyquist bin gets set to zero455 456        spec = np.abs(S).astype(np.float32)457        return spec458 459    def _extract_features(self, wav, fs):460        """Extract log filterbank features from waveform.461        Args:462            wav (1D array): waveform of the input463            fs (int): sampling rate of the waveform, 16000 or 8000.464                If fs=8000, the waveform will be resampled to 16000Hz.465        Output:466            log_fbank (2D array): a TxD matrix of log Mel filterbank features.467                D=80, and T is the number of frames.468        """469        spec = self._extract_spectrogram(wav, fs)470        spec_power = spec**2471 472        fbank_power = np.clip(spec_power.dot(self._mel), 1.0, None)473        log_fbank = np.log(fbank_power).astype(np.float32)474 475        return log_fbank476 477    def _compute_audio_embed_size(self, audio_frames):478        integer = audio_frames // self.compression_rate479        remainder = audio_frames % self.compression_rate480 481        result = integer if remainder == 0 else integer + 1482 483        integer = result // self.qformer_compression_rate484        remainder = result % self.qformer_compression_rate485        result = integer if remainder == 0 else integer + 1  # qformer compression486 487        return result488 489 490class Phi4MMProcessor(ProcessorMixin):491    r"""492    Constructs a Phi4MM processor which raps an image processor, a audio processor, and a GPT tokenizer into a single processor.493 494    [`Phi4MMProcessor`] offers all the functionalities of [`Phi4MMImageProcessor`] and [`GPT2Tokenizer`]. See the495    [`~Phi4MMProcessor.__call__`] and [`~Phi4MMProcessor.decode`] for more information.496 497    Args:498        image_processor ([`Phi4MMImageProcessor`], *optional*):499            The image processor is a required input.500        tokenizer ([`GPT2Tokenizer`], *optional*):501            The tokenizer is a required input.502    """503 504    attributes = ["image_processor", "audio_processor", "tokenizer"]505    tokenizer_class = "GPT2TokenizerFast"506    image_processor_class = "AutoImageProcessor"  # Phi4MMImageProcessor will be registered later507    audio_processor_class = "AutoFeatureExtractor"  # Phi4MMAudioFeatureExtractor will be registered later508 509    def __init__(self, image_processor, audio_processor, tokenizer, chat_template=None, **kwargs):510        self.image_processor = image_processor511        self.audio_processor = audio_processor512        self.tokenizer = tokenizer513 514    def __call__(515        self,516        text: Union[TextInput, List[TextInput]],517        images: Optional[ImageInput] = None,518        audios: Optional[AudioInputs] = None,519        padding: Union[bool, str, PaddingStrategy] = False,520        truncation: Optional[Union[bool, str, TruncationStrategy]] = None,521        max_length=None,522        return_tensors: Optional[Union[str, TensorType]] = TensorType.PYTORCH,523    ) -> BatchFeature:524        """525        Main method to prepare for the model one or several sequences(s) and image(s). This method forards the `text`526        and `kwargs` arguments to GPT2Tokenizer's [`~GPT2Tokenizer.__call__`] if `text` is not `None` to encode527        the text. To prepare the image(s), this method forwards the `images` and `kwrags` arguments to528        Phi4MMImageProcessor's [`~Phi4MMImageProcessor.__call__`] if `images` is not `None`. Please refer to the doctsring529        of the above two methods for more information.530 531        Args:532            text (`str`, `List[str]`, `List[List[str]]`):533                The sequence or batch of sequences to be encoded. Each sequence can be a string or a list of strings534                (pretokenized string). If the sequences are provided as list of strings (pretokenized), you must set535                `is_split_into_words=True` (to lift the ambiguity with a batch of sequences).536            images (`PIL.Image.Image`, `np.ndarray`, `torch.Tensor`, `List[PIL.Image.Image]`, `List[np.ndarray]`, `List[torch.Tensor]`):537                The image or batch of images to be prepared. Each image can be a PIL image, NumPy array or PyTorch538                tensor. Both channels-first and channels-last formats are supported.539            padding (`bool`, `str` or [`~utils.PaddingStrategy`], *optional*, defaults to `False`):540                Select a strategy to pad the returned sequences (according to the model's padding side and padding541                index) among:542                - `True` or `'longest'`: Pad to the longest sequence in the batch (or no padding if only a single543                  sequence if provided).544                - `'max_length'`: Pad to a maximum length specified with the argument `max_length` or to the maximum545                  acceptable input length for the model if that argument is not provided.546                - `False` or `'do_not_pad'` (default): No padding (i.e., can output a batch with sequences of different547                  lengths).548            max_length (`int`, *optional*):549                Maximum length of the returned list and optionally padding length (see above).550            truncation (`bool`, *optional*):551                Activates truncation to cut input sequences longer than `max_length` to `max_length`.552            return_tensors (`str` or [`~utils.TensorType`], *optional*):553                If set, will return tensors of a particular framework. Acceptable values are:554 555                - `'tf'`: Return TensorFlow `tf.constant` objects.556                - `'pt'`: Return PyTorch `torch.Tensor` objects.557                - `'np'`: Return NumPy `np.ndarray` objects.558                - `'jax'`: Return JAX `jnp.ndarray` objects.559 560        Returns:561            [`BatchFeature`]: A [`BatchFeature`] with the following fields:562 563            - **input_ids** -- List of token ids to be fed to a model.564            - **input_image_embeds** -- Pixel values to be fed to a model.565            - **image_sizes** -- List of tuples specifying the size of each image in `input_image_embeds`.566            - **image_attention_mask** -- List of attention masks for each image in `input_image_embeds`.567            - **input_audio_embeds** -- Audio embeddings to be fed to a model.568            - **audio_embed_sizes** -- List of integers specifying the size of each audio in `input_audio_embeds`.569            - **attention_mask** -- List of indices specifying which tokens should be attended to by the model.570        """571        image_inputs = self.image_processor(images, return_tensors=return_tensors) if images is not None else {}572        audio_inputs = self.audio_processor(audios, return_tensors=return_tensors) if audios is not None else {}573        inputs = self._convert_images_audios_text_to_inputs(574            image_inputs,575            audio_inputs,576            text,577            padding=padding,578            truncation=truncation,579            max_length=max_length,580            return_tensors=return_tensors,581        )582 583        # idenfity the input mode584        if len(image_inputs) > 0 and len(audio_inputs) > 0:585            input_mode = InputMode.VISION_SPEECH586        elif len(image_inputs) > 0:587            input_mode = InputMode.VISION588        elif len(audio_inputs) > 0:589            input_mode = InputMode.SPEECH590        else:591            input_mode = InputMode.LANGUAGE592        inputs["input_mode"] = torch.tensor([input_mode.value], dtype=torch.long)593 594        return inputs595 596    @property597    def special_image_token_id(self):598        return self.tokenizer.convert_tokens_to_ids(self.special_image_token)599 600    def get_special_image_token_id(self):601        return self.tokenizer.convert_tokens_to_ids(self.special_image_token)602 603    @property604    def chat_template(self):605        return self.tokenizer.chat_template606 607    def _convert_images_audios_text_to_inputs(608        self, images, audios, text, padding=False, truncation=None, max_length=None, return_tensors=None609    ):610        # prepare image id to image input ids611        if len(images) > 0:612            input_image_embeds = images["input_image_embeds"]613            image_sizes = images["image_sizes"]614            image_attention_mask = images["image_attention_mask"]615            num_img_tokens = images['num_img_tokens']616        else:617            input_image_embeds = torch.tensor([])618            image_sizes = torch.tensor([])619            image_attention_mask = torch.tensor([])620            num_img_tokens = []621 622        # prepare audio id to audio input ids623        if len(audios) > 0:624            input_audio_embeds = audios["input_audio_embeds"]625            audio_embed_sizes = audios["audio_embed_sizes"]626            audio_attention_mask = audios.get("audio_attention_mask", None)627        else:628            input_audio_embeds = torch.tensor([])629            audio_embed_sizes = torch.tensor([])630            audio_attention_mask = None631 632        # Replace certain special tokens for compatibility633        # Ref: https://stackoverflow.com/questions/11475885/python-replace-regex634        if isinstance(text, str):635            text = [text]636        assert isinstance(text, list)637        processed_text = [re.sub(_COMPATIBLE_IMAGE_SPECIAL_TOKEN_PATTERN, _IMAGE_SPECIAL_TOKEN, t) for t in text]638        processed_text = [re.sub(_COMPATIBLE_AUDIO_SPECIAL_TOKEN_PATTERN, _AUDIO_SPECIAL_TOKEN, t) for t in processed_text]639 640        input_ids_list = [self.tokenizer(t).input_ids for t in processed_text]641 642        img_cnt, audio_cnt = 0, 0  # only needed for later assertion643        image_token_count_iter = iter(num_img_tokens)644        audio_embed_size_iter = iter(audio_embed_sizes.tolist())645        new_input_ids_list = []646        for input_ids in input_ids_list:647            i = 0648            while i < len(input_ids):649                token_id = input_ids[i]650                if token_id == _AUDIO_SPECIAL_TOKEN_ID:651                    token_count = next(audio_embed_size_iter)652                    audio_cnt += 1653                elif token_id == _IMAGE_SPECIAL_TOKEN_ID:654                    token_count = next(image_token_count_iter)655                    img_cnt += 1656                else:657                    i += 1658                    continue659                tokens = [token_id] * token_count660                input_ids = input_ids[:i] + tokens + input_ids[i + 1:]661                i += token_count662            input_ids = torch.tensor(input_ids, dtype=torch.long)663            new_input_ids_list.append(input_ids)664        lengths = torch.tensor([len(input_ids) for input_ids in new_input_ids_list])665        max_len = lengths.max()666        input_ids = input_ids.new_full((len(new_input_ids_list), max_len), self.tokenizer.pad_token_id)667        # batched inference requires left padding668        for i in range(len(new_input_ids_list)):669            input_ids[i, max_len - len(new_input_ids_list[i]):] = new_input_ids_list[i]670 671        # If the below assertion fails, it might be that input pure-text672        # messages contain image/audio special tokens literally673        # (<|endoftext10|>, <|endoftext11|>).674        assert (675            img_cnt == len(num_img_tokens)676        ), (677            f"Number of image tokens in prompt_token_ids ({img_cnt}) "678            f"does not match number of images ({len(num_img_tokens)})"679        )680        assert (681            audio_cnt == len(audio_embed_sizes)682        ), (683            f"Number of audio tokens in prompt_token_ids ({audio_cnt}) "684            f"does not match number of audios ({len(audio_embed_sizes)})"685        )686 687        # prepare attention mask688        seq_range = torch.arange(max_len - 1, -1, -1)689        attention_mask = seq_range.unsqueeze(0) < lengths.unsqueeze(1)690 691        # prepare batch feature692        data = {693            "input_ids": input_ids,694            "input_image_embeds": input_image_embeds,695            "image_sizes": image_sizes,696            "image_attention_mask": image_attention_mask,697            "input_audio_embeds": input_audio_embeds,698            "audio_embed_sizes": audio_embed_sizes,699            "audio_attention_mask": audio_attention_mask,700            "attention_mask": attention_mask,701        }702 703        return BatchFeature(704            data=data705        )706 707    # Copied from transformers.models.clip.processing_clip.CLIPProcessor.batch_decode with CLIP->Llama708    def batch_decode(self, *args, **kwargs):709        """710        This method forwards all its arguments to GPT2Tokenizer's [`~PreTrainedTokenizer.batch_decode`]. Please711        refer to the docstring of this method for more information.712        """713        return self.tokenizer.batch_decode(*args, **kwargs)714 715    # Copied from transformers.models.clip.processing_clip.CLIPProcessor.decode with CLIP->Llama716    def decode(self, *args, **kwargs):717        """718        This method forwards all its arguments to GPT2Tokenizer's [`~PreTrainedTokenizer.decode`]. Please refer to719        the docstring of this method for more information.720        """721        return self.tokenizer.decode(*args, **kwargs)722 723    @property724    # Copied from transformers.models.clip.processing_clip.CLIPProcessor.model_input_names725    def model_input_names(self):726        tokenizer_input_names = self.tokenizer.model_input_names727        image_processor_input_names = self.image_processor.model_input_names728        audio_processor_input_names = self.audio_processor.model_input_names729        return list(dict.fromkeys(tokenizer_input_names + image_processor_input_names + audio_processor_input_names))730 731 732AutoImageProcessor.register("Phi4MMImageProcessor", Phi4MMImageProcessor)733AutoFeatureExtractor.register("Phi4MMAudioFeatureExtractor", Phi4MMAudioFeatureExtractor)734