Ë
    Fêñi&`  ã                  óâ   — d dl mZ d dlZd dlmZ ddlmZ ddlmZm	Z	 ddl
mZmZmZ ddlmZ  G d„ d	ej                   «      Z G d
„ de«      Zdd„Zdd„Zddd„Zdd„Z	 	 d	 	 	 	 	 	 	 	 	 dd„Zy)é    )ÚannotationsNé   )ÚLOGGER)Úbbox_iouÚprobiou)Ú	xywh2xyxyÚxywhr2xyxyxyxyÚ	xyxy2xywh)Ú
TORCH_1_11c                  ó°   ‡ — e Zd ZdZddddg d¢ddf	 	 	 	 	 	 	 	 	 	 	 dˆ fd	„Z ej                  «       d
„ «       Zd„ Zd„ Z	d„ Z
d„ Zdd„Zd„ Zdd„Zd„ Zˆ xZS )ÚTaskAlignedAssigneraG  A task-aligned assigner for object detection.

    This class assigns ground-truth (gt) objects to anchors based on the task-aligned metric, which combines both
    classification and localization information.

    Attributes:
        topk (int): The number of top candidates to consider.
        topk2 (int): Secondary topk value for additional filtering.
        num_classes (int): The number of object classes.
        alpha (float): The alpha parameter for the classification component of the task-aligned metric.
        beta (float): The beta parameter for the localization component of the task-aligned metric.
        stride (list): List of stride values for different feature levels.
        stride_val (int): The stride value used for select_candidates_in_gts.
        eps (float): A small value to prevent division by zero.
    é   éP   ç      ð?g      @)é   é   é    ç•Ö&è.>Nc                ó  •— t         ‰| �  «        || _        |xs || _        || _        || _        || _        || _        t        | j                  «      dkD  r| j                  d   n| j                  d   | _	        || _
        y)aÖ  Initialize a TaskAlignedAssigner object with customizable hyperparameters.

        Args:
            topk (int, optional): The number of top candidates to consider.
            num_classes (int, optional): The number of object classes.
            alpha (float, optional): The alpha parameter for the classification component of the task-aligned metric.
            beta (float, optional): The beta parameter for the localization component of the task-aligned metric.
            stride (list, optional): List of stride values for different feature levels.
            eps (float, optional): A small value to prevent division by zero.
            topk2 (int, optional): Secondary topk value for additional filtering.
        r   r   N)ÚsuperÚ__init__ÚtopkÚtopk2Únum_classesÚalphaÚbetaÚstrideÚlenÚ
stride_valÚeps)	Úselfr   r   r   r   r   r    r   Ú	__class__s	           €úW/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/utils/tal.pyr   zTaskAlignedAssigner.__init__   sp   ø€ ô* 	‰ÑÔØˆŒ	Ø’]˜dˆŒ
Ø&ˆÔØˆŒ
ØˆŒ	ØˆŒÜ,/°·±Ó,<¸qÒ,@˜$Ÿ+™+ aš.ÀdÇkÁkÐRSÁnˆŒØˆ�ó    c                óÒ  ‡— |j                   d   | _        |j                   d   | _        |j                  Š| j                  dk(  rzt	        j
                  |d   | j                  «      t	        j                  |«      t	        j                  |«      t	        j                  |d   «      t	        j                  |d   «      fS 	 | j                  ||||||«      S # t        $ r‡}dt        |«      j                  «       v rft        j                  d«       ||||||fD �cg c]  }|j                  «       ‘Œ nc c}w }	} | j                  |	Ž }
t        ˆfd„|
D «       «      cY d}~S ‚ d}~ww xY w)a  Compute the task-aligned assignment.

        Args:
            pd_scores (torch.Tensor): Predicted classification scores with shape (bs, num_total_anchors, num_classes).
            pd_bboxes (torch.Tensor): Predicted bounding boxes with shape (bs, num_total_anchors, 4).
            anc_points (torch.Tensor): Anchor points with shape (num_total_anchors, 2).
            gt_labels (torch.Tensor): Ground truth labels with shape (bs, n_max_boxes, 1).
            gt_bboxes (torch.Tensor): Ground truth boxes with shape (bs, n_max_boxes, 4).
            mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (bs, n_max_boxes, 1).

        Returns:
            target_labels (torch.Tensor): Target labels with shape (bs, num_total_anchors).
            target_bboxes (torch.Tensor): Target bounding boxes with shape (bs, num_total_anchors, 4).
            target_scores (torch.Tensor): Target scores with shape (bs, num_total_anchors, num_classes).
            fg_mask (torch.Tensor): Foreground mask with shape (bs, num_total_anchors).
            target_gt_idx (torch.Tensor): Target ground truth indices with shape (bs, num_total_anchors).

        References:
            https://github.com/Nioolek/PPYOLOE_pytorch/blob/master/ppyoloe/assigner/tal_assigner.py
        r   r   ).r   zout of memoryz7CUDA OutOfMemoryError in TaskAlignedAssigner, using CPUc              3  ó@   •K  — | ]  }|j                  ‰«      –— Œ y ­w©N)Úto)Ú.0ÚtÚdevices     €r#   ú	<genexpr>z.TaskAlignedAssigner.forward.<locals>.<genexpr>i   s   øè ø€ Ò:¨a˜QŸT™T &Ÿ\Ñ:ùs   ƒN)ÚshapeÚbsÚn_max_boxesr+   ÚtorchÚ	full_liker   Ú
zeros_likeÚ_forwardÚRuntimeErrorÚstrÚlowerr   ÚwarningÚcpuÚtuple)r!   Ú	pd_scoresÚ	pd_bboxesÚ
anc_pointsÚ	gt_labelsÚ	gt_bboxesÚmask_gtÚer*   Úcpu_tensorsÚresultr+   s              @r#   ÚforwardzTaskAlignedAssigner.forward>   sB  ø€ ð, —/‘/ !Ñ$ˆŒØ$Ÿ?™?¨1Ñ-ˆÔØ×!Ñ!ˆà×Ñ˜qÒ ä—‘ 	¨&Ñ 1°4×3CÑ3CÓDÜ× Ñ  Ó+Ü× Ñ  Ó+Ü× Ñ  ¨6Ñ!2Ó3Ü× Ñ  ¨6Ñ!2Ó3ðð ð		Ø—=‘= ¨I°zÀ9ÈiÐY`ÓaÐaøÜò 	Ø¤# a£&§,¡,£.Ñ0ä—‘ÐXÔYØ1:¸IÀzÐS\Ð^gÐipÐ0qÖr¨1˜qŸu™u�wÑrùÒr�ÐrØ&˜Ÿ™¨Ð4�ÜÓ:°6Ô:Ó:Õ:Øûð	ús0   Ã C Ã	E&Ã:E!ÄD1Ä0*E!ÅE&Å E!Å!E&c                ó   — | j                  ||||||«      \  }}}	| j                  ||	| j                  |«      \  }
}}| j                  |||
|«      \  }}}||z  }|j	                  dd¬«      }|	|z  j	                  dd¬«      }||z  || j
                  z   z  j	                  d«      j                  d«      }||z  }||||j                  «       |
