§
    ‚Štj0V  ã                   ó   — d dl Z d dlmZ d dlmc mZ ddlmZmZm	Z	 ddl
mZmZmZmZmZ  e¦   «         rd dlmZ  e¦   «         rd dlmZ d„ Z G d	„ d
ej        ¦  «        Z G d„ dej        ¦  «        Z	 	 	 	 	 dd„ZdS )é    Né   )Úis_scipy_availableÚis_vision_availableÚrequires_backendsé   )Úbox_iouÚ	dice_lossÚgeneralized_box_iouÚnested_tensor_from_tensor_listÚsigmoid_focal_loss©Úlinear_sum_assignment)Úcenter_to_corners_formatc                 ó6   — d„ t          | |¦  «        D ¦   «         S )Nc                 ó   — g | ]
\  }}||d œ‘ŒS ))ÚlogitsÚ
pred_boxes© )Ú.0ÚaÚbs      ú\/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/loss/loss_rt_detr.pyú
<listcomp>z!_set_aux_loss.<locals>.<listcomp>'   s$   € ÐYÐYÐY©t¨q°!�q¨Ð*Ð*ÐYÐYÐYó    )Úzip)Úoutputs_classÚoutputs_coords     r   Ú_set_aux_lossr   &   s    € ØYÐYµs¸=È-Ñ7XÔ7XÐYÑYÔYÐYr   c                   óR   ‡ — e Zd ZdZˆ fd„Z ej        ¦   «         d„ ¦   «         Zˆ xZS )ÚRTDetrHungarianMatchera›  This class computes an assignment between the targets and the predictions of the network

    For efficiency reasons, the targets don't include the no_object. Because of this, in general, there are more
    predictions than targets. In this case, we do a 1-to-1 matching of the best predictions, while the others are
    un-matched (and thus treated as non-objects).

    Args:
        config: RTDetrConfig
    c                 óf  •— t          ¦   «                              ¦   «          t          | dg¦  «         |j        | _        |j        | _        |j        | _        |j	        | _	        |j
        | _        |j        | _        | j        | j        cxk    r| j        cxk    rdk    rn d S t          d¦  «        ‚d S )NÚscipyr   z#All costs of the Matcher can't be 0)ÚsuperÚ__init__r   Úmatcher_class_costÚ
class_costÚmatcher_bbox_costÚ	bbox_costÚmatcher_giou_costÚ	giou_costÚuse_focal_lossÚmatcher_alphaÚalphaÚmatcher_gammaÚgammaÚ
ValueError)ÚselfÚconfigÚ	__class__s     €r   r$   zRTDetrHungarianMatcher.__init__5   s¹   ø€ Ý‰Œ×ÒÑÔÐÝ˜$  	Ñ*Ô*Ð*à Ô3ˆŒØÔ1ˆŒØÔ1ˆŒà$Ô3ˆÔØÔ)ˆŒ
ØÔ)ˆŒ
àŒ?˜dœnÐCÐCÒCÐC°´ÐCÐCÒCÐCÀ!ÒCÐCÐCÐCÐCÐCÝÐBÑCÔCÐCð DÐCr   c                 óF  — |d         j         dd…         \  }}|d                              dd¦  «        }t          j        d„ |D ¦   «         ¦  «        }t          j        d„ |D ¦   «         ¦  «        }| j        rŸt          j        |d                              dd¦  «        ¦  «        }|dd…|f         }d| j        z
  || j        z  z  d|z
  d	z    	                    ¦   «          z  }	| j        d|z
  | j        z  z  |d	z    	                    ¦   «          z  }
|
|	z
  }n<|d                              dd¦  «         
                    d
¦  «        }|dd…|f          }t          j        ||d¬¦  «        }t          t          |¦  «        t          |¦  «        ¦  «         }| j        |z  | j        |z  z   | j        |z  z   }|                     ||d
¦  «                             ¦   «         }d„ |D ¦   «         }d„ t'          |                     |d
¦  «        ¦  «        D ¦   «         }d„ |D ¦   «         S )a…  Performs the matching

        Params:
            outputs: This is a dict that contains at least these entries:
                 "logits": Tensor of dim [batch_size, num_queries, num_classes] with the classification logits
                 "pred_boxes": Tensor of dim [batch_size, num_queries, 4] with the predicted box coordinates

            targets: This is a list of targets (len(targets) = batch_size), where each target is a dict containing:
                 "class_labels": Tensor of dim [num_target_boxes] (where num_target_boxes is the number of ground-truth
                           objects in the target) containing the class labels
                 "boxes": Tensor of dim [num_target_boxes, 4] containing the target box coordinates

        Returns:
            A list of size batch_size, containing tuples of (index_i, index_j) where:
                - index_i is the indices of the selected predictions (in order)
                - index_j is the indices of the corresponding selected targets (in order)
            For each batch element, it holds:
                len(index_i) = len(index_j) = min(num_queries, num_target_boxes)
        r   Nr   r   r   r   c                 ó   — g | ]
}|d          ‘ŒS ©Úclass_labelsr   ©r   Úvs     r   r   z2RTDetrHungarianMatcher.forward.<locals>.<listcomp>^   s   € ÐCÐCÐC°a  .Ô 1ÐCÐCÐCr   c                 ó   — g | ]
}|d          ‘ŒS ©Úboxesr   r8   s     r   r   z2RTDetrHungarianMatcher.forward.<locals>.<listcomp>_   s   € Ð =Ð =Ð =°  7¤Ð =Ð =Ð =r   g:Œ0âŽyE>éÿÿÿÿ)Úpc                 ó8   — g | ]}t          |d          ¦  «        ‘ŒS r;   ©Úlenr8   s     r   r   z2RTDetrHungarianMatcher.forward.<locals>.<listcomp>u   s"   € Ð2Ð2Ð2 Q•�Q�w”Z‘”Ð2Ð2Ð2r   c                 ó>   — g | ]\  }}t          ||         ¦  «        ‘ŒS r   r   )r   ÚiÚcs      r   r   z2RTDetrHungarianMatcher.forward.<locals>.<listcomp>v   s)   € ÐcÐcÐc±4°1°aÕ(¨¨1¬Ñ.Ô.ÐcÐcÐcr   c                 ó”   — g | ]E\  }}t          j        |t           j        ¬ ¦  «        t          j        |t           j        ¬ ¦  «        f‘ŒFS )©Údtype)ÚtorchÚ	as_tensorÚint64)r   rC   Újs      r   r   z2RTDetrHungarianMatcher.forward.<locals>.<listcomp>x   sH   € ÐsÐsÐsÑcgÐcdÐfg•” ­%¬+Ð6Ñ6Ô6½¼ÈÕQVÔQ\Ð8]Ñ8]Ô8]Ð^ÐsÐsÐsr   )ÚshapeÚflattenrH   Úcatr+   ÚFÚsigmoidr-   r/   ÚlogÚsoftmaxÚcdistr
   r   r(   r&   r*   ÚviewÚcpuÚ	enumerateÚsplit)r1   ÚoutputsÚtargetsÚ
