Ë
    Fêñi};  ã                  óÈ   — d dl mZ d dlmZ d dlZd dlmZ d dlmc mZ	 d dl
mZ d dlmZ d dlmZmZ  G d„ dej"                  «      Z	 	 	 	 d
	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dd	„Zy)é    )Úannotations)ÚAnyN)Úlinear_sum_assignment)Úbbox_iou)Ú	xywh2xyxyÚ	xyxy2xywhc                  ót   ‡ — e Zd ZdZ	 	 	 	 	 	 d	 	 	 	 	 	 	 	 	 	 	 dˆ fd„Z	 	 d	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dd„Zˆ xZS )ÚHungarianMatcheraÑ  A module implementing the HungarianMatcher for optimal assignment between predictions and ground truth.

    HungarianMatcher performs optimal bipartite assignment over predicted and ground truth bounding boxes using a cost
    function that considers classification scores, bounding box coordinates, and optionally mask predictions. This is
    used in end-to-end object detection models like DETR.

    Attributes:
        cost_gain (dict[str, float]): Dictionary of cost coefficients for 'class', 'bbox', 'giou', 'mask', and 'dice'
            components.
        use_fl (bool): Whether to use Focal Loss for classification cost calculation.
        with_mask (bool): Whether the model makes mask predictions.
        num_sample_points (int): Number of sample points used in mask cost calculation.
        alpha (float): Alpha factor in Focal Loss calculation.
        gamma (float): Gamma factor in Focal Loss calculation.

    Methods:
        forward: Compute optimal assignment between predictions and ground truths for a batch.
        _cost_mask: Compute mask cost and dice cost if masks are predicted.

    Examples:
        Initialize a HungarianMatcher with custom cost gains
        >>> matcher = HungarianMatcher(cost_gain={"class": 2, "bbox": 5, "giou": 2})

        Perform matching between predictions and ground truth
        >>> pred_boxes = torch.rand(2, 100, 4)  # batch_size=2, num_queries=100
        >>> pred_scores = torch.rand(2, 100, 80)  # 80 classes
        >>> gt_boxes = torch.rand(10, 4)  # 10 ground truth boxes
        >>> gt_classes = torch.randint(0, 80, (10,))
        >>> gt_groups = [5, 5]  # 5 GT boxes per image
        >>> indices = matcher(pred_boxes, pred_scores, gt_boxes, gt_classes, gt_groups)
    c                óŠ   •— t         ‰| �  «        |€ddddddœ}|| _        || _        || _        || _        || _        || _        y)aÉ  Initialize HungarianMatcher for optimal assignment of predicted and ground truth bounding boxes.

        Args:
            cost_gain (dict[str, float], optional): Dictionary of cost coefficients for different matching cost
                components. Should contain keys 'class', 'bbox', 'giou', 'mask', and 'dice'.
            use_fl (bool): Whether to use Focal Loss for classification cost calculation.
            with_mask (bool): Whether the model makes mask predictions.
            num_sample_points (int): Number of sample points used in mask cost calculation.
            alpha (float): Alpha factor in Focal Loss calculation.
            gamma (float): Gamma factor in Focal Loss calculation.
        Né   é   é   )ÚclassÚbboxÚgiouÚmaskÚdice)ÚsuperÚ__init__Ú	cost_gainÚuse_flÚ	with_maskÚnum_sample_pointsÚalphaÚgamma)Úselfr   r   r   r   r   r   Ú	__class__s          €ú^/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/models/utils/ops.pyr   zHungarianMatcher.__init__1   sQ   ø€ ô( 	‰ÑÔØÐØ"#¨Q¸À1ÈaÑPˆIØ"ˆŒØˆŒØ"ˆŒØ!2ˆÔØˆŒ
Øˆ�
ó    c           
     ó  — |j                   \  }}	}
t        |«      dk(  rat        |«      D �cg c]L  }t        j                  g t        j
                  ¬«      t        j                  g t        j
                  ¬«      f‘ŒN c}S |j                  «       j                  d|