fS )a�  Compute the task-aligned assignment.

        Args:
            pd_scores (torch.Tensor): Predicted classification scores with shape (bs, num_total_anchors, num_classes).
            pd_bboxes (torch.Tensor): Predicted bounding boxes with shape (bs, num_total_anchors, 4).
            anc_points (torch.Tensor): Anchor points with shape (num_total_anchors, 2).
            gt_labels (torch.Tensor): Ground truth labels with shape (bs, n_max_boxes, 1).
            gt_bboxes (torch.Tensor): Ground truth boxes with shape (bs, n_max_boxes, 4).
            mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (bs, n_max_boxes, 1).

        Returns:
            target_labels (torch.Tensor): Target labels with shape (bs, num_total_anchors).
            target_bboxes (torch.Tensor): Target bounding boxes with shape (bs, num_total_anchors, 4).
            target_scores (torch.Tensor): Target scores with shape (bs, num_total_anchors, num_classes).
            fg_mask (torch.Tensor): Foreground mask with shape (bs, num_total_anchors).
            target_gt_idx (torch.Tensor): Target ground truth indices with shape (bs, num_total_anchors).
        éÿÿÿÿT)ÚdimÚkeepdiméþÿÿÿ)Úget_pos_maskÚselect_highest_overlapsr/   Úget_targetsÚamaxr    Ú	unsqueezeÚbool)r!   r:   r;   r<   r=   r>   r?   Úmask_posÚalign_metricÚoverlapsÚtarget_gt_idxÚfg_maskÚtarget_labelsÚtarget_bboxesÚtarget_scoresÚpos_align_metricsÚpos_overlapsÚnorm_align_metrics                     r#   r3   zTaskAlignedAssigner._forwardl   s  € ð$ ,0×+<Ñ+<Ø�y )¨Y¸