batch_sizeÚnum_queriesÚout_bboxÚ
target_idsÚtarget_bboxÚout_probÚneg_cost_classÚpos_cost_classr&   r(   r*   Úcost_matrixÚsizesÚindicess                    r   ÚforwardzRTDetrHungarianMatcher.forwardD   sE  € ð* #*¨(Ô"3Ô"9¸"¸1¸"Ô"=Ñˆ
�Kð ˜<Ô(×0Ò0°°AÑ6Ô6ˆå”YÐCÐC¸7ÐCÑCÔCÑDÔDˆ
Ý”iÐ =Ð =°WÐ =Ñ =Ô =Ñ>Ô>ˆð Ôð 	2Ý”y ¨Ô!2×!:Ò!:¸1¸aÑ!@Ô!@ÑAÔAˆHØ    : Ô.ˆHØ $¤*™n°¸4¼:Ñ1EÑFÈAÐPXÉLÐ[_ÑL_×KdÒKdÑKfÔKfÐJfÑgˆNØ!œZ¨A°©L¸T¼ZÑ+GÑHÈhÐY]Éo×MbÒMbÑMdÔMdÐLdÑeˆNØ'¨.Ñ8ˆJˆJà˜xÔ(×0Ò0°°AÑ6Ô6×>Ò>¸rÑBÔBˆHØ" 1 1 1 j =Ô1Ð1ˆJõ ”K ¨+¸Ð;Ñ;Ô;ˆ	å(Õ)AÀ(Ñ)KÔ)KÕMeÐfqÑMrÔMrÑsÔsÐsˆ	à”n yÑ0°4´?ÀZÑ3OÑOÐRVÔR`ÐclÑRlÑlˆØ!×&Ò& z°;ÀÑCÔC×GÒGÑIÔIˆà2Ð2¨'Ð2Ñ2Ô2ˆØcÐc½9À[×EVÒEVÐW\Ð^`ÑEaÔEaÑ;bÔ;bÐcÑcÔcˆàsÐsÐkrÐsÑsÔsÐsr   )	Ú__name__Ú
__module__Ú__qualname__Ú__doc__r$   rH   Úno_gradre   Ú__classcell__©r3   s   @r   r    r    *   sq   ø€ € € € € ðð ðDð Dð Dð Dð Dð €U„]�_„_ð3tð 3tñ „_ð3tð 3tð 3tð 3tð 3tr   r    c                   ó¬   ‡ — e Zd ZdZˆ fd„Zdd„Zdd„Z ej        ¦   «         d„ ¦   «         Z	d„ Z
d„ Zdd	„Zd
„ Zd„ Zdd„Zd„ Zed„ ¦   «         Zd„ Zˆ xZS )Ú
RTDetrLossah  
    This class computes the losses for RTDetr. The process happens in two steps: 1) we compute hungarian assignment
    between ground truth boxes and the outputs of the model 2) we supervise each pair of matched ground-truth /
    prediction (supervise class and box).

    Args:
        matcher (`DetrHungarianMatcher`):
            Module able to compute a matching between targets and proposals.
        weight_dict (`Dict`):
            Dictionary relating each loss with its weights. These losses are configured in RTDetrConf as
            `weight_loss_vfl`, `weight_loss_bbox`, `weight_loss_giou`
        losses (`list[str]`):
            List of all the losses to be applied. See `get_loss` for a list of all available losses.
        alpha (`float`):
            Parameter alpha used to compute the focal loss.
        gamma (`float`):
            Parameter gamma used to compute the focal loss.
        eos_coef (`float`):
            Relative classification weight applied to the no-object category.
        num_classes (`int`):
            Number of object categories, omitting the special no-object category.
    c                 óŽ  •— t          ¦   «                              ¦   «          t          |¦  «        | _        |j        | _        |j        |j        |j        dœ| _	        ddg| _
        |j        | _        t          j        |j        dz   ¦  «        }| j        |d<   |                      d|¦  «         |j        | _        |j        | _        d S )N)Úloss_vflÚ	loss_bboxÚ	loss_giouÚvflr<   r   r=   Úempty_weight)r#   r$   r    ÚmatcherÚ
num_labelsÚnum_classesÚweight_loss_vflÚweight_loss_bboxÚweight_loss_giouÚweight_dictÚlossesÚeos_coefficientÚeos_coefrH   ÚonesÚregister_bufferÚfocal_loss_alphar-   Úfocal_loss_gammar/   )r1   r2   rt   r3   s      €r   r$   zRTDetrLoss.__init__“   s½   ø€ Ý‰Œ×ÒÑÔÐå-¨fÑ5Ô5ˆŒØ!Ô,ˆÔàÔ.ØÔ0ØÔ0ð
ð 
ˆÔð
 ˜gÐ&ˆŒØÔ.ˆŒÝ”z &Ô"3°aÑ"7Ñ8Ô8ˆØœ=ˆ�RÑØ×Ò˜^¨\Ñ:Ô:Ð:ØÔ,ˆŒ
ØÔ,ˆŒ
ˆ
ˆ
r   Tc                 óÔ  — d|vrt          d¦  «        ‚d|vrt          d¦  «        ‚|                      |¦  «        }|d         |         }t          j        d„ t	          ||¦  «        D ¦   «         d¬¦  «        }t          t          |                     ¦   «         ¦  «        t          |¦  «        ¦  «        \  }	}
t          j        |	¦  «        }	|d         }|j	        }t          j        d„ t	          ||¦  «        D ¦   «         ¦  «        }t          j
        |j        d d	…         | j        t          j        |j        ¬
¦  «        }|||<   t          j        || j        dz   ¬¦  «        dd d…f         }t          j        ||¬¦  «        }|	                     |¦  «        ||<   |                     d¦  «        |z  }t          j        |                     ¦   «         ¦  «        }| j        |                     | j        ¦  «        z  d|z
  z  |z                        |¦  «        }t          j        |||d¬¦  «        }|                     d¦  «                             ¦   «         |j        d         z  |z  }d|iS )Nr   ú#No predicted boxes found in outputsr   z$No predicted logits found in outputsc                 ó6   — g | ]\  }\  }}|d          |         ‘ŒS r;   r   ©r   Ú_targetÚ_rC   s       r   r   z.RTDetrLoss.loss_labels_vfl.<locals>.<listcomp>­   s*   € Ð!cÐ!cÐ!c¹/¸'Á6ÀAÀq '¨'Ô"2°1Ô"5Ð!cÐ!cÐ!cr   r   ©Údimc                 ó6   — g | ]\  }\  }}|d          |         ‘ŒS r6   r   r†   s       r   r   z.RTDetrLoss.loss_labels_vfl.<locals>.<listcomp>³   ó-   € Ð,uÐ,uÐ,uÉOÈGÑU[ÐVWÐYZ¨W°^Ô-DÀQÔ-GÐ,uÐ,uÐ,ur   r   ©rG   Údevicer   ©rw   .r=   rF   Únone)ÚweightÚ	reductionrp   )ÚKeyErrorÚ_get_source_permutation_idxrH   rN   r   r   r   ÚdetachÚdiagrG   ÚfullrL   rw   rJ   rŽ   rO   Úone_hotÚ
zeros_likeÚtoÚ	unsqueezerP   r-   Úpowr/   Ú binary_cross_entropy_with_logitsÚmeanÚsum)r1   rX   rY   rd   Ú	num_boxesrQ   ÚidxÚ	src_boxesÚtarget_boxesÚiousrˆ   Ú
src_logitsrG   Útarget_classes_originalÚtarget_classesÚtargetÚtarget_score_originalÚtarget_scoreÚ
pred_scorer‘   Úlosss                        r   Úloss_labels_vflzRTDetrLoss.loss_labels_vfl¥   sZ  € Ø˜wÐ&Ð&ÝÐ@ÑAÔAÐAØ˜7Ð"Ð"ÝÐAÑBÔBÐBØ×.Ò.¨wÑ7Ô7ˆà˜LÔ)¨#Ô.ˆ	Ý”yÐ!cÐ!cÍSÐQXÐZaÑMbÔMbÐ!cÑ!cÔ!cÐijÐkÑkÔkˆÝÕ2°9×3CÒ3CÑ3EÔ3EÑFÔFÕH`ÐamÑHnÔHnÑoÔo‰ˆˆaÝŒz˜$ÑÔˆà˜XÔ&ˆ
ØÔ ˆÝ"'¤)Ð,uÐ,uÕ_bÐcjÐlsÑ_tÔ_tÐ,uÑ,uÔ,uÑ"vÔ"vÐÝœØÔ˜R˜a˜RÔ  $Ô"2½%¼+ÈjÔN_ð
ñ 
ô 
ˆð 6ˆ�sÑÝ”˜>°tÔ7GÈ!Ñ7KÐLÑLÔLÈSÐRUÐSUÐRUÈXÔVˆå %Ô 0°ÀuÐ MÑ MÔ MÐØ%)§W¢W¨U¡^¤^Ð˜cÑ"Ø,×6Ò6°rÑ:Ô:¸VÑCˆå”Y˜z×0Ò0Ñ2Ô2Ñ3Ô3ˆ
à”*˜zŸ~š~¨d¬jÑ9Ô9Ñ9¸QÀ¹ZÑHÈ<ÑW×[Ò[Ð\aÑbÔbˆåÔ1°*¸lÐSYÐekÐlÑlÔlˆØ�yŠy˜‰|Œ|×ÒÑ!Ô! JÔ$4°QÔ$7Ñ7¸)ÑCˆØ˜DÐ!Ð!r   c                 ó   — d|vrt          d¦  «        ‚|d         }|                      |¦  «        }t          j        d„ t	          ||¦  «        D ¦   «         ¦  «        }t          j        |j        dd…         | j        t          j        |j	        ¬¦  «        }	||	|<   t          j        |                     dd¦  «        |	| j        ¦  «        }
d|
i}|S )	z‰Classification loss (NLL)
        targets dicts must contain the key "class_labels" containing a tensor of dim [nb_target_boxes]
        r   z#No logits were found in the outputsc                 ó6   — g | ]\  }\  }}|d          |         ‘ŒS r6   r   r†   s       r   r   z*RTDetrLoss.loss_labels.<locals>.<listcomp>Ð   rŒ   r   Nr   r�   r   Úloss_ce)r“   r”   rH   rN   r   r—   rL   rw   rJ   rŽ   rO   Úcross_entropyÚ	transposeÚclass_weight)r1   rX   rY   rd   r    rQ   r¥   r¡   r¦   r§   r°   r|   s               r   Úloss_labelszRTDetrLoss.loss_labelsÆ   sÛ   € ð ˜7Ð"Ð"ÝÐ@ÑAÔAÐAà˜XÔ&ˆ
