Ë
    þÍ:já  ã                   ó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Ú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 )	z1Compute the IoU matrix between two sets of boxes.é   éÿÿÿÿé   z-Expected preds to be of shape (N, 4) but got z.Expected target to be of shape (N, 4) but got r   )Úbox_iou)ÚdeviceÚdtype)
ÚndimÚshapeÚ
ValueErrorÚtorchvision.opsr   ÚnumelÚtorchÚzerosr   Úfloat32)r   r   r   r	   r   Úious         úz/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/detection/iou.pyÚ_iou_updater      s  € ð ‡z�z�Q‚˜%Ÿ+™+ b™/¨QÒ.ÜÐHÈÏÉÈÐVÓWÐWØ‡{�{�aÒ˜6Ÿ<™<¨Ñ+¨qÒ0ÜÐIÈ&Ï,É,ÈÐXÓYÐYå'à‡{�{ƒ}˜ÒÜ�{‰{˜6Ÿ<™<¨™?¨F¯L©L¸©OÀFÇMÁMÔY^×YfÑYfÔgÐgØ‡|�|ƒ~˜ÒÜ�{‰{˜5Ÿ;™; q™>¨5¯;©;°q©>À%Ç,Á,ÔV[×VcÑVcÔdÐdá
�%˜Ó
 €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   Ú_iou_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 Intersection over Union 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 iou is aggregated across all box pairs e.g. mean along the diagonal of the IoU matrix:

        >>> import torch
        >>> from torchmetrics.functional.detection import 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],
        ...     ]
        ... )
        >>> intersection_over_union(preds, target)
        tensor(0.5879)

    Example::
        By setting `aggregate=False` the full IoU matrix is returned:

        >>> import torch
        >>> from torchmetrics.functional.detection import 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],
        ...     ]
        ... )
        >>> intersection_over_union(preds, target, aggregate=False)
        tensor([[0.6898, 0.0000, 0.0000],
                [0.0000, 0.5086, 0.0000],
                [0.0000, 0.0000, 0.5654]])

    ú`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   4   sO   € õL "Ü!ØÔ'×0Ñ0Ð1ð 2Jð Jó
ð 	
ô �e˜V ]°OÓ
D€CÜ˜˜YÓ'Ð'r   )r   )T)Nr   T)Útypingr   r   Útorchmetrics.utilities.importsr   Ú__doctest_skip__ÚTensorÚfloatr   Úboolr#   r   © r   r   ú<module>r/      så   ðõ ã å AáØ1Ð2Ðð ijñØ�<‰<ðØ!&§¡ðØ>FÀu¹oðØ`eðà
‡\�\óñ,Z�e—l‘lð Z¨tð Z¸u¿|¹|ó Zð &*ØØñL(Ø�<‰<ðL(à�L‰LðL(ð ˜E‘?ðL(ð ð	L(ð
 ðL(ð ‡\�\ôL(r   