ÀGó,
Ñ(ˆ�, ð ,0×+GÑ+GØ�h × 0Ñ 0°,ó,
Ñ(ˆ�w ð
 7;×6FÑ6FÀyÐR[Ð]jÐlsÓ6tÑ3ˆ�} mð 	˜Ñ ˆØ(×-Ñ-°"¸dÐ-ÓCÐØ  8Ñ+×1Ñ1°bÀ$Ð1ÓGˆØ)¨LÑ8Ð<MÐPT×PXÑPXÑ<XÑY×_Ñ_Ð`bÓc×mÑmÐnpÓqÐØ%Ð(9Ñ9ˆà˜m¨]¸G¿L¹L»NÈMÐYÐYr$   c                óð   — | j                  |||«      }| j                  ||||||z  «      \  }}	| j                  ||j                  dd| j                  «      j                  «       ¬«      }
|
|z  |z  }|||	fS )aÓ  Get positive mask for each ground truth box.

        Args:
            pd_scores (torch.Tensor): Predicted classification scores with shape (bs, num_total_anchors, num_classes).
            pd_bboxes (torch.Tensor): Predicted bounding boxes with shape (bs, num_total_anchors, 4).
            gt_labels (torch.Tensor): Ground truth labels with shape (bs, n_max_boxes, 1).
            gt_bboxes (torch.Tensor): Ground truth boxes with shape (bs, n_max_boxes, 4).
            anc_points (torch.Tensor): Anchor points with shape (num_total_anchors, 2).
            mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (bs, n_max_boxes, 1).

        Returns:
            mask_pos (torch.Tensor): Positive mask with shape (bs, max_num_obj, h*w).
            align_metric (torch.Tensor): Alignment metric with shape (bs, max_num_obj, h*w).
            overlaps (torch.Tensor): Overlaps between predicted vs ground truth boxes with shape (bs, max_num_obj, h*w).
        rE   )Ú	topk_mask)Úselect_candidates_in_gtsÚget_box_metricsÚselect_topk_candidatesÚexpandr   rN   )r!   r:   r;   r=   r>   r<   r?   Úmask_in_gtsrP   rQ   Ú	mask_topkrO   s               r#   rI   z TaskAlignedAssigner.get_pos_mask’   s�   € ð  ×3Ñ3°JÀ	È7ÓSˆà!%×!5Ñ!5°iÀÈIÐW`ÐbmÐpwÑbwÓ!xÑˆ�hà×/Ñ/°ÈÏÉÐWYÐ[]Ð_c×_hÑ_hÓHi×HnÑHnÓHpÐ/Óqˆ	à˜{Ñ*¨WÑ4ˆà˜ xÐ/Ð/r$   c                óþ  — |j                   d   }|j                  «       }t        j                  | j                  | j
                  |g|j                  |j                  ¬«      }t        j                  | j                  | j
                  |g|j                  |j                  ¬«      }t        j                  d| j                  | j
                  gt        j                  ¬«      }	t        j                  | j                  ¬«      j                  dd«      j                  d| j
                  «      |	d<   |j                  d«      |	d<   ||	d   d	d	…|	d   f   |   ||<   |j                  d«      j                  d| j
                  dd«      |   }
|j                  d«      j                  dd|d«      |   }| j                  ||
«      ||<   |j                  | j                   «      |j                  | j"                  «      z  }||fS )
a/  Compute alignment metric given predicted and ground truth bounding boxes.

        Args:
            pd_scores (torch.Tensor): Predicted classification scores with shape (bs, num_total_anchors, num_classes).
            pd_bboxes (torch.Tensor): Predicted bounding boxes with shape (bs, num_total_anchors, 4).
            gt_labels (torch.Tensor): Ground truth labels with shape (bs, n_max_boxes, 1).
            gt_bboxes (torch.Tensor): Ground truth boxes with shape (bs, n_max_boxes, 4).
            mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (bs, n_max_boxes, h*w).

        Returns:
            align_metric (torch.Tensor): Alignment metric combining classification and localization.
            overlaps (torch.Tensor): IoU overlaps between predicted and ground truth boxes.
        rH   ©Údtyper+   é   )rd   )ÚendrE   r   r   N)r-   rN   r0   Úzerosr.   r/   rd   r+   ÚlongÚarangeÚviewr_   ÚsqueezerM   Úiou_calculationÚpowr   r   )r!   r:   r;   r=   r>   r?   ÚnarQ   Úbbox_scoresÚindÚpd_boxesÚgt_boxesrP   s                r#   r]   z#TaskAlignedAssigner.get_box_metrics¬   s¥  € ð �_‰_˜RÑ ˆØ—,‘,“.ˆÜ—;‘; §¡¨×)9Ñ)9¸2Ð>ÀiÇoÁoÐ^g×^nÑ^nÔoˆÜ—k‘k 4§7¡7¨D×,<Ñ,<¸bÐ"AÈÏÉÐaj×aqÑaqÔrˆä�k‰k˜1˜dŸg™g t×'7Ñ'7Ð8ÄÇ
