Ë
    þÍ:j¢!  ã                   óZ  — d dl Z d dlmZmZ d dlZd dlmZmZ d dlmZm	Z	 d dl
mZ d dlmZ d dlmZ ded	ed
ededeeeef   f
d„Zdedeeeedf   f   defd„Zdededefd„Zdededefd„Zded	ededeeeef   fd„Zded	ed
ededef
d„Z	 	 	 dded	ed
ee   dedeed      defd„Zy)é    N)ÚOptionalÚUnion)ÚTensorÚtensor)Úconv2dÚpad)ÚLiteral)Ú_check_same_shape)ÚreduceÚpredsÚtargetÚ	hp_filterÚwindow_sizeÚreturnc           
      ó   — | j                   |j                   k7  r|j                  | j                   «      }t        | |«       | j                  dvr&t	        d| j
                  › d|j
                  › d�«      ‚t        | j
                  «      dk(  r"| j                  d«      } |j                  d«      }|dkD  st	        d|› d�«      ‚|| j                  d	«      kD  s|| j                  d«      kD  r3t	        d
|› d| j                  d	«      › d| j                  d«      › d�«      ‚| j                  t        j                  «      } |j                  t        j                  «      }|dddd…f   j                  | j                   | j                  ¬«      }| ||fS )a¤  Update and returns variables required to compute Spatial Correlation Coefficient.

    Args:
        preds: Predicted tensor
        target: Ground truth tensor
        hp_filter: High-pass filter tensor
        window_size: Local window size integer

    Return:
        Tuple of (preds, target, hp_filter) tensors

    Raises:
        ValueError:
            If ``preds`` and ``target`` have different number of channels
            If ``preds`` and ``target`` have different shapes
            If ``preds`` and ``target`` have invalid shapes
            If ``window_size`` is not a positive integer
            If ``window_size`` is greater than the size of the image

    )é   é   zŠExpected `preds` and `target` to have batch of colored images with BxCxHxW shape  or batch of grayscale images of BxHxW shape. Got preds: z and target: ú.r   é   r   z5Expected `window_size` to be a positive integer. Got é   z[Expected `window_size` to be less than or equal to the size of the image. Got window_size: z and image size: ÚxN)ÚdtypeÚdevice)r   Útor
   ÚndimÚ