«      }| j                  rt        j                  |«      nt        j                  |d¬«      }|j                  «       j                  dd«      }|dd…|f   }| j                  rqd| j                  z
  || j                  z  z  d|z
  dz   j                  «        z  }| j                  d|z
  | j                  z  z  |dz   j                  «        z  }||z
  }n| }|j                  d«      |j                  d«      z
  j!                  «       j                  d«      }d	t#        |j                  d«      |j                  d«      d
d
¬«      j%                  d«      z
  }| j&                  d   |z  | j&                  d   |z  z   | j&                  d   |z  z   }| j(                  r|| j+                  ||||«      z  }d||j-                  «       |j/                  «       z  <   |j                  ||	d«      j1                  «       }t3        |j5                  |d«      «      D ��cg c]  \  }}t7        ||   «      ‘Œ }}}t        j8                  dg|dd ¢«      j;                  d«      }t3        |«      D ���cg c]X  \  }\  }}t        j                  |t        j
                  ¬«      t        j                  |t        j
                  ¬«      ||   z   f‘ŒZ c}}}S c c}w c c}}w c c}}}w )a  Compute optimal assignment between predictions and ground truth using Hungarian algorithm.

        This method calculates matching costs based on classification scores, bounding box coordinates, and optionally
        mask predictions, then finds the optimal bipartite assignment between predictions and ground truth.

        Args:
            pred_bboxes (torch.Tensor): Predicted bounding boxes with shape (batch_size, num_queries, 4).
            pred_scores (torch.Tensor): Predicted classification scores with shape (batch_size, num_queries,
                num_classes).
            gt_bboxes (torch.Tensor): Ground truth bounding boxes with shape (num_gts, 4).
            gt_cls (torch.Tensor): Ground truth class labels with shape (num_gts,).
            gt_groups (list[int]): Number of ground truth boxes for each image in the batch.
            masks (torch.Tensor, optional): Predicted masks with shape (batch_size, num_queries, height, width).
            gt_mask (list[torch.Tensor], optional): Ground truth masks, each with shape (num_masks, Height, Width).

        Returns:
            (list[tuple[torch.Tensor, torch.Tensor]]): A list of size batch_size, each element is a tuple (index_i,
                index_j), where index_i is the tensor of indices of the selected predictions (in order) and index_j is
                the tensor of indices of the corresponding selected ground truth targets (in order).
            For each batch element, it holds: len(index_i) = len(index_j) = min(num_queries, num_target_boxes).
        r   ©Údtypeéÿÿÿÿ©Údimé   Nr   g:Œ0âŽyE>ç      ð?T)ÚxywhÚGIoUr   r   r   ç        )ÚshapeÚsumÚrangeÚtorchÚtensorÚlongÚdetachÚviewr   ÚFÚsigmoidÚsoftmaxr   r   ÚlogÚ	unsqueezeÚabsr   Úsqueezer   r   Ú
