Ë
    þÍ:j¸  ã                   óf   — d dl mZ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  G d„ de«      Zy)	é    )ÚAnyÚOptionalN)ÚTensorÚtensor)Ú_scc_per_channel_compute)Ú_scc_update)ÚMetricc            	       ó~   ‡ — e Zd ZU dZdZdZdZeed<   eed<   dde	e   de
d	ed
dfˆ fd„Zdeded
dfd„Zd
efd„Zˆ xZS )ÚSpatialCorrelationCoefficienta  Compute Spatial Correlation Coefficient (SCC_).

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

    - ``preds`` (:class:`~torch.Tensor`): Predictions from model of shape ``(N,C,H,W)`` or ``(N,H,W)``.
    - ``target`` (:class:`~torch.Tensor`): Ground truth values of shape ``(N,C,H,W)`` or ``(N,H,W)``.

    As output of `forward` and `compute` the metric returns the following output

    - ``scc`` (:class:`~torch.Tensor`): Tensor with scc score

    Args:
        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.
        kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info.

    Example:
        >>> from torch import randn
        >>> from torchmetrics.image import SpatialCorrelationCoefficient as SCC
        >>> preds = randn([32, 3, 64, 64])
        >>> target = randn([32, 3, 64, 64])
        >>> scc = SCC()
        >>> scc(preds, target)
        tensor(0.0023)

    TFÚ	scc_scoreÚtotalNÚhigh_pass_filterÚwindow_sizeÚkwargsÚreturnc                 óà   •— t        ‰| �  di |¤Ž |€t        g d¢g d¢g d¢g«      }|| _        || _        | j                  dt        d«      d¬«       | j                  dt        d«      d¬«       y )	N)éÿÿÿÿr   r   )r   é   r   r   g        Úsum)ÚdefaultÚdist_reduce_fxr   © )ÚsuperÚ__init__r   Ú	hp_filterÚwsÚ	add_state)Úselfr   r   r   Ú	__class__s       €úk/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/image/scc.pyr   z&SpatialCorrelationCoefficient.__init__;   se   ø€ Ü‰ÑÑ"˜6Ò"àÐ#Ü%¢|²[Â,Ð&OÓPÐà)ˆŒØˆŒà�‰�{¬F°3«KÈˆÔNØ�‰�w¬¨s«ÀEˆÕJó    ÚpredsÚtargetc                 ó4  — t        ||| j                  | j                  «      \  }}}t        |j	                  d«      «      D �cg c]R  }t        |dd…|dd…dd…f   j                  d«      |dd…|dd…dd…f   j                  d«      || j                  «      ‘ŒT }}| xj                  t        j                  t        j                  t        j                  |d¬«      g d¢¬«      «      z  c_        | xj                  |j	                  d«      z  c_        yc c}w )z*Update state with predictions and targets.é   N)Údim)r%   é   é   r   )r   r   r   ÚrangeÚsizeÚ_scc_computeÚ	unsqueezer   Útorchr   ÚmeanÚcatr   )r   r"   r#   r   ÚiÚscc_per_channels         r    Úupdatez$SpatialCorrelationCoefficient.updateG   sÜ   € ä#.¨u°f¸d¿n¹nÈdÏgÉgÓ#VÑ ˆˆv�yô ˜5Ÿ:™: a›=Ó)ö
àô ˜šq !¢Qª˜zÑ*×4Ñ4°QÓ7¸ÂÀ1ÂaÊÀ
Ñ9K×9UÑ9UÐVWÓ9XÐZcÐei×elÑelÕmð
ˆð 
ð 	�Šœ%Ÿ)™)¤E§J¡J¬u¯y©y¸ÈaÔ/PÒV_Ô$`ÓaÑa�Ø�
Š
�e—j‘j “mÑ#Ž
ùò
s   ÁADc                 ó4   — | j                   | j                  z  S )zICompute the VIF score based on inputs passed in to ``update`` previously.)r   r   )r   s    r    Úcomputez%SpatialCorrelationCoefficient.computeQ   s   € à�~‰~ §
¡
Ñ*Ð*r!   )Nr   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__Úis_differentiableÚhigher_is_betterÚfull_state_updater   Ú__annotations__r   Úintr   r   r2   r4   Ú__classcell__)r   s   @r    r   r      sz   ø… ñð6 ÐØÐØÐàÓØƒMñ
K¨°&Ñ)9ð 
KÈsð 
KÐbeð 
KÐjnõ 
Kð$˜Fð $¨Fð $°tó $ð+˜÷ +r!   r   )Útypingr   r   r-   r   r   Ú!torchmetrics.functional.image.sccr   r+   r   Útorchmetrics.metricr	   r   r   r!   r    ú<module>rB      s%   ð÷ !ã ß  å VÝ 9Ý &ô;+ Fõ ;+r!   