Á
ÔKˆÜ—‘ $§'¡'Ô*×/Ñ/°°AÓ6×=Ñ=¸bÀ$×BRÑBRÓSˆˆA‰Ø×"Ñ" 2Ó&ˆˆA‰à(¨¨Q©²°C¸±FÐ):Ñ;¸GÑDˆ�GÑð ×&Ñ& qÓ)×0Ñ0°°T×5EÑ5EÀrÈ2ÓNÈwÑWˆØ×&Ñ& qÓ)×0Ñ0°°R¸¸RÓ@ÀÑIˆØ ×0Ñ0°¸8ÓDˆ�Ñà"—‘ t§z¡zÓ2°X·\±\À$Ç)Á)Ó5LÑLˆØ˜XÐ%Ð%r$   c                ó\   — t        ||dd¬«      j                  d«      j                  d«      S )a
  Calculate IoU for horizontal bounding boxes.

        Args:
            gt_bboxes (torch.Tensor): Ground truth boxes.
            pd_bboxes (torch.Tensor): Predicted boxes.

        Returns:
            (torch.Tensor): IoU values between each pair of boxes.
        FT)ÚxywhÚCIoUrE   r   )r   rk   Úclamp_©r!   r>   r;   s      r#   rl   z#TaskAlignedAssigner.iou_calculationÍ   s,   € ô ˜	 9°5¸tÔD×LÑLÈRÓP×WÑWÐXYÓZÐZr$   c           
     ó   — t        j                  || j                  dd¬«      \  }}|€2|j                  dd¬«      d   | j                  kD  j	                  |«      }|j                  | d«       t        j                  |j                  t         j                  |j                  ¬«      }t        j                  |dd…dd…dd…f   t         j                  |j                  ¬«      }t        | j                  «      D ]$  }|j                  d|dd…dd…||dz   …f   |«       Œ& |j                  |dkD  d«       |j                  |j                  «      S )	aÈ  Select the top-k candidates based on the given metrics.

        Args:
            metrics (torch.Tensor): A tensor of shape (b, max_num_obj, h*w), where b is the batch size, max_num_obj is
                the maximum number of objects, and h*w represents the total number of anchor points.
            topk_mask (torch.Tensor, optional): An optional boolean tensor of shape (b, max_num_obj, topk), where topk
                is the number of top candidates to consider. If not provided, the top-k values are automatically
                computed based on the given metrics.

        Returns:
            (torch.Tensor): A tensor of shape (b, max_num_obj, h*w) containing the selected top-k candidates.
        rE   T©rF   ÚlargestN)rG   r   rc   r   )r0   r   Úmaxr    Ú	expand_asÚmasked_fill_rg   r-   Úint8r+   Ú	ones_likeÚrangeÚscatter_add_r(   rd   )r!   Úmetricsr[   Útopk_metricsÚ	topk_idxsÚcount_tensorÚonesÚks           r#   r^   z*TaskAlignedAssigner.select_topk_candidatesÙ   s  € ô #(§*¡*¨W°d·i±iÀRÐQUÔ"VÑˆ�iØÐØ%×)Ñ)¨"°dÐ)Ó;¸AÑ>ÀÇÁÑI×TÑTÐU^Ó_ˆIà×Ñ 	˜z¨1Ô-ô —{‘{ 7§=¡=¼¿
¹
È9×K[ÑK[Ô\ˆÜ�‰˜yªªA¨r°¨r¨Ñ2¼%¿*¹*ÈY×M]ÑM]Ô^ˆÜ�t—y‘yÓ!ò 	LˆAà×%Ñ% b¨)²A²q¸!¸aÀ!¹e¸)°OÑ*DÀdÕKð	Lð 	×!Ñ! ,°Ñ"2°AÔ6à�‰˜wŸ}™}Ó-Ð-r$   c                óÆ  — t        j                  | j                  t         j                  |j                  ¬«      d   }||| j
                  z  z   }|j                  «       j                  «       |   }|j                  d|j                  d   «      |   }|j                  d«       t        j                  |j                  d   |j                  d   | j                  ft         j                  |j                  ¬«      }|j                  d|j                  d«      d«       |dd…dd…df   j                  dd| j                  «      }	t        j                   |	dkD  |d«      }|||fS )	a@  Compute target labels, target bounding boxes, and target scores for the positive anchor points.

        Args:
            gt_labels (torch.Tensor): Ground truth labels of shape (b, max_num_obj, 1), where b is the batch size and
                max_num_obj is the maximum number of objects.
            gt_bboxes (torch.Tensor): Ground truth bounding boxes of shape (b, max_num_obj, 4).
            target_gt_idx (torch.Tensor): Indices of the assigned ground truth objects for positive anchor points, with
                shape (b, h*w), where h*w is the total number of anchor points.
            fg_mask (torch.Tensor): A boolean tensor of shape (b, h*w) indicating the positive (foreground) anchor
                points.

        Returns:
            target_labels (torch.Tensor): Target labels for positive anchor points with shape (b, h*w).
            target_bboxes (torch.Tensor): Target bounding boxes for positive anchor points with shape (b, h*w, 4).
            target_scores (torch.Tensor): Target scores for positive anchor points with shape (b, h*w, num_classes).
        )rf   rd   r+   ).NrE   r   r   rc   re   N)r0   ri   r.   Úint64r+   r/   rh   Úflattenrj   r-   rv   rg   r   Úscatter_rM   ÚrepeatÚwhere)
