Ë
    þÍ:jF  ã                   ó\   — d dl Z d dl mZ deddfd„Zdedeeef   fd„Zdedeeef   fd„Zy)é    N)ÚTensorÚimgÚreturnc                 ó¢   — t        | t        «      st        dt        | «      › �«      ‚| j                  dk7  rt        d| j                  › d�«      ‚y)z*Validate whether img is a 4D torch Tensor.z3The `img` expects a value of <Tensor> type but got é   z&The `img` expects a 4D tensor but got zD tensorN)Ú
isinstancer   Ú	TypeErrorÚtypeÚndimÚRuntimeError©r   s    ú|/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/image/gradients.pyÚ_image_gradients_validater      sM   € ä�cœ6Ô"ÜÐMÌdÐSVËiÈ[ÐYÓZÐZØ
‡x�x�1‚}ÜÐCÀCÇHÁHÀ:ÈXÐVÓWÐWð ó    c                 ó   — | j                   \  }}}}| ddd…dd…f   | ddd…dd…f   z
  }| ddd…dd…f   | ddd…dd…f   z
  }||d|g}t        j                  |t        j                  || j                  | j
                  ¬«      gd¬«      }|j                  | j                   «      }|||dg}t        j                  |t        j                  || j                  | j
                  ¬«      gd¬«      }|j                  | j                   «      }||fS )	z2Compute image gradients (dy/dx) for a given image..é   Néÿÿÿÿ)ÚdeviceÚdtypeé   )Údimé   )ÚshapeÚtorchÚcatÚzerosr   r   Úview)	r   Ú
batch_sizeÚchannelsÚheightÚwidthÚdyÚdxÚshapeyÚshapexs	            r   Ú_compute_image_gradientsr&      sý   € à*-¯)©)Ñ'€J�˜& %à	ˆS�!‘"’aˆZ‰˜3˜s C R Cª˜{Ñ+Ñ	+€BØ	ˆS’!�Q‘RˆZ‰˜3˜s¢A s¨ s˜{Ñ+Ñ	+€Bà˜( A uÐ-€FÜ	�‰�BœŸ™ F°3·:±:ÀSÇYÁYÔOÐPÐVWÔ	X€BØ	�‰�—‘Ó	€Bà˜( F¨AÐ.€FÜ	�‰�BœŸ™ F°3·:±:ÀSÇYÁYÔOÐPÐVWÔ	X€BØ	�‰�—‘Ó	€Bàˆrˆ6€Mr   c                 ó.   — t        | «       t        | «      S )a„  Compute `Gradient Computation of Image`_ of a given image using finite difference.

    Args:
        img: An ``(N, C, H, W)`` input tensor where ``C`` is the number of image channels

    Return:
        Tuple of ``(dy, dx)`` with each gradient of shape ``[N, C, H, W]``

    Raises:
        TypeError:
            If ``img`` is not of the type :class:`~torch.Tensor`.
        RuntimeError:
            If ``img`` is not a 4D tensor.

    Example:
        >>> from torchmetrics.functional.image import image_gradients
        >>> image = torch.arange(0, 1*1*5*5, dtype=torch.float32)
        >>> image = torch.reshape(image, (1, 1, 5, 5))
        >>> dy, dx = image_gradients(image)
        >>> dy[0, 0, :, :]
        tensor([[5., 5., 5., 5., 5.],
                [5., 5., 5., 5., 5.],
                [5., 5., 5., 5., 5.],
                [5., 5., 5., 5., 5.],
                [0., 0., 0., 0., 0.]])

    .. note::
           The implementation follows the 1-step finite difference method as followed
           by the TF implementation. The values are organized such that the gradient of
           [I(x+1, y)-[I(x, y)]] are at the (x, y) location

    )r   r&   r   s    r   Úimage_gradientsr(   -   s   € ôB ˜cÔ"ä# CÓ(Ð(r   )r   r   r   Útupler&   r(   © r   r   ú<module>r+      s[   ðó Ý ðX 6ð X¨dó Xð &ð ¨U°6¸6°>Ñ-Bó ð$#)˜ð #) E¨&°&¨.Ñ$9ô #)r   