à×.Ò.¨wÑ7Ô7ˆÝ"'¤)Ð,uÐ,uÕ_bÐcjÐlsÑ_tÔ_tÐ,uÑ,uÔ,uÑ"vÔ"vÐÝœØÔ˜R˜a˜RÔ  $Ô"2½%¼+ÈjÔN_ð
ñ 
ô 
ˆð 6ˆ�sÑå”/ *×"6Ò"6°q¸!Ñ"<Ô"<¸nÈdÔN_Ñ`Ô`ˆØ˜WÐ%ˆØˆr   c                 óz  — |d         }|j         }t          j        d„ |D ¦   «         |¬¦  «        }|                     ¦   «                              d¦  «        j        dk                         d¦  «        }t          j         	                    | 
                    ¦   «         | 
                    ¦   «         ¦  «        }	d|	i}
|
S )zá
        Compute the cardinality error, i.e. the absolute error in the number of predicted non-empty boxes. This is not
        really a loss, it is intended for logging purposes only. It doesn't propagate gradients.
        r   c                 ó8   — g | ]}t          |d          ¦  «        ‘ŒS r6   r@   r8   s     r   r   z/RTDetrLoss.loss_cardinality.<locals>.<listcomp>â   s%   € Ð)RÐ)RÐ)RÀQ­#¨a°Ô.?Ñ*@Ô*@Ð)RÐ)RÐ)Rr   )rŽ   r=   g      à?r   Úcardinality_error)rŽ   rH   rI   rP   ÚmaxÚvaluesrŸ   ÚnnÚ