r!   r=   r>   rR   rS   Ú	batch_indrT   rU   rV   Úfg_scores_masks
             r#   rK   zTaskAlignedAssigner.get_targetsø   s7  € ô$ —L‘L T§W¡W´E·K±KÈ	×HXÑHXÔYÐZcÑdˆ	Ø%¨	°D×4DÑ4DÑ(DÑDˆØ!Ÿ™Ó(×0Ñ0Ó2°=ÑAˆð "Ÿ™ r¨9¯?©?¸2Ñ+>Ó?ÀÑNˆð 	×Ñ˜QÔô Ÿ™Ø× Ñ  Ñ# ]×%8Ñ%8¸Ñ%;¸T×=MÑ=MÐNÜ—+‘+Ø ×'Ñ'ô
ˆð
 	×Ñ˜q -×"9Ñ"9¸"Ó"=¸qÔAà ¢¢A t Ñ,×3Ñ3°A°q¸$×:JÑ:JÓKˆÜŸ™ N°QÑ$6¸ÀqÓIˆà˜m¨]Ð:Ð:r$   c                ól  — t        |«      }|ddd…f   | j                  d   k  }t        j                  ||z  j	                  «       t        j
                  | j                  |j                  |j                  ¬«      |ddd…f   «      |ddd…f<   t        |«      }|j                  d   }|j                  \  }}	}
|j                  ddd«      j                  dd«      \  }}t        j                  |d   |z
  ||d   z
  fd¬	«      j                  ||	|d«      }|j                  d
«      j                  |«      S )a¿  Select positive anchor centers within ground truth bounding boxes.

        Args:
            xy_centers (torch.Tensor): Anchor center coordinates, shape (h*w, 2).
            gt_bboxes (torch.Tensor): Ground truth bounding boxes, shape (b, n_boxes, 4).
            mask_gt (torch.Tensor): Mask for valid ground truth boxes, shape (b, n_boxes, 1).
            eps (float, optional): Small value for numerical stability.

        Returns:
            (torch.Tensor): Boolean mask of positive anchors, shape (b, n_boxes, h*w).

        Notes:
            - b: batch size, n_boxes: number of ground truth boxes, h: height, w: width.
            - Bounding box format: [x_min, y_min, x_max, y_max].
        .re   Nr   rc   rE   r   é   ©rF   é   )r
   r   r0   r�   rN   Útensorr   rd   r+   r   r-   rj   ÚchunkÚcatÚaminÚgt_)r!   Ú
xy_centersr>   r?   r    Úgt_bboxes_xywhÚwh_maskÚ	n_anchorsr.   Ún_boxesÚ_ÚltÚrbÚbbox_deltass                 r#   r\   z,TaskAlignedAssigner.select_candidates_in_gts!  s,  € ô  # 9Ó-ˆØ   a¡b Ñ)¨D¯K©K¸©NÑ:ˆÜ"'§+¡+Ø�wÑ×$Ñ$Ó&Ü�L‰L˜Ÿ™°×0DÑ0DÈ^×MbÑMbÔcØ˜3 ¡˜7Ñ#ó#
ˆ�s˜A™B�wÑô
 ˜nÓ-ˆ	à×$Ñ$ QÑ'ˆ	Ø"Ÿ™‰ˆˆG�QØ—‘  A qÓ)×/Ñ/°°1Ó5‰ˆˆBÜ—i‘i ¨DÑ!1°BÑ!6¸¸ZÈÑ=MÑ8MÐ NÐTUÔV×[Ñ[Ð\^Ð`gÐirÐtvÓwˆØ×Ñ Ó"×&Ñ& sÓ+Ð+r$   c                óR  — |j                  d«      }|j                  «       dkD  rÄ|j                  d«      dkD  j                  d|d«      }|j	                  d«      }t        j                  |j                  |j                  |j                  ¬«      }|j                  d|j                  d«      d«       t        j                  |||«      j                  «       }|j                  d«      }| j                  | j                  k7  r‘||z  }t        j                  || j                  dd¬«      j                  }t        j                  |j                  |j                  |j                  ¬«      }	|	j                  d|d«       ||	z  }|j                  d«      }|j	                  d«      }
