Ë
    þÍ:j¼/  ã                   ó~  — d dl Z d dlmZ d dlZd dlmZmZ d dlmZ d dl	m
Z
 d dlmZmZ erd dlmZ d dlmZ d	d
dœZe
d   ZdZeresdgZ G d„ dej,                  «      Z G d„ de«      Z	 ddedej,                  dededeeeeef   f   f
d„Z	 d dedeeef   de
d   defd„Z	 	 	 	 d!dedede
d   dededefd„Zy)"é    N)ÚUnion)ÚTensorÚnn©Ú	normalize)ÚLiteral)Ú_TORCH_GREATER_EQUAL_2_2Ú_TORCHVISION_AVAILABLE)Ú
transforms)Úresnet50)é   é   )r   éd   )Úkadid10kÚkoniq10kz=https://github.com/miccunifi/ARNIQA/releases/download/weightsÚarniqac            	       ó|   ‡ — e Zd ZdZddeddfˆ fd„Zdd„Zddedede	eef   fd	„Z
d
edefd„Zddededefd„Zˆ xZS )Ú_ARNIQAzÌInitializes a No-Reference Image Quality Assessment ARNIQA torch.nn.Module.

    Args:
        regressor_dataset: dataset used for training the regressor, choose between [``koniq10k``, ``kadid10k``]

    Úregressor_datasetÚreturnNc                 ó€  •— t         ‰| �  «        t        st        d«      ‚t        st        d«      ‚t        j                  «       }||vrt        d|› d|› d�«      ‚|| _	        g d¢| _
        g d¢| _        t        «       }|j                  j                  | _        t!        j"                  t%        |j'                  «       «      d d Ž }|| _        t!        j*                  | j                  d	z  d
«      | _        | j/                  «        dt         j0                  dd fd„} || j(                  «        || j,                  «       y )Nz'ARNIQA metric requires PyTorch >= 2.2.0z‡ARNIQA metric requires that torchvision is installed. Either install as `pip install torchmetrics[image]` or `pip install torchvision`.z,Argument `regressor_dataset` must be one of ú
, but got ú.)g
×£p=
ß?gÉv¾Ÿ/Ý?g–C‹lçûÙ?)gZd;ßOÍ?gyé&1¬Ì?gÍÌÌÌÌÌÌ?éÿÿÿÿé   r   Úmoduler   c                 ó\   — | j                  «        | j                  «       D ]	  }d|_        Œ y )NF)ÚevalÚ
parametersÚrequires_grad)r   Úps     úy/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/image/arniqa.pyÚ_freezez!_ARNIQA.__init__.<locals>._freezeU   s*   € Ø�K‰KŒMØ×&Ñ&Ó(ò (�Ø"'�•ñ(ó    )ÚsuperÚ__init__r	   ÚRuntimeErrorr
   ÚModuleNotFoundErrorÚ_AVAILABLE_REGRESSOR_DATASETSÚkeysÚ
ValueErrorr   Úimagenet_norm_meanÚimagenet_norm_stdr   ÚfcÚin_featuresÚfeat_dimr   Ú
SequentialÚlistÚchildrenÚencoderÚLinearÚ	regressorÚ_load_weightsÚModule)Úselfr   Úvalid_regressor_datasetsr4   r#   Ú	__class__s        €r"   r&   z_ARNIQA.__init__8   s'  ø€ Ü‰ÑÔå'ÜÐHÓIÐIå%Ü%ðeóð ô
 $A×#EÑ#EÓ#GÐ ØÐ$<Ñ<ÜØ>Ð?WÐ>XÐXbÐctÐbuÐuvÐwóð ð "3ˆÔÚ"7ˆÔÚ!6ˆÔä“*ˆØŸ
™
×.Ñ.ˆŒÜ—-‘-¤ g×&6Ñ&6Ó&8Ó!9¸#¸2Ð!>Ð?ˆØˆŒÜŸ™ 4§=¡=°1Ñ#4°aÓ8ˆŒØ×ÑÔð	(œBŸI™Ið 	(¨$ó 	(ñ
 	�—‘ÔÙ�—‘Õr$   c                 óÆ  — t         j                  j                  t        › d�dd¬«      }|j	                  «       D ��ci c]  \  }}d|vsŒ|j                  dd«      |“Œ }}}| j                  j                  |d¬«       t        j                  «       5  t        j                  d	t        d