_cost_maskÚisnanÚisinfÚcpuÚ	enumerateÚsplitr   Ú	as_tensorÚcumsum_)r   Úpred_bboxesÚpred_scoresÚ	gt_bboxesÚgt_clsÚ	gt_groupsÚmasksÚgt_maskÚbsÚnqÚncÚ_Úneg_cost_classÚpos_cost_classÚ
cost_classÚ	cost_bboxÚ	cost_giouÚCÚiÚcÚindicesÚkÚjs                          r   ÚforwardzHungarianMatcher.forwardO   s&  € ð> !×&Ñ&‰
ˆˆB�äˆy‹>˜QÒÜfkÐlnÓfoÖpÐab”U—\‘\ "¬E¯J©JÔ7¼¿¹ÀbÔPU×PZÑPZÔ9[Ò\ÒpÐpð "×(Ñ(Ó*×/Ñ/°°BÓ7ˆØ04·²”a—i‘i Ô,ÄÇÁÈ;Ð\^ÔA_ˆØ!×(Ñ(Ó*×/Ñ/°°AÓ6ˆð "¢! V )Ñ,ˆØ�;Š;Ø $§*¡*™n°¸d¿j¹jÑ1HÑIÈqÐS^ÉÐaeÑOe×NjÑNjÓNlÐMlÑmˆNØ!ŸZ™Z¨A°©OÀÇ
Á
Ñ+JÑKÐQ\Ð_cÑQc×PhÑPhÓPjÐOjÑkˆNØ'¨.Ñ8‰Jà%˜ˆJð !×*Ñ*¨1Ó-°	×0CÑ0CÀAÓ0FÑF×KÑKÓM×QÑQÐRTÓUˆ	ð œ( ;×#8Ñ#8¸Ó#;¸Y×=PÑ=PÐQRÓ=SÐZ^ÐeiÔj×rÑrÐsuÓvÑvˆ	ð �N‰N˜7Ñ# jÑ0Ø�n‰n˜VÑ$ yÑ0ñ1à�n‰n˜VÑ$ yÑ0ñ1ð 	
ð �>Š>Ø�—‘  Y°°wÓ?Ñ?ˆAð $'ˆˆ!�'‰'‹)�a—g‘g“iÑ
Ñ à�F‰F�2�r˜2Ó×"Ñ"Ó$ˆÜ;DÀQÇWÁWÈYÐXZÓE[Ó;\×]±4°1°aÔ(¨¨1©Õ.Ð]ˆÑ]Ü—O‘O QÐ$8¨°3°B¨Ð$8Ó9×AÑAÀ!ÓDˆ	ô ' wÓ/÷
ð 
á�‘6�A�qô �\‰\˜!¤5§:¡:Ô.´·±¸QÄeÇjÁjÔ0QÐT]Ð^_ÑT`Ñ0`Òaô
ð 	
ùòO qùóJ ^ùô
s   ¬AM4Ê>M9ÌAM?)NTFi 1  g      Ð?ç       @)r   zdict[str, float] | Noner   Úboolr   rZ   r   Úintr   Úfloatr   r\   )NN)rB   útorch.TensorrC   r]   rD   r]   rE   r]   rF   z	list[int]rG   ztorch.Tensor | NonerH   zlist[torch.Tensor] | NoneÚreturnz'list[tuple[torch.Tensor, torch.Tensor]])Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   rX   Ú__classcell__)r   s   @r   r
   r
      sÁ   ø„ ñðD .2ØØØ!&ØØðà*ðð ðð ð	ð
 ðð ðð õðJ &*Ø-1ðL
à!ðL
ð "ðL
ð  ð	L
ð
 ðL
ð ðL
ð #ðL
ð +ðL
ð 
1÷L
r   r
   c           
     óŒ	  — |r|dk  s| €y| d   }t        |«      }	t        |«      }
|
dk(  ry||
z  }|dk(  rdn|}t        |«      }| d   }| d   }| d   }|j                  d	|z  «      }|j                  d	|z  d«      }|j                  d	|z  «      j	                  d
«      }t        j                  |	|z  t
        j                  |j                  ¬«      ||	z  z   }|dkD  r|t        j                  |j                  «      |dz  k  }t        j                  |«      j                  d
«      }t        j                  |d||j                  |j                  ¬«      }|||<   |dkD  r«t        |«      }|dd	d…f   dz  j                  dd	«      |z  }t        j                  |dd	«      dz  dz
  }t        j                   |«      }||xx   dz  cc<   ||z  }|||z  z  }|j#                  dd¬«       t%        |«      }t        j&                  |d¬«      }t)        |
d	z  |z  «      }||   }t        j*                  |||j                  d
   |j                  ¬«      }t        j*                  ||d|j                  ¬«      }t        j,                  |D �cg c]0  }t        j.                  t1        |«      t
        j                  ¬«      ‘Œ2 c}«      }t        j2                  t1        |«      D � cg c]
  } ||
| z  z   ‘Œ c} d¬«      }!t        j,                  t1        d	|z  «      D � cg c]
  } ||
| z  z   ‘Œ c} «      }||||f<   ||||f<   ||z   }"t        j*                  |"|"gt
        j4                  ¬«      }#d|#|d…d|…f<   t1        |«      D ]–  } | dk(  r#d|#|