|
||fS )a®  Select anchor boxes with highest IoU when assigned to multiple ground truths.

        Args:
            mask_pos (torch.Tensor): Positive mask, shape (b, n_max_boxes, h*w).
            overlaps (torch.Tensor): IoU overlaps, shape (b, n_max_boxes, h*w).
            n_max_boxes (int): Maximum number of ground truth boxes.
            align_metric (torch.Tensor): Alignment metric for selecting best matches.

        Returns:
            target_gt_idx (torch.Tensor): Indices of assigned ground truths, shape (b, h*w).
            fg_mask (torch.Tensor): Foreground mask, shape (b, h*w).
            mask_pos (torch.Tensor): Updated positive mask, shape (b, n_max_boxes, h*w).
        rH   r   rE   rc   Try   r   )Úsumr{   rM   r_   Úargmaxr0   rg   r-   rd   r+   r‹   r�   Úfloatr   r   Úindices)r!   rO   rQ   r/   rP   rS   Úmask_multi_gtsÚmax_overlaps_idxÚis_max_overlapsÚtopk_idxrR   s              r#   rJ   z+TaskAlignedAssigner.select_highest_overlaps@  s]  € ð —,‘,˜rÓ"ˆØ�;‰;‹=˜1ÒØ%×/Ñ/°Ó2°QÑ6×>Ñ>¸rÀ;ÐPRÓSˆNà'Ÿ™¨qÓ1ÐÜ#Ÿk™k¨(¯.©.ÀÇÁÐW_×WfÑWfÔgˆOØ×$Ñ$ QÐ(8×(BÑ(BÀ1Ó(EÀqÔIÜ—{‘{ >°?ÀHÓM×SÑSÓUˆHà—l‘l 2Ó&ˆGà�:‰:˜Ÿ™Ò"Ø'¨(Ñ2ˆLÜ$Ÿz™z¨,¸¿