functionalÚl1_lossÚfloat)r1   rX   rY   rd   r    r   rŽ   Útarget_lengthsÚ	card_predÚcard_errr|   s              r   Úloss_cardinalityzRTDetrLoss.loss_cardinalityÚ   s¨   € ð ˜Ô"ˆØ”ˆÝœÐ)RÐ)RÈ'Ð)RÑ)RÔ)RÐ[aÐbÑbÔbˆà—^’^Ñ%Ô%×)Ò)¨"Ñ-Ô-Ô4°sÒ:×?Ò?ÀÑBÔBˆ	Ý”=×(Ò(¨¯ªÑ):Ô):¸N×<PÒ<PÑ<RÔ<RÑSÔSˆØ% xÐ0ˆØˆr   c           	      óæ  — d|vrt          d¦  «        ‚|                      |¦  «        }|d         |         }t          j        d„ t	          ||¦  «        D ¦   «         d¬¦  «        }i }t          j        ||d¬¦  «        }	|	                     ¦   «         |z  |d<   d	t          j        t          t          |¦  «        t          |¦  «        ¦  «        ¦  «        z
  }
|
                     ¦   «         |z  |d
<   |S )a;  
        Compute the losses related to the bounding boxes, the L1 regression loss and the GIoU loss. Targets dicts must
        contain the key "boxes" containing a tensor of dim [nb_target_boxes, 4]. The target boxes are expected in
        format (center_x, center_y, w, h), normalized by the image size.
        r   r„   c                 ó6   — g | ]\  }\  }}|d          |         ‘ŒS r;   r   )r   Útrˆ   rC   s       r   r   z)RTDetrLoss.loss_boxes.<locals>.<listcomp>ó   s(   € Ð!WÐ!WÐ!W±I°A±v¸¸1 ! G¤*¨Q¤-Ð!WÐ!WÐ!Wr   r   r‰   r�   ©r’   rq   r   rr   )r“   r”   rH   rN   r   rO   r¼   rŸ   r–   r
   r   )r1   rX   rY   rd   r    r¡   r¢   r£   r|   rq   rr   s              r   Ú
loss_boxeszRTDetrLoss.loss_boxesé   sö   € ð ˜wÐ&Ð&ÝÐ@ÑAÔAÐAØ×.Ò.¨wÑ7Ô7ˆØ˜LÔ)¨#Ô.ˆ	Ý”yÐ!WÐ!WÅÀWÈgÑAVÔAVÐ!WÑ!WÔ!WÐ]^Ð_Ñ_Ô_ˆàˆå”I˜i¨ÀÐHÑHÔHˆ	Ø'Ÿmšm™oœo°	Ñ9ˆˆ{Ñà�œ
ÝÕ 8¸Ñ CÔ CÕE]Ð^jÑEkÔEkÑlÔlñ
ô 
ñ 
ˆ	ð (Ÿmšm™oœo°	Ñ9ˆˆ{ÑØˆr   c                 ó�  — d|vrt          d¦  «        ‚|                      |¦  «        }|                      |¦  «        }|d         }||         }d„ |D ¦   «         }t          |¦  «                             ¦   «         \  }	}
|	                     |¦  «        }	|	|         }	t          j                             |dd…df         |	j	        dd…         dd¬¦  «        }|dd…d	f          
                    d
