Ë
    þÍ:j2  ã                   ó8  — d dl mZ d dlZd dlmZ esdgZ	 ddej                  dej                  dee   ded	ej                  f
d
„Zddej                  de	d	ej                  fd„Z
	 	 	 ddej                  dej                  dee   dede	d	ej                  fd„Zy)é    )ÚOptionalN)Ú_TORCHVISION_AVAILABLEÚ distance_intersection_over_unionÚpredsÚtargetÚiou_thresholdÚreplacement_valÚreturnc                 ó”  — | j                   dk7  s| j                  d   dk7  rt        d| j                  › �«      ‚|j                   dk7  s|j                  d   dk7  rt        d|j                  › �«      ‚ddlm} | j                  «       dk(  rKt        j                  |j                  d   |j                  d   |j                  t        j                  ¬«      S |j                  «       dk(  rKt        j                  | j                  d   | j                  d   | j                  t        j                  ¬«      S  || |«      }|�||||k  <   |S )	Né   éÿÿÿÿé   z-Expected preds to be of shape (N, 4) but got z.Expected target to be of shape (N, 4) but got r   )Údistance_box_iou)ÚdeviceÚdtype)
ÚndimÚshapeÚ
ValueErrorÚtorchvision.opsr   ÚnumelÚtorchÚzerosr   Úfloat32)r   r   r   r	   r   Úious         ú{/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/detection/diou.pyÚ_diou_updater      s  € ð ‡z�z�Q‚˜%Ÿ+™+ b™/¨QÒ.ÜÐHÈÏÉÈÐVÓWÐWØ‡{�{�aÒ˜6Ÿ<™<¨Ñ+¨qÒ0ÜÐIÈ&Ï,É,ÈÐXÓYÐYå0à‡{�{ƒ}˜ÒÜ�{‰{˜6Ÿ<™<¨™?¨F¯L©L¸©OÀFÇMÁMÔY^×YfÑYfÔgÐgØ‡|�|ƒ~˜ÒÜ�{‰{˜5Ÿ;™; q™>¨5¯;©;°q©>À%Ç,Á,ÔV[×VcÑVcÔdÐdá
˜5 &Ó
)€CØÐ Ø#2ˆˆC�-ÑÑ Ø€Jó    r   Ú	aggregatec                 ó®   — |s| S | j                  «       dkD  r| j                  «       j                  «       S t        j                  d| j
                  ¬«      S )Nr   g        )r   )r   ÚdiagÚmeanr   Útensorr   )r   r   s     r   Ú_diou_computer#   -   s=   € ÙØˆ
Ø #§	¡	£¨a¢ˆ3�8‰8‹:�?‰?ÓÐY´U·\±\À#ÈcÏjÉjÔ5YÐYr   c                 ó|   — t         st        dt        j                  › d�«      ‚t	        | |||«      }t        ||«      S )a[	  Compute Distance Intersection over Union (`DIOU`_) between two sets of boxes.

    Both sets of boxes are expected to be in (x1, y1, x2, y2) format with 0 <= x1 < x2 and 0 <= y1 < y2.

    Args:
        preds:
            The input tensor containing the predicted bounding boxes.
        target:
            The tensor containing the ground truth.
        iou_threshold:
            Optional IoU thresholds for evaluation. If set to `None` the threshold is ignored.
        replacement_val:
            Value to replace values under the threshold with.
        aggregate:
            Return the average value instead of the full matrix of values

    Example::
        By default diou is aggregated across all box pairs e.g. mean along the diagonal of the dIoU matrix:

        >>> import torch
        >>> from torchmetrics.functional.detection import distance_intersection_over_union
        >>> preds = torch.tensor(
        ...     [
        ...         [296.55, 93.96, 314.97, 152.79],
        ...         [328.94, 97.05, 342.49, 122.98],
        ...         [356.62, 95.47, 372.33, 147.55],
        ...     ]
        ... )
        >>> target = torch.tensor(
        ...     [
        ...         [300.00, 100.00, 315.00, 150.00],
        ...         [330.00, 100.00, 350.00, 125.00],
        ...         [350.00, 100.00, 375.00, 150.00],
        ...     ]
        ... )
        >>> distance_intersection_over_union(preds, target)
        tensor(0.5793)

    Example::
        By setting `aggregate=False` the IoU score per prediction and target boxes is returned:

        >>> import torch
        >>> from torchmetrics.functional.detection import distance_intersection_over_union
        >>> preds = torch.tensor(
        ...     [
        ...         [296.55, 93.96, 314.97, 152.79],
        ...         [328.94, 97.05, 342.49, 122.98],
        ...         [356.62, 95.47, 372.33, 147.55],
        ...     ]
        ... )
        >>> target = torch.tensor(
        ...     [
        ...         [300.00, 100.00, 315.00, 150.00],
        ...         [330.00, 100.00, 350.00, 125.00],
        ...         [350.00, 100.00, 375.00, 150.00],
        ...     ]
        ... )
        >>> distance_intersection_over_union(preds, target, aggregate=False)
        tensor([[ 0.6883, -0.2043, -0.3351],
                [-0.2214,  0.4886, -0.1913],
                [-0.3971, -0.1510,  0.5609]])

    ú`zf` requires that `torchvision` is installed. Please install with `pip install torchmetrics[detection]`.)r   ÚModuleNotFoundErrorr   Ú__name__r   r#   )r   r   r   r	   r   r   s         r   r   r   3   sO   € õL "Ü!ØÔ0×9Ñ9Ð:ð ;Jð Jó
ð 	
ô �u˜f m°_Ó
E€CÜ˜˜iÓ(Ð(r   )r   )T)Nr   T)Útypingr   r   Útorchmetrics.utilities.importsr   Ú__doctest_skip__ÚTensorÚfloatr   Úboolr#   r   © r   r   ú<module>r/      så   ðõ ã å AáØ:Ð;Ðð ijñØ�<‰<ðØ!&§¡ðØ>FÀu¹oðØ`eðà
‡\�\óñ*Z�u—|‘|ð Z°ð ZÀÇÁó Zð &*ØØñL)Ø�<‰<ðL)à�L‰LðL)ð ˜E‘?ðL)ð ð	L)ð
 ðL)ð ‡\�\ôL)r   