
    i '                     x    d dl mZmZmZmZ d dlmZ d dlmZ d dl	m
Z
 d dlmZ d dlmZ esdgZ G d de
      Zy	)
    )AnyDictListUnion)Tensor)$video_multi_method_assessment_fusion)Metric)dim_zero_cat)_TORCH_VMAF_AVAILABLE VideoMultiMethodAssessmentFusionc                   n    e Zd ZU dZdZeed<   dZeed<   dZeed<   dZ	e
ed<   d	Ze
ed
<   ee   ed<   ee   ed<   ee   ed<   ee   ed<   ee   ed<   ee   ed<   ee   ed<   ee   ed<   ee   ed<   ee   ed<   ee   ed<   ee   ed<   d dededdf fdZdededdfdZdeeeeef   f   fdZ xZS )!r   aY  Calculates Video Multi-Method Assessment Fusion (VMAF) metric.

    VMAF is a full-reference video quality assessment algorithm that combines multiple quality assessment features
    such as detail loss, motion, and contrast using a machine learning model to predict human perception of video
    quality more accurately than traditional metrics like PSNR or SSIM.

    The metric works by:

       1. Converting input videos to luma component (grayscale)
       2. Computing multiple elementary features:
          - Additive Detail Measure (ADM): Evaluates detail preservation at different scales
          - Visual Information Fidelity (VIF): Measures preservation of visual information across frequency bands
          - Motion: Quantifies the amount of motion in the video
       3. Combining these features using a trained SVM model to predict quality

    .. note::
       This implementation requires you to have vmaf-torch installed: https://github.com/alvitrioliks/VMAF-torch.
       Install either by cloning the repository and running ``pip install .``
       or with ``pip install torchmetrics[video]``.

    As input to ``forward`` and ``update`` the metric accepts the following input:

        - ``preds`` (:class:`~torch.Tensor`): Video tensor of shape ``(batch, channels, frames, height, width)``.
          Expected to be in RGB format with values in range [0, 1].
        - ``target`` (:class:`~torch.Tensor`): Video tensor of shape ``(batch, channels, frames, height, width)``.
          Expected to be in RGB format with values in range [0, 1].

    As output of ``forward`` and ``compute`` the metric returns the following output ``vmaf`` (:class:`~torch.Tensor`):

        - If ``features`` is False, returns a tensor with shape (batch, frame)
          of VMAF score for each frame in each video. Higher scores indicate better quality, with typical values
          ranging from 0 to 100.
        - If ``features`` is True, returns a dictionary where each value is a (batch, frame) tensor of the
          corresponding feature. The keys are:

            - 'integer_motion2': Integer motion feature
            - 'integer_motion': Integer motion feature
            - 'integer_adm2': Integer ADM feature
            - 'integer_adm_scale0': Integer ADM feature at scale 0
            - 'integer_adm_scale1': Integer ADM feature at scale 1
            - 'integer_adm_scale2': Integer ADM feature at scale 2
            - 'integer_adm_scale3': Integer ADM feature at scale 3
            - 'integer_vif_scale0': Integer VIF feature at scale 0
            - 'integer_vif_scale1': Integer VIF feature at scale 1
            - 'integer_vif_scale2': Integer VIF feature at scale 2
            - 'integer_vif_scale3': Integer VIF feature at scale 3
            - 'vmaf': VMAF score for each frame in each video

    Args:
        features: If True, all the elementary features (ADM, VIF, motion) are returned along with the VMAF score in
            a dictionary. This corresponds to the output you would get from the VMAF command line tool with
            the ``--csv`` option enabled. If False, only the VMAF score is returned as a tensor.
        kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info.

    Raises:
        RuntimeError:
            If vmaf-torch is not installed.
        ValueError:
            If ``features`` is not a boolean.

    Example:
        >>> import torch
        >>> from torchmetrics.video import VideoMultiMethodAssessmentFusion
        >>> # 2 videos, 3 channels, 10 frames, 32x32 resolution
        >>> preds = torch.rand(2, 3, 10, 32, 32, generator=torch.manual_seed(42))
        >>> target = torch.rand(2, 3, 10, 32, 32, generator=torch.manual_seed(43))
        >>> vmaf = VideoMultiMethodAssessmentFusion()
        >>> torch.round(vmaf(preds, target), decimals=2)
        tensor([[ 9.9900, 15.9000, 14.2600, 16.6100, 15.9100, 14.3000, 13.5800, 13.4900, 15.4700, 20.2800],
                [ 6.2500, 11.3000, 17.3000, 11.4600, 19.0600, 14.9300, 14.0500, 14.4100, 12.4700, 14.8200]])
        >>> vmaf = VideoMultiMethodAssessmentFusion(features=True)
        >>> vmaf_dict = vmaf(preds, target)
        >>> vmaf_dict['vmaf'].round(decimals=2)
        tensor([[ 9.9900, 15.9000, 14.2600, 16.6100, 15.9100, 14.3000, 13.5800, 13.4900, 15.4700, 20.2800],
                [ 6.2500, 11.3000, 17.3000, 11.4600, 19.0600, 14.9300, 14.0500, 14.4100, 12.4700, 14.8200]])
        >>> vmaf_dict['integer_adm2'].round(decimals=2)
        tensor([[0.4500, 0.4500, 0.3600, 0.4700, 0.4300, 0.3600, 0.3900, 0.4100, 0.3700, 0.4700],
                [0.4200, 0.3900, 0.4400, 0.3700, 0.4500, 0.3900, 0.3800, 0.4800, 0.3900, 0.3900]])

    Fis_differentiableThigher_is_betterfull_state_updateg        plot_lower_boundg      Y@plot_upper_bound