¦  «        }|	 
                    d
¦  «        }	|	                     |j	        ¦  «        }	t          ||	|¦  «        t          ||	|¦  «        dœ}|S )zÃ
        Compute the losses related to the masks: the focal loss and the dice loss. Targets dicts must contain the key
        "masks" containing a tensor of dim [nb_target_boxes, h, w].
        Ú
pred_masksz#No predicted masks found in outputsc                 ó   — g | ]
}|d          ‘ŒS )Úmasksr   ©r   rÄ   s     r   r   z)RTDetrLoss.loss_masks.<locals>.<listcomp>  s   € Ð-Ð-Ð- ��7”Ð-Ð-Ð-r   NéþÿÿÿÚbilinearF)ÚsizeÚmodeÚalign_cornersr   r   )Ú	loss_maskÚ	loss_dice)r“   r”   Ú_get_target_permutation_idxr   Ú	decomposerš   rº   r»   ÚinterpolaterL   rM   rT   r   r	   )r1   rX   rY   rd   r    Ú
source_idxÚ
target_idxÚsource_masksrÊ   Útarget_masksÚvalidr|   s               r   Ú
loss_maskszRTDetrLoss.loss_masks   sa  € ð
 ˜wÐ&Ð&ÝÐ@ÑAÔAÐAà×5Ò5°gÑ>Ô>ˆ
Ø×5Ò5°gÑ>Ô>ˆ
Ø˜|Ô,ˆØ# JÔ/ˆØ-Ð- WÐ-Ñ-Ô-ˆÝ<¸UÑCÔC×MÒMÑOÔOÑˆ�eØ#—’ |Ñ4Ô4ˆØ# JÔ/ˆõ ”}×0Ò0Ø˜˜˜˜D˜Ô!¨Ô(:¸2¸3¸3Ô(?ÀjÐ`eð 1ñ 
ô 
ˆð $ A A A q DÔ)×1Ò1°!Ñ4Ô4ˆà#×+Ò+¨AÑ.Ô.ˆØ#×(Ò(¨Ô);Ñ<Ô<ˆå+¨L¸,È	ÑRÔRÝ" <°¸yÑIÔIð
ð 
ˆð ˆr   c                 ó  — |d         }|                       |¦  «        }t          j        d„ t          ||¦  «        D ¦   «         ¦  «        }t          j        |j        d d…         | j        t          j        |j        ¬¦  «        }	||	|<   t          j
        |	| j        dz   ¬¦  «        dd d…f         }
t          j        ||
d	z  d
¬¦  «        }|                     d¦  «                             ¦   «         |j        d         z  |z  }d|iS )Nr   c                 ó6   — g | ]\  }\  }}|d          |         ‘ŒS r6   r   r†   s       r   r   z.RTDetrLoss.loss_labels_bce.<locals>.<listcomp>"  rŒ   r   r   r�   r   r�   .r=   g      ð?r�   rÅ   Úloss_bce)r”   rH   rN   r   r—   rL   rw   rJ   rŽ   rO   r˜   r�   rž   rŸ   ©r1   rX   rY   rd   r    rQ   r¥   r¡   r¦   r§   r¨   r¬   s               r   Úloss_labels_bcezRTDetrLoss.loss_labels_bce  s
  € Ø˜XÔ&ˆ