¬«       t         j                  j                  t        › d| j                  › d�dd¬«      j                  «       }|j                  d«      |d<   |j                  d«      j                  d«      |d<   | j                   j                  |d¬«       ddd«       yc c}}w # 1 sw Y   yxY w)z/Loads the weights of the encoder and regressor.z/ARNIQA.pthTÚcpu)ÚprogressÚmap_locationÚ	projectorzmodel.Ú )ÚstrictÚignoreztorch.serialization)Úcategoryr   z/regressor_z.pthÚweightsÚweightÚbiasesr   ÚbiasN)ÚtorchÚhubÚload_state_dict_from_urlÚ	_base_urlÚitemsÚreplacer4   Úload_state_dictÚwarningsÚcatch_warningsÚfilterwarningsÚUserWarningr   Ú
state_dictÚpopÚ	unsqueezer6   )r9   Úencoder_state_dictÚkÚvÚfiltered_encoder_state_dictÚregressor_state_dicts         r"   r7   z_ARNIQA._load_weights]   s`  € ä"ŸY™Y×?Ñ?Üˆk˜Ð%°À5ð @ó 
Ðð 4F×3KÑ3KÓ3M÷'
Ù+/¨1¨aÐQ\ÐdeÒQeˆA�I‰I�h Ó# QÑ&ð'
Ð#ñ '
ð 	�‰×$Ñ$Ð%@ÈÐ$ÔNä×$Ñ$Ó&ñ 	NÜ×#Ñ# H´{ÐK`ÕaÜ#(§9¡9×#EÑ#EÜ�+˜[¨×)?Ñ)?Ð(@ÀÐEÐPTÐchð $Fó $ç‰j‹lð !ð .B×-EÑ-EÀiÓ-PÐ  Ñ*Ø+?×+CÑ+CÀHÓ+M×+WÑ+WÐXYÓ+ZÐ  Ñ(Ø�N‰N×*Ñ*Ð+?ÈÐ*ÔM÷	Nð 	Nùó'
÷
	Nð 	Nús   ½EÁ
EÂB5EÅE Úimgr   c                 ó<  — |j                   dd \  }} t        j                  |dz  |dz  f«      |«      }|rb t        j                  | j                  | j
                  ¬«      |«      } t        j                  | j                  | j
                  ¬«      |«      }||fS )zŽPreprocesses the input to the model.

        Obtains the half-scale version of the input image and applies normalization if needed.

        éþÿÿÿNr   )ÚmeanÚstd)Úshaper   ÚResizeÚ	Normalizer,   r-   )r9   r\   r   ÚhÚwÚimg_dss         r"   Ú_preprocess_inputz_ARNIQA._preprocess_inputq   s”   € ð �y‰y˜˜ˆ~‰ˆˆ1Ø4”×"Ñ" A¨¡F¨A°©FÐ#3Ó4°SÓ9ˆÙØ`”*×&Ñ&¨D×,CÑ,CÈ×I_ÑI_Ô`ÐadÓeˆCØc”Z×)Ñ)¨t×/FÑ/FÈD×LbÑLbÔcÐdjÓkˆFØ�Fˆ{Ðr$   Úscorec                 óD   — t         | j                     \  }}||z
  ||z
  z  S )zKScales the quality score to be in the [0, 1] range, where higher is better.)r)   r   )r9   rh   Ú	min_scoreÚ	max_scores       r"   Ú_scale_scorez_ARNIQA._scale_score~   s,   € ä<¸T×=SÑ=SÑTÑˆ	�9Ø˜	Ñ! i°)Ñ&;Ñ<Ð<r$   c                 ó†  — | j                  ||«      \  }}| j                  |«      }|j                  d| j                  «      }t	        |d¬«      }| j                  |«      }|j                  d| j                  «      }t	        |d¬«      }t        j                  ||f«      }| j                  |«      }| j                  |«      S )Nr   r   )Údim)	rg   r4   Úviewr0   Únormalize_fnrI   Úhstackr6   rl   )r9   r\   r   rf   Úimg_fÚimg_ds_fÚfrh   s           r"   Úforwardz_ARNIQA.forwardƒ   sª   € à×,Ñ,¨S°)Ó<‰ˆˆVð —‘˜SÓ!ˆØ—