ValueErrorÚshapeÚlenÚ	unsqueezeÚsizeÚtorchÚfloat32r   )r   r   r   r   s       úv/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/image/scc.pyÚ_scc_updater$      sq  € ð* ‡{�{�f—l‘lÒ"Ø—‘˜5Ÿ;™;Ó'ˆÜ�e˜VÔ$Ø‡z�z˜ÑÜðà Ÿ;™;˜- }°V·\±\°NÀ!ðEó
ð 	
ô ˆ5�;‰;Ó˜1ÒØ—‘ Ó"ˆØ×!Ñ! !Ó$ˆà˜Š?ÜÐPÐQ\ÐP]Ð]^Ð_Ó`Ð`à�U—Z‘Z “]Ò" k°E·J±J¸q³MÒ&AÜð!Ø!, Ð->¸u¿z¹zÈ!»}¸oÈQÈuÏzÉzÐZ[Ë}ÈoÐ]^ð`ó
ð 	
ð
 �H‰H”U—]‘]Ó#€EØ�Y‰Y”u—}‘}Ó%€FØ˜$ ¢a˜-Ñ(×+Ñ+°%·+±+ÀeÇlÁlÐ+ÓS€IØ�&˜)Ð#Ð#ó    Ú	input_imgr   .c                 ó  — t        |t        «      r||||f}t        |«      dk7  rt        dt        |«      › �«      ‚| dd…dd…dd…d|d   …f   j	                  dg¬«      }| dd…dd…dd…|d    d…f   j	                  dg¬«      }t        j                  || |gd¬«      }|dd…dd…d|d	   …dd…f   j	                  d	g¬«      }|dd…dd…|d    d…dd…f   j	                  d	g¬«      }t        j                  |||gd	¬«      S )
zlApplies symmetric padding to the 2D image tensor input using ``reflect`` mode (d c b a | a b c d | d c b a).r   z+Expected padding to have length 4, but got Nr   r   )Údimsr   ©Údimr   )Ú
isinstanceÚintr   r   Úflipr!   Úcat)r&   r   Úleft_padÚ	right_padÚpaddedÚtop_padÚ
bottom_pads          r#   Ú_symmetric_reflect_pad_2dr4   L   s  € ä�#”sÔØ�C˜˜cÐ"ˆÜ
ˆ3ƒx�1‚}ÜÐFÄsÈ3ÃxÀjÐQÓRÐRàššAšq ! c¨!¡f *Ð,Ñ-×2Ñ2¸¸Ð2Ó<€HØš!šQ¢ C¨¡F 7¡9Ð,Ñ-×2Ñ2¸¸Ð2Ó<€IÜ�Y‰Y˜ )¨YÐ7¸QÔ?€Fà’Qš˜1˜s 1™v˜:¢qÐ(Ñ)×.Ñ.°Q°CÐ.Ó8€GØšš1˜s 1™v˜g™iªÐ*Ñ+×0Ñ0°q°cÐ0Ó:€JÜ�9‰9�g˜v zÐ2¸Ô:Ð:r%   Úkernelc                 ó¸  — t        j                  |j                  d«      dz
  dz  «      }t        j                  |j                  d«      dz
  dz  «      }t        j                  |j                  d«      dz
  dz  «      }t        j                  |j                  d«      dz
  dz  «      }t	        | ||||f¬«      }|j                  ddg«      }t        ||dd¬«      S )zHApplies 2D signal convolution to the input tensor with the given kernel.r   r   r   )r   r   ©ÚstrideÚpadding)ÚmathÚfloorr    Úceilr4   r-   r   )r&   r5   Úleft_paddingÚright_paddingÚtop_paddingÚbottom_paddingr1   s          r#   Ú_signal_convolve_2drA   \   s¼   € ä—:‘:˜vŸ{™{¨1›~°Ñ1°QÑ6Ó7€LÜ—I‘I˜vŸ{™{¨1›~°Ñ1°QÑ6Ó7€MÜ—*‘*˜fŸk™k¨!›n¨qÑ0°AÑ5Ó6€KÜ—Y‘Y §¡¨A£°Ñ 2°aÑ7Ó8€Nä& y°|À]ÐT_ÐaoÐ6pÔq€FØ�[‰[˜!˜Q˜Ó €FÜ�&˜&¨°AÔ6Ð6r%   c                 ó    — t        | |«      dz  S )zOApplies 2-D Laplace filter to the input tensor with the given high pass filter.g       @)rA   )r&   r5   s     r#   Ú_hp_2d_laplacianrC   h   s   € ä˜y¨&Ó1°CÑ7Ð7r%   Úwindowc                 óÀ  — t        j                  |j                  d«      dz
  dz  «      }t        j                  |j                  d«      dz
  dz  «      }t	        | ||||f«      } t	        |||||f«      }t        | |dd¬«      }t        ||dd¬«      }t        | dz  |dd¬«      |dz  z
  }t        |dz  |dd¬«      |dz  z
  }t        || z  |dd¬«      ||z  z
  }	|||	fS )z<Computes local variance and covariance of the input tensors.r   r   r   r   r7   )r:   r<   r    r;   r   r   )
r   r   rD   r=   r>   Ú
preds_meanÚtarget_meanÚ	preds_varÚ
target_varÚtarget_preds_covs
             r#   Ú_local_variance_covariancerK   m   sö   € ô
 —9‘9˜fŸk™k¨!›n¨qÑ0°AÑ5Ó6€LÜ—J‘J §¡¨A£°Ñ 2°aÑ7Ó8€Mä�˜ m°\À=ÐQÓR€EÜ�˜,¨°|À]ÐSÓT€Fä˜˜v¨a¸Ô;€JÜ˜ °¸1Ô=€Kä�u˜a‘x °¸1Ô=À
ÈAÁÑM€IÜ˜ ™	 6°!¸QÔ?À+ÈqÁ.ÑP€JÜ˜f u™n¨f¸QÈÔJÈ[Ð[eÑMeÑeÐà�jÐ"2Ð2Ð2r%   c                 óp  — | j                   }| j                  }t        j                  dd||f||¬«      |dz  z  }t	        | |«      }t	        ||«      }t        |||«      \  }	}
}d|	|	dk  <   d|
|
dk  <   t        j                  |
«      t        j                  |	«      z  }|dk(  }d||dk(  <   ||z  }d||<   |S )a[  Computes per channel Spatial Correlation Coefficient.

    Args:
        preds: estimated image of Bx1xHxW shape.
        target: ground truth image of Bx1xHxW shape.
        hp_filter: 2D high-pass filter.
        window_size: size of window for local mean calculation.

    Return:
        Tensor with Spatial Correlation Coefficient score

    r   )r    r   r   r   r   )r   r   r!   ÚonesrC   rK   Úsqrt)r   r   r   r   r   r   rD   Úpreds_hpÚ	target_hprH   rI   rJ   ÚdenÚidxÚsccs                  r#   Ú_scc_per_channel_computerT   ‚   sÖ   € ð �K‰K€EØ�\‰\€Fô
 �Z‰Z˜a  K°Ð=ÀUÐSYÔZÐ^iÐklÑ^lÑm€Fä  yÓ1€HÜ  ¨Ó3€Iä.HÈÐS\Ð^dÓ.eÑ+€IˆzÐ+à €Iˆi˜!‰mÑØ!"€Jˆz˜A‰~Ñä
�*‰*�ZÓ
 ¤5§:¡:¨iÓ#8Ñ
8€CØ
�‰(€CØ€Cˆˆq‰�MØ
˜SÑ
 €CØ€Cˆ�HØ€Jr%   Ú	reduction)ÚmeanÚnoneNc                 ó  — |€t        g d¢g d¢g d¢g«      }|€d}|dvrt        d|› �«      ‚t        | |||«      \  } }}t        | j	                  d«      «      D �cg c]H  }t        | dd…|dd…dd…f   j                  d«      |dd…|dd…dd…f   j                  d«      ||«      ‘ŒJ }}|dk(  r.t        j                  t        j                  |d¬«      g d	¢¬«      S |d
k(  r"t        t        j                  |d¬«      d¬«      S yc c}w )a  Compute Spatial Correlation Coefficient (SCC_).

    Args:
        preds: predicted images of shape ``(N,C,H,W)`` or ``(N,H,W)``.
        target: ground truth images of shape ``(N,C,H,W)`` or ``(N,H,W)``.
        hp_filter: High-pass filter tensor. default: tensor([[-1,-1,-1],[-1,8,-1],[-1,-1,-1]])
        window_size: Local window size integer. default: 8,
        reduction: Reduction method for output tensor. If ``None`` or ``"none"``,
                   returns a tensor with the per sample results. default: ``"mean"``.

    Return:
        Tensor with scc score

    Example:
        >>> from torch import randn
        >>> from torchmetrics.functional.image import spatial_correlation_coefficient as scc
        >>> x = randn(5, 3, 16, 16)
        >>> scc(x, x)
        tensor(1.)
        >>> x = randn(5, 16, 16)
        >>> scc(x, x)
        tensor(1.)
        >>> x = randn(5, 3, 16, 16)
        >>> y = randn(5, 3, 16, 16)
        >>> scc(x, y, reduction="none")
        tensor([0.0223, 0.0256, 0.0616, 0.0159, 0.0170])

    N)éÿÿÿÿrY   rY   )rY   é   rY   rW   )rV   rW   z3Expected reduction to be 'mean' or 'none', but got r   r)   )r   r   r   rV   Úelementwise_mean)rU   )r   r   r$   Úranger    rT   r   r!   rV   r.   r   )r   r   r   r   rU   ÚiÚper_channels          r#   Úspatial_correlation_coefficientr_   §   s  € ðF ÐÜšLª+²|ÐDÓEˆ	ØÐØˆ	ØÐ(Ñ(ÜÐNÈyÈkÐZÓ[Ð[Ü*¨5°&¸)À[ÓQÑ€Eˆ6�9ô �u—z‘z !“}Ó%ö	ð ô 	!Ø’!�Qšš1�*Ñ×'Ñ'¨Ó*¨F²1°aººA°:Ñ,>×,HÑ,HÈÓ,KÈYÐXcõ	
ð€Kð ð �FÒÜ�z‰zœ%Ÿ)™) K°QÔ7ºYÔGÐGØ�FÒÜ”e—i‘i °Ô3Ð?QÔRÐRØùòs   ÁAD)NrZ   rV   )r:   Útypingr   r   r!   r   r   Útorch.nn.functionalr   r   Útyping_extensionsr	   Útorchmetrics.utilities.checksr
   Ú"torchmetrics.utilities.distributedr   r,   Útupler$   r4   rA   rC   rK   rT   r_   © r%   r#   ú<module>rg      sm  ðó ß "ã ß  ß +Ý %å ;Ý 5ð/$�vð /$ vð /$¸&ð /$Èsð /$ÐW\Ð]cÐekÐmsÐ]sÑWtó /$ðd;¨ð ;°e¸CÀÀsÈCÀxÁÐ<PÑ6Qð ;ÐV\ó ;ð 	7 6ð 	7°6ð 	7¸fó 	7ð8 ð 8°ð 8¸6ó 8ð
3 fð 3°fð 3Àfð 3ÐQVÐW]Ð_eÐgmÐWmÑQnó 3ð*" Fð "°Fð "Àvð "Ð\_ð "Ðdjó "ðP #'ØØ9?ñ5Øð5àð5ð ˜Ñð5ð ð	5ð
 ˜Ð 4Ñ5Ñ6ð5ð ô5r%   