¹
ÈÐTXÔY×aÑaÐÜ—{‘{ 8§>¡>¸¿¹ÐPX×P_ÑP_Ô`ˆHØ×Ñ˜bÐ"2°CÔ8Ø˜Ñ ˆHØ—l‘l 2Ó&ˆGà Ÿ™¨Ó+ˆØ˜g xÐ/Ð/r$   )r   Úintr   r«   r   r¥   r   r¥   r   Úlistr    r¥   r'   )r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r0   Úno_gradrC   r3   rI   r]   rl   r^   rK   r\   rJ   Ú__classcell__)r"   s   @r#   r   r      s¨   ø„ ñð$ ØØØÚ"ØØðàðð ðð ð	ð
 ðð ðð õð> €U‡]�]ƒ_ñ+ó ð+òZ$ZòL0ò4&òB
[ó.ò>';óR,ö>#0r$   r   c                  ó   — e Zd ZdZd„ Zd„ Zy)ÚRotatedTaskAlignedAssignerzSAssigns ground-truth objects to rotated bounding boxes using a task-aligned metric.c                óV   — t        ||«      j                  d«      j                  d«      S )z)Calculate IoU for rotated bounding boxes.rE   r   )r   rk   rv   rw   s      r#   rl   z*RotatedTaskAlignedAssigner.iou_calculationi  s%   € ä�y )Ó,×4Ñ4°RÓ8×?Ñ?ÀÓBÐBr$   c                ób  — |j                  «       }|ddd…f   | j                  d   k  }t        j                  ||z  j	                  «       t        j
                  | j                  |j                  |j                  ¬«      |ddd…f   «      |ddd…f<   t        |«      }|j                  dd¬«      \  }}}	}
||z
  }|
|z
  }||z
  }||z  j                  d	¬«      }||z  j                  d	¬«      }||z  j                  d	¬«      }||z  j                  d	¬«      }|dk\  ||k  z  |dk\  z  ||k  z  S )
aÞ  Select the positive anchor center in gt for rotated bounding boxes.

        Args:
            xy_centers (torch.Tensor): Anchor center coordinates with shape (h*w, 2).
            gt_bboxes (torch.Tensor): Ground truth bounding boxes with shape (b, n_boxes, 5).
            mask_gt (torch.Tensor): Mask for valid ground truth boxes with shape (b, n_boxes, 1).

        Returns:
            (torch.Tensor): Boolean mask of positive anchors with shape (b, n_boxes, h*w).
        .re   r‘   r   rc   r   rH   r’   rE   )Úcloner   r0   r�   rN   r”   r   rd   r+   r	   Úsplitr£   )r!   r™   r>   r?   Úgt_bboxes_cloner›   ÚcornersÚaÚbrž   ÚdÚabÚadÚapÚnorm_abÚnorm_adÚ	ap_dot_abÚ	ap_dot_ads                     r#   r\   z3RotatedTaskAlignedAssigner.select_candidates_in_gtsm  sF  € ð $Ÿ/™/Ó+ˆØ! # q¨ s (Ñ+¨d¯k©k¸!©nÑ<ˆÜ$)§K¡KØ�wÑ×$Ñ$Ó&Ü�L‰L˜Ÿ™°×0EÑ0EÈo×NdÑNdÔeØ˜C  1 ˜HÑ%ó%
ˆ˜˜Q˜q˜S˜Ñ!ô ! Ó1ˆà—]‘] 1¨"�]Ó-‰
ˆˆ1ˆa�Ø�‰UˆØ�‰Uˆð ˜!‰^ˆØ˜‘7—-‘- B�-Ó'ˆØ˜‘7—-‘- B�-Ó'ˆØ˜"‘W—M‘M b�MÓ)ˆ	Ø˜"‘W—M‘M b�MÓ)ˆ	Ø˜Q‘ 9°Ñ#7Ñ8¸IÈ¹NÑKÈyÐ\cÑOcÑdÐdr$   N)r­   r®   r¯   r°   rl   r\   © r$   r#   r´   r´   f  s   „ Ù]òCó er$   r´   c           	     ó  — g g }}| €J ‚| d   j                   | d   j                  }}t        t        | «      «      D �]   }||   }t	        | t
        «      r| |   j                  dd n!t        | |   d   «      t        | |   d   «      f\  }	}
t        j                  |
||¬«      |z   }t        j                  |	||¬«      |z   }t        rt        j                  ||d¬«      nt        j                  ||«      \  }}|j                  t        j                  ||fd«      j                  dd«      «       |j                  t        j                  |	|
z  df|||¬	«      «       �Œ# t        j                   |«      t        j                   |«      fS )
zGenerate anchors from features.Nr   re   r   )rf   r+   rd   Úij)ÚindexingrE   rc   )rd   r+   r€   r   Ú
isinstancer¬   r-   r«   r0   ri   r   ÚmeshgridÚappendÚstackrj   Úfullr–   )ÚfeatsÚstridesÚgrid_cell_offsetÚanchor_pointsÚstride_tensorrd   r+   Úir   ÚhÚwÚsxÚsys                r#   Úmake_anchorsrØ   �  sb  € à#% r�=€MØÐÐÐØ˜!‘H—N‘N E¨!¡H§O¡Oˆ6€EÜ”3�u“:Óó YˆØ˜‘ˆÜ%/°´tÔ%<ˆu�Q‰x�~‰~˜a˜bÑ!Ä3ÀuÈQÁxÐPQÁ{ÓCSÔUXÐY^Ð_`ÑYaÐbcÑYdÓUeÐBf‰ˆˆ1Ü�\‰\˜a¨°eÔ<Ð?OÑOˆÜ�\‰\˜a¨°eÔ<Ð?OÑOˆÝ:D”—‘  B°Õ6Ì%Ï.É.ÐY[Ð]_ÓJ`‰ˆˆBØ×ÑœUŸ[™[¨"¨b¨°2Ó6×;Ñ;¸BÀÓBÔCØ×ÑœUŸZ™Z¨¨Q©°¨
°FÀ%ÐPVÔWÖXðYô �9‰9�]Ó#¤U§Y¡Y¨}Ó%=Ð=Ð=r$   c                ó¾   — | j                  d|«      \  }}||z
  }||z   }|r%||z   dz  }||z
  }	t        j                  ||	g|«      S t        j                  ||f|«      S )z.Transform distance(ltrb) to box(xywh or xyxy).re   )r•   r0   r–   )
ÚdistancerÑ   rt   rF   rŸ   r    Úx1y1Úx2y2Úc_xyÚwhs
             r#   Ú	dist2bboxrß      sn   € à�^‰^˜A˜sÓ#�F€BˆØ˜2Ñ€DØ˜2Ñ€DÙØ�t‘˜qÑ ˆØ�D‰[ˆÜ�y‰y˜$ ˜ SÓ)Ð)Ü�9‰9�d˜D�\ 3Ó'Ð'r$   c                óš   — |j                  dd«      \  }}t        j                  | |z
  || z
  fd«      }|�|j                  d|dz
  «      }|S )z#Transform bbox(xyxy) to dist(ltrb).re   rE   r   ç{®Gáz„?)r•   r0   r–   rv   )rÑ   ÚbboxÚreg_maxrÛ   rÜ   Údists         r#   Ú	bbox2distrå   ¬  sT   € à—‘˜A˜rÓ"�J€Dˆ$Ü�9‰9�m dÑ*¨D°=Ñ,@ÐAÀ2ÓF€DØÐØ�{‰{˜1˜g¨™nÓ-ˆØ€Kr$   c                óV  — | j                  d|¬«      \  }}t        j                  |«      t        j                  |«      }}||z
  dz  j                  d|¬«      \  }}	||z  |	|z  z
  ||z  |	|z  z   }}
t        j                  |
|g|¬«      |z   }t        j                  |||z   g|¬«      S )aî  Decode predicted rotated bounding box coordinates from anchor points and distribution.

    Args:
        pred_dist (torch.Tensor): Predicted rotated distance with shape (bs, h*w, 4).
        pred_angle (torch.Tensor): Predicted angle with shape (bs, h*w, 1).
        anchor_points (torch.Tensor): Anchor points with shape (h*w, 2).
        dim (int, optional): Dimension along which to split.

    Returns:
        (torch.Tensor): Predicted rotated bounding boxes with shape (bs, h*w, 4).
    re   r’   r   )r¸   r0   ÚcosÚsinr–   )Ú	pred_distÚ
pred_anglerÑ   rF   rŸ   r    rç   rè   ÚxfÚyfÚxÚyÚxys                r#   Ú	dist2rboxrð   µ  s«   € ð �_‰_˜Q Cˆ_Ó(�F€BˆÜ�y‰y˜Ó$¤e§i¡i°
Ó&;ˆ€Cà�B‰w˜!‰m×"Ñ" 1¨#Ð"Ó.�F€BˆØ�‰8�b˜3‘hÑ  S¡¨2°©8Ñ 3€q€AÜ	�‰�A�q�6˜sÔ	# mÑ	3€BÜ�9‰9�b˜"˜r™'�]¨Ô,Ð,r$   c                óº  — | j                  d|¬«      \  }}||z
  }|j                  d|¬«      \  }}	t        j                  |«      t        j                  |«      }}
||
z  |	|z  z   }| |z  |	|
z  z   }|j                  d|¬«      \  }}|dz  |z
  }|dz  |z
  }|dz  |z   }|dz  |z   }t        j                  ||||g|¬«      }|�|j                  d|dz
  «      }|S )a[  Transform rotated bounding box (xywh) to distance (ltrb). This is the inverse of dist2rbox.

    Args:
        target_bboxes (torch.Tensor): Target rotated bounding boxes with shape (bs, h*w, 4), format [x, y, w, h].
        anchor_points (torch.Tensor): Anchor points with shape (h*w, 2).
        target_angle (torch.Tensor): Target angle with shape (bs, h*w, 1).
        dim (int, optional): Dimension along which to split.
        reg_max (int, optional): Maximum regression value for clamping.

    Returns:
        (torch.Tensor): Rotated distance with shape (bs, h*w, 4), format [l, t, r, b].
    re   r’   r   r   rá   )r¸   r0   rç   rè   r–   rv   )rU   rÑ   Útarget_anglerF   rã   rï   rÞ   ÚoffsetÚoffset_xÚoffset_yrç   rè   rë   rì   rÕ   rÔ   Útarget_lÚtarget_tÚtarget_rÚtarget_brä   s                        r#   Ú	rbox2distrú   Ê  s  € ð& × Ñ  ¨Ð Ó,�F€BˆØ�-Ñ€FØŸ™ a¨S˜Ó1Ñ€HˆhÜ�y‰y˜Ó&¬¯	©	°,Ó(?ˆ€CØ	�C‰˜( S™.Ñ	(€BØ
ˆ�S‰˜8 c™>Ñ	)€Bà�8‰8�A˜3ˆ8Ó�D€A€qØ�1‰u�r‰z€HØ�1‰u�r‰z€HØ�1‰u�r‰z€HØ�1‰u�r‰z€Hä�9‰9�h ¨(°HÐ=À3ÔG€DØÐØ�{‰{˜1˜g¨™nÓ-ˆà€Kr$   )g      à?)TrE   r'   )rÑ   útorch.Tensorrâ   rû   rã   ú
int | NoneÚreturnrû   )rE   )rE   N)
rU   rû   rÑ   rû   rò   rû   rF   r«   rã   rü   )Ú
__future__r   r0   Útorch.nnÚnnÚ r   r‚   r   r   Úopsr   r	   r
   Útorch_utilsr   ÚModuler   r´   rØ   rß   rå   rð   rú   rÅ   r$   r#   ú<module>r     s‘   ðõ #ã Ý å ß &ß 5Ñ 5Ý #ôU0˜"Ÿ)™)ô U0ôp
'eÐ!4ô 'eóT>ó 	(ôó-ð2 Øð$Øð$àð$ð ð$ð 
ð	$ð
 ô$r$   