Ë
    þÍ:j '  ã                   ó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/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/video/vmaf.pyr(   z)VideoMultiMethodAssessmentFusion.__init__€   s-  ø€ Ü‰ÑÑ"˜6Ò"Ý$ÜÐtÓuÐuä˜(¤DÔ)ÜÐfÓgÐgØ ˆŒà�‰�|¨RÀˆÔFØ�=Š=Ø�N‰NÐ,°bÈˆNÔOØ�N‰NÐ+°RÈˆNÔNØ�N‰N˜>°2ÀeˆNÔLØ�N‰NÐ/¸ÈEˆNÔRØ�N‰NÐ/¸ÈEˆNÔRØ�N‰NÐ/¸ÈEˆNÔRØ�N‰NÐ/¸ÈEˆNÔRØ�N‰NÐ/¸ÈEˆNÔRØ�N‰NÐ/¸ÈEˆNÔRØ�N‰NÐ/¸ÈEˆNÔRØ�N‰NÐ/¸ÈEˆNÕ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  € ä4°U¸FÀDÇMÁMÓRˆØ�=‹=œZ¨¬tÕ4Ø�O‰O×"Ñ" 5¨¡=Ô1Ø× Ñ ×'Ñ'¨Ð.?Ñ(@ÔAØ×Ñ×&Ñ& uÐ-=Ñ'>Ô?Ø×Ñ×$Ñ$ U¨>Ñ%:Ô;Ø×#Ñ#×*Ñ*¨5Ð1EÑ+FÔGØ×#Ñ#×*Ñ*¨5Ð1EÑ+FÔGØ×#Ñ#×*Ñ*¨5Ð1EÑ+FÔGØ×#Ñ#×*Ñ*¨5Ð1EÑ+FÔGØ×#Ñ#×*Ñ*¨5Ð1EÑ+FÔGØ×#Ñ#×*Ñ*¨5Ð1EÑ+FÔGØ×#Ñ#×*Ñ*¨5Ð1EÑ+FÔGØ×#Ñ#×*Ñ*¨5Ð1EÑ+FÕGÜ˜œvÔ&Ø�O‰O×"Ñ" 5Õ)ð 'r1   c                 ó*  — | j                   rót        | 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Ü#/°×0DÑ0DÓ#EÜ".¨t×/BÑ/BÓ"CÜ ,¨T×->Ñ->Ó ?Ü&2°4×3JÑ3JÓ&KÜ&2°4×3JÑ3JÓ&KÜ&2°4×3JÑ3JÓ&KÜ&2°4×3JÑ3JÓ&KÜ&2°4×3JÑ3JÓ&KÜ&2°4×3JÑ3JÓ&KÜ&2°4×3JÑ3JÓ&KÜ&2°4×3JÑ3JÓ&Kñð ô ˜DŸO™OÓ,Ð,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ó *ð&-˜˜v t¨C°¨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   