vmaf_scoreinteger_motion2integer_motioninteger_adm2integer_adm_scale0integer_adm_scale1integer_adm_scale2integer_adm_scale3integer_vif_scale0integer_vif_scale1integer_vif_scale2integer_vif_scale3featureskwargsreturnNc                    t        |   di | t        st        d      t	        |t
              st        d      || _        | j                  dg d       | j                  r| j                  dg d       | j                  dg d       | j                  dg d       | j                  d	g d       | j                  d
g d       | j                  dg d       | j                  dg d       | j                  dg d       | j                  dg d       | j                  dg d       | j                  dg d       y y )NzSvmaf-torch is not installed. Please install with `pip install torchmetrics[video]`.zGArgument `elementary_features` should be a boolean, but got {features}.r   cat)defaultdist_reduce_fxr   r   r   r   r   r   r   r   r   r   r    )	super__init__r   RuntimeError
isinstancebool
ValueErrorr   	add_state)selfr   r    	__class__s      l/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchmetrics/video/vmaf.pyr(   z)VideoMultiMethodAssessmentFusion.__init__   s-   "6"$tuu(D)fgg |RF==NN,bNONN+RNNNN>2eNLNN/ENRNN/ENRNN/ENRNN/ENRNN/ENRNN/ENRNN/ENRNN/ENR     predstargetc                    t        ||| j                        }| j                  rzt        |t              ri| j                  j                  |d          | j                  j                  |d          | j                  j                  |d          | j                  j                  |d          | j                  j                  |d          | j                  j                  |d          | j                  j                  |d          | j                  j                  |d          | j                  j                  |d	          | j                  j                  |d
          | j                  j                  |d          | j                   j                  |d          yt        |t"              r| j                  j                  |       yy)z*Update state with predictions and targets.vmafr   r   r   r   r   r   r   r   r   r   r   N)r   r   r*   dictr   appendr   r   r   r   r   r   r   r   r   r   r   r   )r.   r2   r3   scores       r0   updatez'VideoMultiMethodAssessmentFusion.update   sx   4UFDMMR==Zt4OO""5=1  ''.?(@A&&u-='>?$$U>%:;##**51E+FG##**51E+FG##**51E+FG##**51E+FG##**51E+FG##**51E+FG##**51E+FG##**51E+FGv&OO""5) 'r1   c                 *   | j                   rt        | j                        t        | j                        t        | j                        t        | j
                        t        | j                        t        | j                        t        | j                        t        | j                        t        | j                        t        | j                        t        | j                        t        | j                        dS t        | j                        S )zCompute final VMAF score.)r5   r   r   r   r   r   r   r   r   r   r   r   )r   r
   r   r   r   r   r   r   r   r   r   r   r   r   )r.   s    r0   computez(VideoMultiMethodAssessmentFusion.compute   s    ==$T__5#/0D0D#E".t/B/B"C ,T->-> ?&243J3J&K&243J3J&K&243J3J&K&243J3J&K&243J3J&K&243J3J&K&243J3J&K&243J3J&K  DOO,,r1   )F)__name__
__module____qualname____doc__r   r+   __annotations__r   r   r   floatr   r   r   r   r(   r9   r   r   strr;   __classcell__)r/   s   @r0   r   r      s   Ob $t#!d!#t#!e!#e#V&\!L v,V$V$V$V$V$V$V$V$S S S S.*F *F *t *&-vtCK'889 -r1   N)typingr   r   r   r   torchr   "torchmetrics.functional.video.vmafr   torchmetrics.metricr	   torchmetrics.utilities.datar
   torchmetrics.utilities.importsr   __doctest_skip__r   r&   r1   r0   <module>rK      s6    * )  S & 4 @:;`-v `-r1   