Ø×.Ò.¨wÑ7Ô7ˆÝ"'¤)Ð,uÐ,uÕ_bÐcjÐlsÑ_tÔ_tÐ,uÑ,uÔ,uÑ"vÔ"vÐÝœØÔ˜R˜a˜RÔ  $Ô"2½%¼+ÈjÔN_ð
ñ 
ô 
ˆð 6ˆ�sÑå”˜>°tÔ7GÈ!Ñ7KÐLÑLÔLÈSÐRUÐSUÐRUÈXÔVˆÝÔ1°*¸fÀs¹lÐV\Ð]Ñ]Ô]ˆØ�yŠy˜‰|Œ|×ÒÑ!Ô! JÔ$4°QÔ$7Ñ7¸)ÑCˆØ˜DÐ!Ð!r   c                 óœ   — t          j        d„ t          |¦  «        D ¦   «         ¦  «        }t          j        d„ |D ¦   «         ¦  «        }||fS )Nc                 óD   — g | ]\  }\  }}t          j        ||¦  «        ‘ŒS r   ©rH   Ú	full_like)r   rC   Úsourcerˆ   s       r   r   z:RTDetrLoss._get_source_permutation_idx.<locals>.<listcomp>/  s,   € ÐcÐcÐc¹n¸aÁÀ&È!�uœ¨v°qÑ9Ô9ÐcÐcÐcr   c                 ó   — g | ]\  }}|‘ŒS r   r   )r   rå   rˆ   s      r   r   z:RTDetrLoss._get_source_permutation_idx.<locals>.<listcomp>0  s   € ÐBÐBÐB©;¨F°A ÐBÐBÐBr   ©rH   rN   rV   )r1   rd   Ú	batch_idxrÖ   s       r   r”   z&RTDetrLoss._get_source_permutation_idx-  óS   € å”IÐcÐcÕPYÐZaÑPbÔPbÐcÑcÔcÑdÔdˆ	Ý”YÐBÐB¸'ÐBÑBÔBÑCÔCˆ
Ø˜*Ð$Ð$r   c                 óœ   — t          j        d„ t          |¦  «        D ¦   «         ¦  «        }t          j        d„ |D ¦   «         ¦  «        }||fS )Nc                 óD   — g | ]\  }\  }}t          j        ||¦  «        ‘ŒS r   rã   )r   rC   rˆ   r¨   s       r   r   z:RTDetrLoss._get_target_permutation_idx.<locals>.<listcomp>5  s,   € ÐcÐcÐc¹n¸aÁÀ!ÀV�uœ¨v°qÑ9Ô9ÐcÐcÐcr   c                 ó   — g | ]\  }}|‘ŒS r   r   )r   rˆ   r¨   s      r   r   z:RTDetrLoss._get_target_permutation_idx.<locals>.<listcomp>6  s   € ÐBÐBÐB©;¨A¨v ÐBÐBÐBr   rç   )r1   rd   rè   r×   s       r   rÓ   z&RTDetrLoss._get_target_permutation_idx3  ré   r   c                 ó6  — d|vrt          d¦  «        ‚|d         }|                      |¦  «        }t          j        d„ t	          ||¦  «        D ¦   «         ¦  «        }t          j        |j        d d…         | j        t          j        |j	        ¬¦  «        }	||	|<   t          j        |	| j        dz   ¬¦  «        dd d	…f         }
t          ||
| j        | j        ¦  «        }|                     d¦  «                             ¦   «         |j        d         z  |z  }d
|iS )Nr   zNo logits found in outputsc                 ó6   — g | ]\  }\  }}|d          |         ‘ŒS r6   r   r†   s       r   r   z0RTDetrLoss.loss_labels_focal.<locals>.<listcomp>@  rŒ   r   r   r�   r   r�   .r=   Ú
loss_focal)r“   r”   rH   rN   r   r—   rL   rw   rJ   rŽ   rO   r˜   r   r-   r/   rž   rŸ   rß   s               r   Úloss_labels_focalzRTDetrLoss.loss_labels_focal9  s  € Ø˜7Ð"Ð"ÝÐ7Ñ8Ô8Ð8à˜XÔ&ˆ
à×.Ò.¨wÑ7Ô7ˆÝ"'¤)Ð,uÐ,uÕ_bÐcjÐlsÑ_tÔ_tÐ,uÑ,uÔ,uÑ"vÔ"vÐÝœØÔ˜R˜a˜RÔ  $Ô"2½%¼+ÈjÔN_ð
ñ 
ô 
ˆð 6ˆ�sÑå”˜>°tÔ7GÈ!Ñ7KÐLÑLÔLÈSÐRUÐSUÐRUÈXÔVˆÝ! *¨f°d´jÀ$Ä*ÑMÔMˆØ�yŠy˜‰|Œ|×ÒÑ!Ô! JÔ$4°QÔ$7Ñ7¸)ÑCˆØ˜dÐ#Ð#r   c                 ó²   — | j         | j        | j        | j        | j        | j        | j        dœ}||vrt          d|› d�¦  «        ‚ ||         ||||¦  «        S )N)ÚlabelsÚcardinalityr<   rÊ   ÚbceÚfocalrs   zLoss z not supported)r´   rÁ   rÆ   rÛ   rà   rð   r­   r0   )r1   r¬   rX   rY   rd   r    Úloss_maps          r   Úget_losszRTDetrLoss.get_lossK  sw   € àÔ&ØÔ0Ø”_Ø”_ØÔ'ØÔ+ØÔ'ð
ð 
ˆð �xÐÐÝÐ9 TÐ9Ð9Ð9Ñ:Ô:Ð:Øˆx˜Œ~˜g w°¸ÑCÔCÐCr   c           	      ó@  — | d         | d         }}d„ |D ¦   «         }|d         d         j         }g }t          |¦  «        D ]Ü\  }}|dk    r|t          j        |t          j        |¬¦  «        }	|	                     |¦  «        }	t          ||         ¦  «        t          |	¦  «        k    sJ ‚|                     ||         |	f¦  «         Œ‡|                     t          j        dt          j        |¬¦  «        t          j        dt          j        |¬¦  «        f¦  «         ŒÝ|S )NÚdn_positive_idxÚdn_num_groupc                 ó8   — g | ]}t          |d          ¦  «        ‘ŒS r6   r@   rË   s     r   r   z6RTDetrLoss.get_cdn_matched_indices.<locals>.<listcomp>\  s%   € Ð;Ð;Ð;¨a•3�q˜Ô(Ñ)Ô)Ð;Ð;Ð;r   r   r7   r�   )	rŽ   rV   rH   ÚarangerJ   ÚtilerA   ÚappendÚzeros)
Údn_metarY   rù   rú   Únum_gtsrŽ   Údn_match_indicesrC   Únum_gtÚgt_idxs
             r   Úget_cdn_matched_indicesz"RTDetrLoss.get_cdn_matched_indicesY  s(  € à(/Ð0AÔ(BÀGÈNÔD[˜ˆØ;Ð;°7Ð;Ñ;Ô;ˆØ˜”˜NÔ+Ô2ˆàÐÝ" 7Ñ+Ô+ð 	ð 	‰IˆAˆvØ˜ŠzˆzÝœ fµE´KÈÐOÑOÔO�ØŸš \Ñ2Ô2�Ý˜?¨1Ô-Ñ.Ô.µ#°f±+´+Ò=Ð=Ð=Ð=Ø ×'Ò'¨¸Ô);¸VÐ(DÑEÔEÐEÐEà ×'Ò'åœ A­U¬[ÀÐHÑHÔHÝœ A­U¬[ÀÐHÑHÔHðñô ð ð ð  Ðr   c           
      ó  ‡ ‡