‘
˜2˜tŸ}™}Ó-ˆÜ˜U¨Ô*ˆØ—<‘< Ó'ˆØ—=‘=  T§]¡]Ó3ˆÜ ¨aÔ0ˆÜ�L‰L˜% Ð*Ó+ˆð —‘˜qÓ!ˆØ× Ñ  Ó'Ð'r$   )r   )r   N©F)Ú__name__Ú
__module__Ú__qualname__Ú__doc__Ú_TYPE_REGRESSOR_DATASETr&   r7   r   ÚboolÚtuplerg   rl   ru   Ú__classcell__©r;   s   @r"   r   r   0   su   ø„ ññ# Ð*Að # ÐSWõ # óJNñ( Vð ¸ð ÈÈvÐW]È~ÑI^ó ð= &ð =¨Vó =ñ
(˜6ð (¨dð (¸v÷ (r$   r   c                   ó,   ‡ — e Zd ZdZdedd fˆ fd„Zˆ xZS )Ú_NoTrainArniqaz9Wrapper to make sure ARNIQA never leaves evaluation mode.Úmoder   c                 ó"   •— t         ‰| �  d«      S )z.Force network to always be in evaluation mode.F)r%   Útrain)r9   r‚   r;   s     €r"   r„   z_NoTrainArniqa.train˜   s   ø€ ä‰w‰}˜UÓ#Ð#r$   )rw   rx   ry   rz   r|   r„   r~   r   s   @r"   r�   r�   •   s   ø„ ÙCð$˜$ð $Ð#3÷ $ñ $r$   r�   r\   Úmodelr   Úautocastr   c                 ój  — | j                   dk(  r| j                  d   dk(  st        d| j                  › d�«      ‚| j                  «       dk  r| j	                  «       dk\  s0|r.t        d| j	                  «       › d	| j                  «       › d�«      ‚|rSt
        j                  j                  | j                  j                  | j                  ¬
«      5   || |¬«      }ddd«       n$ |j                  | j                  ¬«      | |¬«      }j                  «       | j                  d   fS # 1 sw Y   Œ(xY w)a  Update step for ARNIQA metric.

    Args:
        img: the input image
        model: the pre-trained model
        normalize: boolean indicating whether the input image is normalized
        autocast: boolean indicating whether to use automatic mixed precision

    é   r   é   z?Input image must have shape [N, 3, H, W]. Got input with shape r   g      ð?g        zdInput image values must be in the [0, 1] range when normalize==True. Got input with values in range z and )Údevice_typeÚdtyper   N)r‹   r   )Úndimra   r+   ÚmaxÚminrI   Úampr†   ÚdeviceÚtyper‹   ÚtoÚsqueeze)r\   r…   r   r†   Úlosss        r"   Ú_arniqa_updater•   �   s  € ð �H‰H˜ŠM˜cŸi™i¨™l¨aÒ/ÜÐZÐ[^×[dÑ[dÐZeÐefÐgÓhÐhØ�G‰G‹I˜Ò §¡£¨cÒ!1±yÜðØŸ™›˜ 5¨¯©«¨°1ð6ó
ð 	
ñ
 Ü�Y‰Y×Ñ¨C¯J©J¯O©OÀ3Ç9Á9ÐÓMñ 	3Ù˜¨	Ô2ˆD÷	3ð 	3ð )ˆu�x‰x˜cŸi™iˆxÓ(¨¸	ÔBˆØ�<‰<‹>˜3Ÿ9™9 Q™<Ð'Ð'÷		3ð 	3ús   ÃD)Ä)D2ÚscoresÚ
