MahmoodAnaam/MSP-Processor-With-LM
1
1import torch2import torchvision.transforms.v2.functional as tvF3from torchcodec.decoders import VideoDecoder4from transformers.image_processing_utils import BatchFeature5from transformers.image_utils import PILImageResampling6from transformers.processing_utils import Unpack, VideosKwargs7from transformers.video_processing_utils import BaseVideoProcessor, VideoMetadata8from transformers.video_utils import VideoInput9 10 11class MSPVisualVideoProcessor(BaseVideoProcessor):12 resample = PILImageResampling.BILINEAR13 14 def __init__(15 self,16 do_convert_rgb_to_grayscale: bool = True,17 do_rescale: bool = True,18 rescale_factor: float = 1 / 255.0,19 image_mean=0.421,20 image_std=0.165,21 do_normalize: bool = True,22 do_resize: bool = True,23 size: dict[str, int] = {"height": 96, "width": 96},24 do_center_crop: bool = True,25 crop_size: dict[str, int] = {"height": 88, "width": 88},26 **kwargs: Unpack[VideosKwargs],27 ):28 super().__init__(29 do_rescale=do_rescale,30 rescale_factor=rescale_factor,31 image_mean=image_mean,32 image_std=image_std,33 do_normalize=do_normalize,34 do_resize=do_resize,35 size=size,36 do_center_crop=do_center_crop,37 crop_size=crop_size,38 **kwargs,39 )40 self.do_convert_rgb_to_grayscale = do_convert_rgb_to_grayscale41 42 def sample_frames(43 self,44 metadata: VideoMetadata,45 num_frames: int | None = None,46 fps: int | float | None = None,47 **kwargs,48 ):49 if num_frames:50 total_frames = metadata.total_num_frames51 num_frames = num_frames if num_frames is not None else self.num_frames52 assert num_frames is not None, (53 "`num_frames` must be specified if `fixed_len_video == True`"54 )55 frame_idxs = [56 int(i * (total_frames - 1) / (num_frames - 1))57 for i in range(num_frames)58 ]59 return torch.tensor(frame_idxs)60 else:61 return super().sample_frames(metadata, num_frames, fps, **kwargs)62 63 def _load_video(self, src: str | bytes) -> torch.Tensor:64 """65 Load video from a file path or bytes and return as a 4D torch.Tensor.66 Args:67 src (str | bytes): Path to the video file or bytes of the video file.68 Returns:69 torch.Tensor: Loaded video as a 4D tensor (num_frames, height, width, num_channels).70 """71 vd = VideoDecoder(src)72 video = vd.get_frames_in_range(0, vd.metadata.num_frames).data73 return video74 75 def __call__(76 self, videos: VideoInput | str | list[str] | bytes | list[bytes], **kwargs77 ):78 """Overrides the __call__ method to handle video input as file paths or bytes."""79 if isinstance(videos, (str, bytes)):80 videos = self._load_video(videos)81 elif isinstance(videos, list) and isinstance(videos[0], (str, bytes)):82 videos = [self._load_video(v) for v in videos]83 84 # remove kwargs not in VideosKwargs85 # for key in list(kwargs.keys()):86 # if key not in VideosKwargs.__optional_keys__:87 # kwargs.pop(key, None)88 return super().__call__(videos, **kwargs)89 90 def convert_rgb_to_grayscale(self, video: torch.Tensor) -> torch.Tensor:91 """92 Convert a video to grayscale.93 """94 video = tvF.rgb_to_grayscale(video)95 return video96 97 def _preprocess(98 self,99 videos: VideoInput,100 **kwargs: Unpack[VideosKwargs],101 ) -> BatchFeature:102 """103 Preprocesses a video or a batch of videos.104 Args:105 videos (VideoInput): Video to preprocess.106 See `VideoInput` for details.107 **kwargs: Additional keyword arguments.108 Returns:109 BatchFeature: A BatchFeature with the following fields:110 - pixel_values_videos: Pixel values to be fed to a model, of shape (batch_size,num_channels, num_frames, height, width).111 - padding_mask_videos (optional): Mask to be used for padding, of shape (batch_size, num_frames).112 """113 114 # Always set `return_tensors` to `None` since it won't pad variable length videos115 # We'll handle this after we call the parent' method116 return_tensors = kwargs.pop("return_tensors", None)117 result = super()._preprocess(videos, **kwargs)118 pixels = result.pixel_values_videos119 if self.do_convert_rgb_to_grayscale:120 pixels = [self.convert_rgb_to_grayscale(video) for video in pixels]121 data = {"pixel_values_videos": pixels}122 if return_tensors:123 lengths = torch.tensor([video.size(0) for video in pixels])124 pixels = torch.nn.utils.rnn.pad_sequence(125 pixels, batch_first=True, padding_value=0.0126 )127 data["pixel_values_videos"] = pixels128 if lengths.unique().size(0) > 1:129 mask = torch.arange(lengths.max())[None] < lengths[:, None]130 data["padding_mask_videos"] = mask131 # pixel_values_videos shape [batch_size, num_channels, num_frames, height, width]132 data["pixel_values_videos"] = data["pixel_values_videos"].permute(0, 2, 1, 3, 4)133 134 return BatchFeature(data=data, tensor_type=return_tensors)135 