d	z  | z  |
d	z  | dz   z  …|
d	z  | dz   z  |…f<   | |dz
  k(  r!d|#|
d	z  | z  |
d	z  | dz   z  …d|
| z  d	z  …f<   ŒTd|#|
d	z  | z  |
d	z  | dz   z  …|
d	z  | dz   z  |…f<   d|#|
d	z  | z  |
d	z  | dz   z  …d|
d	z  | z  …f<   Œ˜ |!j7                  «       j9                  t;        |«      d¬«      D �$cg c]  }$|$j=                  d
«      ‘Œ c}$|||gdœ}%|j?                  |j                  «      |j?                  |j                  «      |#j?                  |j                  «      |%fS c c}w c c} w c c} w c c}$w )a»  Generate contrastive denoising training group with positive and negative samples from ground truths.

    This function creates denoising queries for contrastive denoising training by adding noise to ground truth bounding
    boxes and class labels. It generates both positive and negative samples to improve model robustness.

    Args:
        batch (dict[str, Any]): Batch dictionary containing 'cls' (torch.Tensor with shape (num_gts,)), 'bboxes'
            (torch.Tensor with shape (num_gts, 4)), 'batch_idx' (torch.Tensor), and 'gt_groups' (list[int]) indicating
            number of ground truths per image.
        num_classes (int): Total number of object classes.
        num_queries (int): Number of object queries.
        class_embed (torch.Tensor): Class embedding weights to map labels to embedding space.
        num_dn (int): Number of denoising queries to generate.
        cls_noise_ratio (float): Noise ratio for class labels.
        box_noise_scale (float): Noise scale for bounding box coordinates.
        training (bool): Whether model is in training mode.

    Returns:
        padding_cls (torch.Tensor | None): Modified class embeddings for denoising with shape (bs, num_dn, embed_dim).
        padding_bbox (torch.Tensor | None): Modified bounding boxes for denoising with shape (bs, num_dn, 4).
        attn_mask (torch.Tensor | None): Attention mask for denoising with shape (tgt_size, tgt_size).
        dn_meta (dict[str, Any] | None): Meta information dictionary containing denoising parameters.

    Examples:
        Generate denoising group for training
        >>> batch = {
        ...     "cls": torch.tensor([0, 1, 2]),
        ...     "bboxes": torch.rand(3, 4),
        ...     "batch_idx": torch.tensor([0, 0, 1]),
        ...     "gt_groups": [2, 1],
        ... }
        >>> class_embed = torch.rand(80, 256)  # 80 classes, 256 embedding dim
        >>> cdn_outputs = get_cdn_group(batch, 80, 100, class_embed, training=True)
    r   N)NNNNrF   r   ÚclsÚbboxesÚ	batch_idxr   r#   )r"   Údeviceç      à?.rY   r'   r*   )ÚminÚmaxg�íµ ÷Æ°>)Úeps)rh   r&   r!   r$   T)Ú