num_scoresÚ	reduction)Úsumr_   Únonec                 óH   — | j                  «       }|dk(  r| S |dk(  r||z  S |S )zCompute step for ARNIQA metric.rš   r_   )r™   )r–   r—   r˜   Ú
sum_scoress       r"   Ú_arniqa_computer�   º   s5   € ð —‘“€JØ�FÒØˆØ�FÒØ˜JÑ&Ð&ØÐr$   r   c                 ó  — d}||vrt        d|› d|› �«      ‚t        |t        «      st        d|› �«      ‚t        |¬«      j	                  | j
                  | j                  ¬«      }t        | |||¬«      \  }}t        |||«      S )a  ARNIQA: leArning distoRtion maNifold for Image Quality Assessment metric.

    `ARNIQA`_ is a No-Reference Image Quality Assessment metric that predicts the technical quality of an image with
    a high correlation with human opinions. ARNIQA consists of an encoder and a regressor. The encoder is a ResNet-50
    model trained in a self-supervised way to model the image distortion manifold to generate similar representation for
    images with similar distortions, regardless of the image content. The regressor is a linear model trained on IQA
    datasets using the ground-truth quality scores. ARNIQA extracts the features from the full- and half-scale versions
    of the input image and then outputs a quality score in the [0, 1] range, where higher is better.

    The input image is expected to have shape ``(N, 3, H, W)``. The image should be in the [0, 1] range if `normalize`
    is set to ``True``, otherwise it should be normalized with the ImageNet mean and standard deviation.

    .. note::
        Using this metric requires you to have ``torchvision`` package installed. Either install as
        ``pip install torchmetrics[image]`` or ``pip install torchvision``.

    Args:
        img: the input image
        regressor_dataset: dataset used for training the regressor. Choose between [``koniq10k``, ``kadid10k``].
            ``koniq10k`` corresponds to the `KonIQ-10k`_ dataset, which consists of real-world images with authentic
            distortions. ``kadid10k`` corresponds to the `KADID-10k`_ dataset, which consists of images with
            synthetically generated distortions.
        reduction: indicates how to reduce over the batch dimension. Choose between [``sum``, ``mean``, ``none``].
        normalize: by default this is ``True`` meaning that the input is expected to be in the [0, 1] range. If set
            to ``False`` will instead expect input to be already normalized with the ImageNet mean and standard
            deviation.
        autocast: boolean indicating whether to use automatic mixed precision

    Returns:
        A tensor in the [0, 1] range, where higher is better, representing the ARNIQA score of the input image. If
        `reduction` is set to ``none``, the output will have shape ``(N,)``, otherwise it will be a scalar tensor.

    Raises:
        ModuleNotFoundError:
            If ``torchvision`` package is not installed
        ValueError:
            If ``regressor_dataset`` is not in [``"kadid10k"``, ``"koniq10k"``]
        ValueError:
            If ``reduction`` is not in [``"sum"``, ``"mean"``, ``"none"``]
        ValueError:
            If ``normalize`` is not a bool
        ValueError:
            If the input image is not a valid image tensor with shape [N, 3, H, W].
        ValueError:
            If the input image values are not in the [0, 1] range when ``normalize`` is set to ``True``

    Examples:
        >>> from torch import rand
        >>> from torchmetrics.functional.image.arniqa import arniqa
        >>> img = rand(8, 3, 224, 224)
        >>> # Non-normalized input
        >>> arniqa(img, regressor_dataset='koniq10k', normalize=True)
        tensor(0.5308)


        >>> from torch import rand
        >>> from torchmetrics.functional.image.arniqa import arniqa
        >>> from torchvision.transforms import Normalize
        >>> img = rand(8, 3, 224, 224)
        >>> img = Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])(img)
        >>> # Normalized input
        >>> arniqa(img, regressor_dataset='koniq10k', normalize=False)
        tensor(0.5065)

    )r_   r™   rš   z$Argument `reduction` must be one of r   z.Argument `normalize` should be a bool but got )r   )r�   r‹   )r   r†   )	r+   Ú
isinstancer|   r�   r’   r�   r‹   r•   r�   )	r\   r   r˜   r   r†   Úvalid_reductionr…   r”   r—   s	            r"   r   r   Æ   s—   € ðP .€OØ˜Ñ'ÜÐ?ÀÐ?PÐPZÐ[dÐZeÐfÓgÐgä�i¤Ô&ÜÐIÈ)ÈÐUÓVÐVäÐ->Ô?×BÑBÈ#Ï*É*Ð\_×\eÑ\eÐBÓf€EÜ% c¨5¸IÐPXÔYÑ€Dˆ*Ü˜4 ¨YÓ7Ð7r$   rv   )r_   )r   r_   TF)rP   Útypingr   rI   r   r   Útorch.nn.functionalr   rp   Útyping_extensionsr   Útorchmetrics.utilities.importsr	   r
   Útorchvisionr   Útorchvision.modelsr   r)   r{   rL   Ú__doctest_skip__r8   r   r�   r|   r}   Úintr•   r�   r   © r$   r"   ú<module>rª      sU  ðó( Ý ã ß Ý 9Ý %ç [áÝ&Ý+ð Øñ!Ð ð
 "Ð"8Ñ9Ð àK€	ñ 	!Ñ%;Ø �zÐôb(ˆb�i‰iô b(ôJ$�Wô $ð FKñ(Ø	ð(ØŸ	™	ð(Ø.2ð(Ø>Bð(à
ˆ6�5˜˜f˜Ñ%Ð%Ñ&ó(ð< agñ	Øð	Ø % f¨c kÑ 2ð	Ø?FÐG\Ñ?]ð	àó	ð 2<Ø06ØØñQ8Ø	ðQ8à.ðQ8ð Ð,Ñ-ðQ8ð ð	Q8ð
 ðQ8ð ôQ8r$   