‡— d„ |                      ¦   «         D ¦   «         }‰                      ||¦  «        }t          d„ |D ¦   «         ¦  «        }t          j        |gt          j        t          t          |                     ¦   «         ¦  «        ¦  «        j	        ¬¦  «        }t          j
        |d¬¦  «                             ¦   «         }i }‰ j        D ]?}‰                      |||||¦  «        Šˆˆ fd„‰D ¦   «         Š|                     ‰¦  «         Œ@d|v rŸt          |d         ¦  «        D ]‰\  Š
}‰                      ||¦  «        }‰ j        D ]f}|dk    rŒ	‰                      |||||¦  «        Šˆˆ fd	„‰D ¦   «         Šˆ
fd
„‰                      ¦   «         D ¦   «         Š|                     ‰¦  «         ŒgŒŠd|v rÄd|vrt!          d¦  «        ‚‰                      |d         |¦  «        }||d         d         z  }t          |d         ¦  «        D ]n\  Š
}‰ j        D ]a}|dk    rŒ	i }	 ‰ j        |||||fi |	¤ŽŠˆˆ fd„‰D ¦   «         Šˆ
fd„‰                      ¦   «         D ¦   «         Š|                     ‰¦  «         ŒbŒo|S )aª  
        This performs the loss computation.

        Args:
             outputs (`dict`, *optional*):
                Dictionary of tensors, see the output specification of the model for the format.
             targets (`list[dict]`, *optional*):
                List of dicts, such that `len(targets) == batch_size`. The expected keys in each dict depends on the
                losses applied, see each loss' doc.
        c                 ó"   — i | ]\  }}d |v¯	||“ŒS )Úauxiliary_outputsr   )r   Úkr9   s      r   ú