dn_pos_idxÚdn_num_groupÚdn_num_split) r,   rk   ÚlenÚrepeatr2   r.   Úaranger0   rh   Úrandr+   Únonzeror9   Úrandint_liker"   r   Ú	rand_likeÚclip_r   Úlogitr[   ÚzerosÚcatr/   r-   ÚstackrZ   r=   r?   ÚlistÚreshapeÚto)&ÚbatchÚnum_classesÚnum_queriesÚclass_embedÚnum_dnÚcls_noise_ratioÚbox_noise_scaleÚtrainingrF   Ú	total_numÚmax_numsÚ	num_grouprI   rE   Úgt_bboxÚb_idxÚdn_clsÚdn_bboxÚdn_b_idxÚneg_idxr   ÚidxÚ	new_labelÚ
known_bboxÚdiffÚ	rand_signÚ	rand_partÚdn_cls_embedÚpadding_clsÚpadding_bboxÚnumÚmap_indicesrS   Úpos_idxÚtgt_sizeÚ	attn_maskÚpÚdn_metas&                                         r   Úget_cdn_groupr    ¼   sö  € ñX ˜ 1š¨¨Ø%Ø�kÑ"€IÜ�I“€IÜ�9‹~€HØ�1‚}Ø%à˜(Ñ"€IØ !’^‘¨€Iä	ˆY‹€BØ�5‰\€FØ�H‰o€GØ�+Ñ€Eð �]‰]˜1˜y™=Ó)€FØ�n‰n˜Q ™]¨AÓ.€GØ�|‰|˜A 	™MÓ*×/Ñ/°Ó3€Hô �l‰l˜9 yÑ0¼¿
¹
È7Ï>É>ÔZÐ]fÐirÑ]rÑr€Gà˜Òä�z‰z˜&Ÿ,™,Ó'¨?¸SÑ+@ÑAˆÜ�m‰m˜DÓ!×)Ñ)¨"Ó-ˆä×&Ñ& s¨A¨{À&Ç,Á,ÐW]×WdÑWdÔeˆ	Øˆˆs‰à˜ÒÜ˜wÓ'ˆ
à˜˜Q™R˜Ñ  3Ñ&×.Ñ.¨q°!Ó4°ÑFˆä×&Ñ& w°°1Ó5¸Ñ;¸cÑAˆ	Ü—O‘O GÓ,ˆ	Ø�'Ó˜cÑ!ÓØ�YÑˆ	Ø�i $Ñ&Ñ&ˆ
Ø×Ñ˜S cÐÔ*Ü˜JÓ'ˆÜ—+‘+˜g¨4Ô0ˆä�˜A‘ 	Ñ)Ó*€FØ˜vÑ&€LÜ—+‘+˜b &¨,×*<Ñ*<¸RÑ*@ÈÏÉÔW€KÜ—;‘;˜r 6¨1°W·^±^ÔD€Lä—)‘)ÐS\Ö]ÈCœUŸ\™\¬%°«*¼E¿J¹JÖGÒ]Ó^€KÜ�k‰k¼uÀYÓ?OÖP¸!˜;¨°A©Ó5ÒPÐVWÔX€Gä—)‘)ÄÀqÈ9Á}ÓAUÖV¸A˜[¨8°a©<Ó7ÒVÓW€KØ+7€K�˜;Ð'Ñ(Ø,3€L�(˜KÐ(Ñ)à˜Ñ#€HÜ—‘˜X xÐ0¼¿
¹
ÔC€Ià"&€Iˆf‰g�w˜�wÐÑä�9Óò \ˆØ�Š6ØdhˆI�h ‘l QÑ&¨°A©¸¸Q¹Ñ)?Ð?ÀÈAÁÐQRÐUVÑQVÑAWÐZ`ÐA`Ð`ÑaØ�	˜A‘ÒØW[ˆI�h ‘l QÑ&¨°A©¸¸Q¹Ñ)?Ð?ÐASÀ8ÈaÁ<ÐRSÑCSÐASÐSÒTàdhˆI�h ‘l QÑ&¨°A©¸¸Q¹Ñ)?Ð?ÀÈAÁÐQRÐUVÑQVÑAWÐZ`ÐA`Ð`ÑaØW[ˆI�h ‘l QÑ&¨°A©¸¸Q¹Ñ)?Ð?ÐASÀ8ÈaÁ<ÐRSÑCSÐASÐSÒTð\ð /6¯k©k«m×.AÑ.AÄ$ÀyÃ/ÐWXÐ.AÓ.YÖZ¨�q—y‘y •}ÒZØ!Ø Ð-ñ€Gð 	�‰�{×)Ñ)Ó*Ø�‰˜×*Ñ*Ó+Ø�‰�[×'Ñ'Ó(Øð	ð ùò5 ^ùÚPùâVùò$ [s   Ê5R2ËR7ÌR<Ñ S)éd   ri   r'   F)r   zdict[str, Any]r€   r[   r�   r[   r‚   r]   rƒ   r[   r„   r\   r…   r\   r†   rZ   r^   z[tuple[torch.Tensor | None, torch.Tensor | None, torch.Tensor | None, dict[str, Any] | None])Ú
__future__r   Útypingr   r.   Útorch.nnÚnnÚtorch.nn.functionalÚ
functionalr3   Úscipy.optimizer   Úultralytics.utils.metricsr   Úultralytics.utils.opsr   r   ÚModuler
   r    © r   r   ú<module>r­      s©   ðõ #å ã Ý ß Ð Ý 0å .ß 6ôK
�r—y‘yô K
ðb Ø Ø ØðØðàðð ðð ð	ð
 ðð ðð ðð ðð aôr   