<dictcomp>z&RTDetrLoss.forward.<locals>.<dictcomp>{  s*   € Ð`Ð`Ð`©¨¨1ÐCVÐ^_ÐC_ÐC_˜q !ÐC_ÐC_ÐC_r   c              3   ó@   K  — | ]}t          |d          ¦  «        V — ŒdS )r7   Nr@   rË   s     r   ú	<genexpr>z%RTDetrLoss.forward.<locals>.<genexpr>�  s/   è è € Ð@Ð@°1�˜A˜nÔ-Ñ.Ô.Ð@Ð@Ð@Ð@Ð@Ð@r   r�   r   )Úminc                 óP   •— i | ]"}|‰j         v ¯|‰|         ‰j         |         z  “Œ#S r   ©r{   ©r   r	  Úl_dictr1   s     €€r   r
  z&RTDetrLoss.forward.<locals>.<dictcomp>‰  s:   ø€ ÐbÐbÐb¸QÈAÐQUÔQaÐLaÐLa�a˜ œ TÔ%5°aÔ%8Ñ8ÐLaÐLaÐLar   r  rÊ   c                 óP   •— i | ]"}|‰j         v ¯|‰|         ‰j         |         z  “Œ#S r   r  r  s     €€r   r
  z&RTDetrLoss.forward.<locals>.<dictcomp>•  ó;   ø€ ÐjÐjÐjÀQÐTUÐY]ÔYiÐTiÐTi˜a ¨¤¨TÔ-=¸aÔ-@Ñ!@ÐTiÐTiÐTir   c                 ó(   •— i | ]\  }}|d ‰› �z   |“ŒS )Ú_aux_r   ©r   r	  r9   rC   s      €r   r
  z&RTDetrLoss.forward.<locals>.<dictcomp>–  s)   ø€ ÐLÐLÐL±T°Q¸˜a +¨! + +™o¨qÐLÐLÐLr   Údn_auxiliary_outputsÚdenoising_meta_valuesz}The output must have the 'denoising_meta_values` key. Please, ensure that 'outputs' includes a 'denoising_meta_values' entry.rú   c                 óP   •— i | ]"}|‰j         v ¯|‰|         ‰j         |         z  “Œ#S r   r  r  s     €€r   r
  z&RTDetrLoss.forward.<locals>.<dictcomp>ª  r  r   c                 ó(   •— i | ]\  }}|d ‰› �z   |“ŒS )Ú_dn_r   r  s      €r   r
  z&RTDetrLoss.forward.<locals>.<dictcomp>«  s)   ø€ ÐKÐKÐK±D°A°q˜a *¨ * *™n¨aÐKÐKÐKr   )Úitemsru   rŸ   rH   rI   r½   ÚnextÚiterr¹   rŽ   ÚclampÚitemr|   r÷   ÚupdaterV   r0   r  )r1   rX   rY   Úoutputs_without_auxrd   r    r|   r¬   r  ÚkwargsrC   r  s   `         @@r   re   zRTDetrLoss.forwardp  s  øøø€ ð aÐ`°·²±´Ð`Ñ`Ô`Ðð —,’,Ð2°GÑ<Ô<ˆõ Ð@Ð@¸Ð@Ñ@Ô@Ñ@Ô@ˆ	Ý”O Y Kµu´{Í4ÕPTÐU\×UcÒUcÑUeÔUeÑPfÔPfÑKgÔKgÔKnÐoÑoÔoˆ	Ý”K 	¨qÐ1Ñ1Ô1×6Ò6Ñ8Ô8ˆ	ð ˆØ”Kð 	"ð 	"ˆDØ—]’] 4¨°'¸7ÀIÑNÔNˆFØbÐbÐbÐbÐbÀ&ÐbÑbÔbˆFØ�MŠM˜&Ñ!Ô!Ð!Ð!ð  'Ð)Ð)Ý(1°'Ð:MÔ2NÑ(OÔ(Oð 	*ð 	*Ñ$�Ð$ØŸ,š,Ð'8¸'ÑBÔB�Ø œKð *ð *�DØ˜w’�à Ø!Ÿ]š]¨4Ð1BÀGÈWÐV_Ñ`Ô`�FØjÐjÐjÐjÐjÈ&ÐjÑjÔj�FØLÐLÐLÐL¸V¿\º\¹^¼^ÐLÑLÔL�FØ—M’M &Ñ)Ô)Ð)Ð)ð*ð " WÐ,Ð,Ø&¨gÐ5Ð5Ý ð Tñô ð ð ×2Ò2°7Ð;RÔ3SÐU\Ñ]Ô]ˆGØ! GÐ,CÔ$DÀ^Ô$TÑTˆIå(1°'Ð:PÔ2QÑ(RÔ(Rð 
*ð 
*Ñ$�Ð$à œKð *ð *�DØ˜w’�à Ø�FØ*˜Tœ]¨4Ð1BÀGÈWÐV_ÐjÐjÐciÐjÐj�FØjÐjÐjÐjÐjÈ&ÐjÑjÔj�FØKÐKÐKÐK¸F¿LºL¹N¼NÐKÑKÔK�FØ—M’M &Ñ)Ô)Ð)Ð)ð*ð ˆr   )T)rf   rg   rh   ri   r$   r­   r´   rH   rj   rÁ   rÆ   rÛ   rà   r”   rÓ   rð   r÷   Ústaticmethodr  re   rk   rl   s   @r   rn   rn   {   s.  ø€ € € € € ðð ð.-ð -ð -ð -ð -ð$"ð "ð "ð "ðBð ð ð ð( €U„]�_„_ðð ñ „_ððð ð ð.ð ð ð>"ð "ð "ð "ð%ð %ð %ð%ð %ð %ð$ð $ð $ð $ð$Dð Dð Dð ð ð  ñ „\ð ð,>ð >ð >ð >ð >ð >ð >r   rn   c
                 óâ  — t          |¦  «        }|                     |¦  «         i }| |d<   ||d<   d }|j        �r|	�@t          j        ||	d         d¬¦  «        \  }}t          j        ||	d         d¬¦  «        \  }}t          |d d …d d…f                              dd¦  «        |d d …d d…f                              dd¦  «        ¦  «        }||d	<   |d	                              t          |g|g¦  «        ¦  «         |	�@t          |                     dd¦  «        |                     dd¦  «        ¦  «        |d
<   |	|d<    |||¦  «        }t          | 	                    ¦   «         ¦  «        }|||fS )Nr   r   Údn_num_splitr   r‰   r=   r   r   r  r  r  )
rn   rš   Úauxiliary_lossrH   rW   r   r²   ÚextendrŸ   r¹   )r   rò   rŽ   r   r2   r   r   Úenc_topk_logitsÚenc_topk_bboxesr  r#  Ú	criterionÚoutputs_lossr  Údn_out_coordÚdn_out_classÚ	loss_dictr¬   s                     r   ÚRTDetrForObjectDetectionLossr0  ±  s²  € õ ˜6Ñ"Ô"€IØ‡L‚L�ÑÔÐà€LØ#€L�ÑØ!+€L�ÑØÐØÔñ JØ Ð,Ý*/¬+°mÐEZÐ[iÔEjÐpqÐ*rÑ*rÔ*rÑ'ˆL˜-Ý*/¬+°mÐEZÐ[iÔEjÐpqÐ*rÑ*rÔ*rÑ'ˆL˜-å)¨-¸¸¸¸3¸B¸3¸Ô*?×*IÒ*IÈ!ÈQÑ*OÔ*OÐQ^Ð_`Ð_`Ð_`ÐbeÐceÐbeÐ_eÔQf×QpÒQpÐqrÐtuÑQvÔQvÑwÔwÐØ,=ˆÐ(Ñ)ØÐ(Ô)×0Ò0µÀÐ?PÐSbÐRcÑ1dÔ1dÑeÔeÐeØ Ð,Ý3@Ø×&Ò& q¨!Ñ,Ô,¨l×.DÒ.DÀQÈÑ.JÔ.Jñ4ô 4ˆLÐ/Ñ0ð 5JˆLÐ0Ñ1à�	˜,¨Ñ/Ô/€Iåˆy×ÒÑ!Ô!Ñ"Ô"€DØ�Ð-Ð-Ð-r   )NNNNN)rH   Útorch.nnrº   Útorch.nn.functionalr»   rO   Úutilsr   r   r   Úloss_for_object_detectionr   r	   r
   r   r   Úscipy.optimizer   Útransformers.image_transformsr   r   ÚModuler    rn   r0  r   r   r   ú<module>r8     sŸ  ðð €€€Ø Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð Ð Ð Ð à NÐ NÐ NÐ NÐ NÐ NÐ NÐ NÐ NÐ Nðð ð ð ð ð ð ð ð ð ð ð ð ð ð ÐÑÔð 5Ø4Ð4Ð4Ð4Ð4Ð4ð ÐÑÔð GØFÐFÐFÐFÐFÐFðZð Zð ZðNtð Ntð Ntð Ntð Nt˜RœYñ Ntô Ntð Ntðbsð sð sð sð s�”ñ sô sð sðx	 ØØØØð%.ð %.ð %.ð %.ð %.ð %.r   