§
    ‚Štj¾M  ã            
       óª  — d dl Zd dlZd dlmZ d dlmZmZ ddlm	Z	 ddl
mZ ddlmZmZmZ ddlmZ  e¦   «         rd d	lmZ d
ej        dej        dedej        fd„Zd
ej        dej        dej        fd„Zd
ededefd„Z	 d dej        dej        dej        fd„Zdededededef
d„Z G d„ de¦  «        Z G d„ de¦  «        Zd„ Z	 	 	 	 	 	 d!d„ZdS )"é    N)ÚTensorÚnné   )Úcenter_to_corners_format)Úis_scipy_availableé   )ÚHungarianMatcherÚ	dice_lossÚgeneralized_box_iou)ÚLwDetrImageLoss©Úlinear_sum_assignmentÚinputsÚlabelsÚ	num_masksÚreturnc                 óœ   — t          j        d¬¦  «        } || |¦  «        }|                     d¦  «                             ¦   «         |z  }|S )a|  
    Args:
        inputs (`torch.Tensor`):
            A float tensor of arbitrary shape.
        labels (`torch.Tensor`):
            A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
            (0 for the negative class and 1 for the positive class).

    Returns:
        loss (`torch.Tensor`): The computed loss.
    Únone©Ú	reductionr   )r   ÚBCEWithLogitsLossÚmeanÚsum)r   r   r   Ú	criterionÚcross_entropy_lossÚlosss         ú\/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/loss/loss_rf_detr.pyÚsigmoid_cross_entropy_lossr   #   sR   € õ Ô$¨vÐ6Ñ6Ô6€IØ"˜ 6¨6Ñ2Ô2Ðà×"Ò" 1Ñ%Ô%×)Ò)Ñ+Ô+¨iÑ7€DØ€Kó    c                 óF  — | j         d         }t          j        d¬¦  «        } || t          j        | ¦  «        ¦  «        } || t          j        | ¦  «        ¦  «        }t          j        ||z  |j        ¦  «        }t          j        ||z  d|z
  j        ¦  «        }||z   }|S )aê  
    A pair wise version of the cross entropy loss, see `sigmoid_cross_entropy_loss` for usage.

    Args:
        inputs (`torch.Tensor`):
            A tensor representing a mask.
        labels (`torch.Tensor`):
            A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
            (0 for the negative class and 1 for the positive class).

    Returns:
        loss (`torch.Tensor`): The computed loss between each pairs.
    r   r   r   )Úshaper   r   ÚtorchÚ	ones_likeÚ
zeros_likeÚmatmulÚT)	r   r   Úheight_and_widthr   Úcross_entropy_loss_posÚcross_entropy_loss_negÚloss_posÚloss_negr   s	            r   Ú$pair_wise_sigmoid_cross_entropy_lossr,   7   s£   € ð ”| A”ÐåÔ$¨vÐ6Ñ6Ô6€IØ&˜Y v­u¬¸vÑ/FÔ/FÑGÔGÐØ&˜Y v­uÔ/?ÀÑ/GÔ/GÑHÔHÐåŒ|Ð2Ð5EÑEÀvÄxÑPÔP€HÝŒ|Ð2Ð5EÑEÈÈFÉ
Ä~ÑVÔV€HØ�hÑ€DØ€Kr   c                 ó(  — |                       ¦   «                              d¦  «        } dt          j        | |j        ¦  «        z  }|                      d¦  «        dd…df         |                     d¦  «        ddd…f         z   }d|dz   |dz   z  z
  }|S )aÈ  
    A pair wise version of the dice loss, see `dice_loss` for usage.

    Args:
        inputs (`torch.Tensor`):
            A tensor representing a mask
        labels (`torch.Tensor`):
            A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
            (0 for the negative class and 1 for the positive class).

    Returns:
        `torch.Tensor`: The computed loss between each pairs.
    r   r   éÿÿÿÿN)ÚsigmoidÚflattenr"   r%   r&   r   )r   r   Ú	numeratorÚdenominatorr   s        r   Úpair_wise_dice_lossr3   S   s�   € ð �^Š^ÑÔ×%Ò% aÑ(Ô(€FØ•E”L ¨¬Ñ2Ô2Ñ2€Ià—*’*˜R‘.”.    D Ô)¨F¯JªJ°r©N¬N¸4ÀÀÀ¸7Ô,CÑC€KØ�	˜A‘ +°¡/Ñ2Ñ2€DØ€Kr   FÚinput_featuresÚpoint_coordinatesc                 óØ   — |                      ¦   «         dk    rd}|                     d¦  «        }t          j        j        j        | d|z  dz
  fi |¤Ž}|r|                     d¦  «        }|S )a(  
    A wrapper around `torch.nn.functional.grid_sample` to support 3D point_coordinates tensors.

    Args:
        input_features (`torch.Tensor` of shape (batch_size, channels, height, width)):
            A tensor that contains features map on a height * width grid
        point_coordinates (`torch.Tensor` of shape (batch_size, num_points, 2) or (batch_size, grid_height, grid_width,:
        2)):
            A tensor that contains [0, 1] * [0, 1] normalized point coordinates
        add_dim (`bool`):
            boolean value to keep track of added dimension

    Returns:
        point_features (`torch.Tensor` of shape (batch_size, channels, num_points) or (batch_size, channels,
        height_grid, width_grid):
            A tensor that contains features for points in `point_coordinates`.
    é   Tr   ç       @g      ð?)ÚdimÚ	unsqueezer"   r   Ú
functionalÚgrid_sampleÚsqueeze)r4   r5   Úadd_dimÚkwargsÚpoint_featuress        r   Úsample_pointrA   j   sƒ   € ð( ×ÒÑÔ !Ò#Ð#ØˆØ-×7Ò7¸Ñ:Ô:Ðõ ”XÔ(Ô4°^ÀSÐK\ÑE\Ð_bÑEbÐmÐmÐflÐmÐm€NØð 3Ø'×/Ò/°Ñ2Ô2ˆàÐr   ÚlogitsÚ
num_pointsÚoversample_ratioÚimportance_sample_ratioc           	      ó<  — | j         d         }t          ||z  ¦  «        }t          j        ||d| j        ¬¦  «        }t          | |d¬¦  «        }t          j        |¦  «         }t          ||z  ¦  «        }	||	z
  }
t          j        |dd…ddd…f         |	d¬¦  «        d         }t          j        |d| 	                    d	¦  «         
                    d	d	d¦  «        ¦  «        }|
dk    r3t          j        |t          j        ||
d| j        ¬¦  «        gd¬
¦  «        }|S )a8  
    This function is meant for sampling points in [0, 1] * [0, 1] coordinate space based on their uncertainty. The
    uncertainty is calculated for each point using the passed `uncertainty function` that takes points logit
    prediction as input.

    Args:
        logits (`float`):
            Logit predictions for P points.
        uncertainty_function:
            A function that takes logit predictions for P points and returns their uncertainties.
        num_points (`int`):
            The number of points P to sample.
        oversample_ratio (`int`):
            Oversampling parameter.
        importance_sample_ratio (`float`):
            Ratio of points that are sampled via importance sampling.

    Returns:
        point_coordinates (`torch.Tensor`):
            Coordinates for P sampled points.
    r   r   ©ÚdeviceF©Úalign_cornersNr   )Úkr9   r.   ©r9   )r!   Úintr"   ÚrandrH   rA   ÚabsÚtopkÚgatherr:   ÚexpandÚcat)rB   rC   rD   rE   Ú	num_boxesÚnum_points_sampledr5   Úpoint_logitsÚpoint_uncertaintiesÚnum_uncertain_pointsÚnum_random_pointsÚidxs               r   Úsample_points_using_uncertaintyr[   ‹   s=  € ð2 ”˜Q”€IÝ˜ZÐ*:Ñ:Ñ;Ô;Ðõ œ
 9Ð.@À!ÈFÌMÐZÑZÔZÐå Ð(9ÈÐOÑOÔO€Lå!œI lÑ3Ô3Ð4ÐåÐ6¸ÑCÑDÔDÐØ"Ð%9Ñ9Ðå
Œ*Ð(¨¨¨¨A¨q¨q¨q¨Ô1Ð5IÈqÐ
QÑ
QÔ
QÐRSÔ
T€CÝœÐ%6¸¸3¿=º=ÈÑ;LÔ;L×;SÒ;SÐTVÐXZÐ\]Ñ;^Ô;^Ñ_Ô_Ðà˜1ÒÐÝ!œIØ¥¤
¨9Ð6GÈÐSYÔS`Ð aÑ aÔ aÐbØð
ñ 
ô 
Ðð Ðr   c                   óv   ‡ — e Zd Z	 	 	 	 	 	 ddedededededefˆ fd	„Z ej        ¦   «         d
„ ¦   «         Zˆ xZ	S )ÚRfDetrHungarianMatcherr   é   Ú
class_costÚ	bbox_costÚ	giou_costÚmask_point_sample_ratioÚcost_mask_class_costÚcost_mask_dice_costc                 óx   •— t          ¦   «                              |||¦  «         || _        || _        || _        d S ©N)ÚsuperÚ__init__rb   Úcost_mask_classÚcost_mask_dice)Úselfr_   r`   ra   rb   rc   rd   Ú	__class__s          €r   rh   zRfDetrHungarianMatcher.__init__½   s?   ø€ õ 	‰Œ×Ò˜ Y°	Ñ:Ô:Ð:à'>ˆÔ$Ø3ˆÔØ1ˆÔÐÐr   c                 ó¢  ‡#‡$— |d         j         dd…         \  }}|d                              dd¦  «                             ¦   «         }|d                              dd¦  «        }|d                              dd¦  «        }t          j        d„ |D ¦   «         ¦  «        }	t          j        d	„ |D ¦   «         ¦  «        }
t          j        d
„ |D ¦   «         ¦  «        }d}d}d|z
  ||z  z  d|z
  dz                        ¦   «          z  }|d|z
  |z  z  |dz                        ¦   «          z  }|dd…|	f         |dd…|	f         z
  }t          j        |                     t          j        ¦  «        |
                     t          j        ¦  «        d¬¦  «         	                    |¦  «        }t          t          |¦  «        t          |
¦  «        ¦  «         }|j         dd…         \  }}||z  | j        z  }t          j        d|d|j        ¬¦  «        }|                     |j         d         dd¦  «        }|                     d¦  «        }t#          ||d¬¦  «        }t          j        |d¦  «        }|                     |j        ¦  «        }|                     |j         d         dd¦  «        }|                     d¦  «        }t#          ||dd¬¦  «        }t          j        |d¦  «        }t)          ||¦  «        }t+          ||¦  «        }| j        |z  | j        |z  z   | j        |z  z   | j        |z  z   | j        |z  z   }|                     ||d¦  «                             ¦   «         }t          j        |j        ¦  «        j        ||                     ¦   «         |                      ¦   «         z  <   d„ |D ¦   «         }g }||z  Š$| !                    ‰$d¬¦  «        }tE          |¦  «        D ]]Š#|‰#         } d„ tG          |  !                    |d¦  «        ¦  «        D ¦   «         }!‰#dk    r|!}Œ@ˆ#ˆ$fd„tI          ||!¦  «        D ¦   «         }Œ^d„ |D ¦   «         }"|"S )a  
        Differences:
        - out_prob = outputs["logits"].flatten(0, 1).sigmoid() instead of softmax
        - class_cost uses alpha and gamma
        - Additionally, mask cost is computed using pair-wise sigmoid cross entropy loss and dice loss
        rB   Nr   r   r   Ú
pred_boxesÚ
pred_masksc                 ó   — g | ]
}|d          ‘ŒS )Úclass_labels© ©Ú.0Úvs     r   ú
<listcomp>z2RfDetrHungarianMatcher.forward.<locals>.<listcomp>Ü   s   € ÐCÐCÐC°a  .Ô 1ÐCÐCÐCr   c                 ó   — g | ]
}|d          ‘ŒS )Úboxesrr   rs   s     r   rv   z2RfDetrHungarianMatcher.forward.<locals>.<listcomp>Ý   s   € Ð =Ð =Ð =°  7¤Ð =Ð =Ð =r   c                 ó   — g | ]
}|d          ‘ŒS ©Úmasksrr   rs   s     r   rv   z2RfDetrHungarianMatcher.forward.<locals>.<listcomp>Þ   s   € Ð!>Ð!>Ð!>° ! G¤*Ð!>Ð!>Ð!>r   g      Ð?r8   g:Œ0âŽyE>)ÚprG   FrI   )r.   r   Únearest©rJ   Úmoder.   c                 ó8   — g | ]}t          |d          ¦  «        ‘ŒS rz   ©Úlenrs   s     r   rv   z2RfDetrHungarianMatcher.forward.<locals>.<listcomp>  s"   € Ð2Ð2Ð2 Q•�Q�w”Z‘”Ð2Ð2Ð2r   rL   c                 ó>   — g | ]\  }}t          ||         ¦  «        ‘ŒS rr   r   )rt   ÚiÚcs      r   rv   z2RfDetrHungarianMatcher.forward.<locals>.<listcomp>  s)   € ÐsÐsÐs¹T¸QÀÕ2°1°Q´4Ñ8Ô8ÐsÐsÐsr   c                 óª   •— g | ]O\  }}t          j        |d          |d          ‰‰z  z   g¦  «        t          j        |d         |d         g¦  «        f‘ŒPS )r   r   )ÚnpÚconcatenate)rt   Úindice1Úindice2Úgroup_idÚgroup_num_queriess      €€r   rv   z2RfDetrHungarianMatcher.forward.<locals>.<listcomp>  sq   ø€ ð ð ð ñ
 )˜ õ œ¨°¬
°G¸A´JÐARÐU]ÑA]Ñ4]Ð'^Ñ_Ô_Ýœ¨°¬
°G¸A´JÐ'?Ñ@Ô@ððð ð r   c                 ó”   — g | ]E\  }}t          j        |t           j        ¬ ¦  «        t          j        |t           j        ¬ ¦  «        f‘ŒFS ))Údtype)r"   Ú	as_tensorÚint64)rt   r„   Újs      r   rv   z2RfDetrHungarianMatcher.forward.<locals>.<listcomp>  sR   € ð 
ð 
ð 
Ù_cÐ_`Ðbc�UŒ_˜Q¥e¤kÐ2Ñ2Ô2µE´OÀAÍUÌ[Ð4YÑ4YÔ4YÐZð
ð 
ð 
r   )%r!   r0   r/   r"   rS   ÚlogÚcdistÚtoÚfloat32Útype_asr   r   rb   rN   rH   Úrepeatr:   rA   r=   rŽ   r,   r3   r`   r_   ra   ri   rj   ÚviewÚcpuÚfinfoÚmaxÚisinfÚisnanÚsplitÚrangeÚ	enumerateÚzip)%rk   ÚoutputsÚtargetsÚ
group_detrÚ
batch_sizeÚnum_queriesÚout_probÚout_bboxÚ	out_masksÚ
target_idsÚtarget_bboxÚtarget_masksÚalphaÚgammaÚneg_cost_classÚpos_cost_classr_   r`   ra   ÚheightÚwidthrC   Úpoint_coordsÚpred_point_coordsÚpred_masks_logitsÚtarget_point_coordsri   rj   Úcost_matrixÚsizesÚindicesÚcost_matrix_listÚgroup_cost_matrixÚgroup_indicesÚmatched_indicesr‹   rŒ   s%                                      @@r   ÚforwardzRfDetrHungarianMatcher.forwardÌ   s°  øø€ ð #*¨(Ô"3Ô"9¸"¸1¸"Ô"=Ñˆ
�Kð ˜8Ô$×,Ò,¨Q°Ñ2Ô2×:Ò:Ñ<Ô<ˆØ˜<Ô(×0Ò0°°AÑ6Ô6ˆØ˜LÔ)×1Ò1°!°QÑ7Ô7ˆ	õ ”YÐCÐC¸7ÐCÑCÔCÑDÔDˆ
Ý”iÐ =Ð =°WÐ =Ñ =Ô =Ñ>Ô>ˆÝ”yÐ!>Ð!>°gÐ!>Ñ!>Ô!>Ñ?Ô?ˆð ˆØˆØ˜e™)¨°%©Ñ8¸aÀ(¹lÈTÑ>Q×=VÒ=VÑ=XÔ=XÐ<XÑYˆØ 1 x¡<°EÑ"9Ñ:ÀÈ4Á×?TÒ?TÑ?VÔ?VÐ>VÑWˆØ# A A A z MÔ2°^ÀAÀAÀAÀzÀMÔ5RÑRˆ
õ ”K §¢­E¬MÑ :Ô :¸K¿NºNÍ5Ì=Ñ<YÔ<YÐ]^Ð_Ñ_Ô_×gÒgÐhpÑqÔqˆ	õ )Õ)AÀ(Ñ)KÔ)KÕMeÐfqÑMrÔMrÑsÔsÐsˆ	ð !œ r¨ rÔ*‰ˆ�Ø˜e‘^ tÔ'CÑCˆ
Ý”z ! Z°¸9Ô;KÐLÑLÔLˆà(×/Ò/°	´ÀÔ0BÀAÀqÑIÔIÐØ×'Ò'¨Ñ*Ô*ˆ	Ý(¨Ð4EÐUZÐ[Ñ[Ô[ÐÝ!œMÐ*;¸WÑEÔEÐà#—’ y¤Ñ7Ô7ˆØ*×1Ò1°,Ô2DÀQÔ2GÈÈAÑNÔNÐØ#×-Ò-¨aÑ0Ô0ˆÝ# LÐ2EÐUZÐajÐkÑkÔkˆÝ”} \°7Ñ;Ô;ˆå>Ð?PÐR^Ñ_Ô_ˆÝ,Ð->ÀÑMÔMˆð ŒN˜YÑ&ØŒo 
Ñ*ñ+àŒn˜yÑ(ñ)ð Ô" _Ñ4ñ5ð Ô! NÑ2ñ	3ð 	ð "×&Ò& z°;ÀÑCÔC×GÒGÑIÔIˆõ BGÄÈ[ÔM^ÑA_ÔA_ÔAcˆ�K×%Ò%Ñ'Ô'¨+×*;Ò*;Ñ*=Ô*=Ñ=Ñ>ð 3Ð2¨'Ð2Ñ2Ô2ˆØˆØ'¨:Ñ5ÐØ&×,Ò,Ð->ÀAÐ,ÑFÔFÐÝ˜jÑ)Ô)ð 	ð 	ˆHØ 0°Ô :ÐØsÐsÅYÐO`×OfÒOfÐglÐnpÑOqÔOqÑErÔErÐsÑsÔsˆMØ˜1Š}ˆ}Ø'��ðð ð ð ð õ
 -0°¸Ñ,GÔ,Gðñ ô ��ð
ð 
Øgnð
ñ 
ô 
ˆð Ðr   )r   r   r   r^   r   r   )
Ú__name__Ú
__module__Ú__qualname__ÚfloatrM   rh   r"   Úno_gradr¾   Ú__classcell__©rl   s   @r   r]   r]   ¼   sº   ø€ € € € € ð ØØØ')Ø&'Ø%&ð2ð 2àð2ð ð2ð ð	2ð
 "%ð2ð $ð2ð #ð2ð 2ð 2ð 2ð 2ð 2ð €U„]�_„_ðUð Uñ „_ðUð Uð Uð Uð Ur   r]   c                   ó*   ‡ — e Zd Zˆ fd„Zd„ Zd„ Zˆ xZS )ÚRfDetrImageLossc                 ó`   •— t          ¦   «                              |||||¦  «         || _        d S rf   )rg   rh   rb   )rk   ÚmatcherÚnum_classesÚfocal_alphaÚlossesr¤   rb   rl   s          €r   rh   zRfDetrImageLoss.__init__&  s1   ø€ Ý‰Œ×Ò˜ +¨{¸FÀJÑOÔOÐOØ'>ˆÔ$Ð$Ð$r   c                 ó–  — d|vrt          d¦  «        ‚|                      |¦  «        }|d         |         }|                     ¦   «         dk    r)t          j        |¦  «        t          j        |¦  «        dœS t          j        d„ t          ||¦  «        D ¦   «         d¬¦  «        }|                     d¦  «        }|                     d¦  «                             ¦   «         }t          |j
        d         |j
        d         |j
        d	         z  | j        z  ¦  «        }t          j        ¦   «         5  t          ||d
d¦  «        }	t          ||	dd¬¦  «                             d¦  «        }
ddd¦  «         n# 1 swxY w Y   t          ||	d¬¦  «                             d¦  «        }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].
        ro   z#No predicted masks found in outputsr   )Úloss_mask_ceÚloss_mask_dicec                 ó6   — g | ]\  }\  }}|d          |         ‘ŒS rz   rr   )rt   ÚtÚ_r‘   s       r   rv   z.RfDetrImageLoss.loss_masks.<locals>.<listcomp><  s(   € Ð!WÐ!WÐ!W±I°A±v¸¸1 ! G¤*¨Q¤-Ð!WÐ!WÐ!Wr   rL   r   éþÿÿÿr.   r7   g      è?Fr}   r~   NrI   )ÚKeyErrorÚ_get_source_permutation_idxÚnumelr"   r$   rS   r¡   r:   rÂ   r›   r!   rb   rÃ   r[   rA   r=   r   r
   )rk   r¢   r£   r¹   rT   Ú
source_idxÚsource_masksr¬   rC   r³   Úpoint_labelsrV   rÌ   s                r   Ú
loss_maskszRfDetrImageLoss.loss_masks*  s*  € ð ˜wÐ&Ð&ÝÐ@ÑAÔAÐAà×5Ò5°gÑ>Ô>ˆ
Ø˜|Ô,¨ZÔ8ˆØ×ÒÑÔ 1Ò$Ð$å %Ô 0°Ñ >Ô >Ý"'Ô"2°<Ñ"@Ô"@ðð ð õ ”yÐ!WÐ!WÅÀWÈgÑAVÔAVÐ!WÑ!WÔ!WÐ]^Ð_Ñ_Ô_ˆà#×-Ò-¨aÑ0Ô0ˆØ#×-Ò-¨aÑ0Ô0×6Ò6Ñ8Ô8ˆõ ØÔ˜rÔ" LÔ$6°rÔ$:¸\Ô=OÐPRÔ=SÑ$SÐW[ÔWsÑ$sñ
ô 
ˆ
õ Œ]‰_Œ_ð 	tð 	tå:¸<ÈÐUVÐX\Ñ]Ô]ˆLå'¨°lÐRWÐ^gÐhÑhÔh×pÒpÐqrÑsÔsˆLð		tð 	tð 	tñ 	tô 	tð 	tð 	tð 	tð 	tð 	tð 	tøøøð 	tð 	tð 	tð 	tõ $ L°,ÈeÐTÑTÔT×\Ò\Ð]^Ñ_Ô_ˆõ 7°|À\ÐS\Ñ]Ô]Ý'¨°lÀIÑNÔNð
ð 
ˆð ˆs   Ä19E6Å6E:Å=E:c           
      ó
  ‡— | j         r| j        nd}d„ |                     ¦   «         D ¦   «         }|                      |||¦  «        }t	          d„ |D ¦   «         ¦  «        }||z  }t          j        |gt
          j        t          t          | 
                    ¦   «         ¦  «        ¦  «        j        ¬¦  «        }d}t          j        ¦   «         rKt          j        ¦   «         r8t          j        |t          j        j        ¬¦  «         t          j        ¦   «         }t          j        ||z  d¬¦  «                             ¦   «         }i }| j        D ].}	|                     |                      |	||||¦  «        ¦  «         Œ/d|v rŠt1          |d         ¦  «        D ]t\  Š}
|                      |
||¦  «        }| j        D ]P}	|                      |	|
|||¦  «        }ˆfd„|                     ¦   «         D ¦   «         }|                     |¦  «         ŒQŒud	|v rv|d	         }|                      |||¬
¦  «        }| j        D ]N}	|                      |	||||¦  «        }d„ |                     ¦   «         D ¦   «         }|                     |¦  «         Œ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.
        r   c                 ó2   — i | ]\  }}|d k    ¯|dk    ¯||“ŒS )Úenc_outputsÚauxiliary_outputsrr   ©rt   rK   ru   s      r   ú
<dictcomp>z+RfDetrImageLoss.forward.<locals>.<dictcomp>`  s:   € ð '
ð '
ð '
Ù�Q˜°°]Ò0BÐ0BÀqÐL_ÒG_ÐG_ˆAˆqÐG_ÐG_ÐG_r   c              3   ó@   K  — | ]}t          |d          ¦  «        V — ŒdS )rq   Nr�   )rt   rÑ   s     r   ú	<genexpr>z*RfDetrImageLoss.forward.<locals>.<genexpr>h  s/   è è € Ð@Ð@°1�˜A˜nÔ-Ñ.Ô.Ð@Ð@Ð@Ð@Ð@Ð@r   )rŽ   rH   )Úop)ÚminrÞ   c                 ó(   •— i | ]\  }}|d ‰› �z   |“ŒS ©rÒ   rr   ©rt   rK   ru   r„   s      €r   rà   z+RfDetrImageLoss.forward.<locals>.<dictcomp>}  s)   ø€ ÐHÐHÐH±°°A˜a ' a ' '™k¨1ÐHÐHÐHr   rÝ   )r¤   c                 ó    — i | ]\  }}|d z   |“ŒS ©Ú_encrr   rß   s      r   rà   z+RfDetrImageLoss.forward.<locals>.<dictcomp>…  s"   € ÐCÐCÐC©D¨A¨q˜!˜f™* aÐCÐCÐCr   )Útrainingr¤   ÚitemsrÉ   r   r"   r�   rÂ   ÚnextÚiterÚvaluesrH   ÚdistÚis_availableÚis_initializedÚ
all_reduceÚReduceOpÚSUMÚget_world_sizeÚclampÚitemrÌ   ÚupdateÚget_lossr    )rk   r¢   r£   r¤   Úoutputs_without_aux_and_encr¹   rT   Ú
world_sizerÌ   r   rÞ   Úl_dictrÝ   r„   s                @r   r¾   zRfDetrImageLoss.forwardT  s©  ø€ ð )-¬Ð<�T”_�_¸1ˆ
ð'
ð '
Ø$Ÿ]š]™_œ_ð'
ñ '
ô '
Ð#ð
 —,’,Ð:¸GÀZÑPÔPˆõ Ð@Ð@¸Ð@Ñ@Ô@Ñ@Ô@ˆ	Ø 
Ñ*ˆ	Ý”O Y Kµu´{Í4ÕPTÐU\×UcÒUcÑUeÔUeÑPfÔPfÑKgÔKgÔKnÐoÑoÔoˆ	Øˆ
ÝÔÑÔð 	/¥4Ô#6Ñ#8Ô#8ð 	/ÝŒO˜I­$¬-Ô*;Ð<Ñ<Ô<Ð<ÝÔ,Ñ.Ô.ˆJÝ”K 	¨JÑ 6¸AÐ>Ñ>Ô>×CÒCÑEÔEˆ	ð ˆØ”Kð 	Uð 	UˆDØ�MŠM˜$Ÿ-š-¨¨g°wÀÈÑSÔSÑTÔTÐTÐTð  'Ð)Ð)Ý(1°'Ð:MÔ2NÑ(OÔ(Oð *ð *Ñ$�Ð$ØŸ,š,Ð'8¸'À:ÑNÔN�Ø œKð *ð *�DØ!Ÿ]š]¨4Ð1BÀGÈWÐV_Ñ`Ô`�FØHÐHÐHÐH¸¿º¹¼ÐHÑHÔH�FØ—M’M &Ñ)Ô)Ð)Ð)ð*ð
 ˜GÐ#Ð#Ø! -Ô0ˆKØ—l’l ;°ÀJ�lÑOÔOˆGØœð &ð &�ØŸš t¨[¸'À7ÈIÑVÔV�ØCÐC°F·L²L±N´NÐCÑCÔC�Ø—’˜fÑ%Ô%Ð%Ð%àˆr   )r¿   rÀ   rÁ   rh   rÚ   r¾   rÄ   rÅ   s   @r   rÇ   rÇ   %  sW   ø€ € € € € ð?ð ?ð ?ð ?ð ?ð(ð (ð (ðT4ð 4ð 4ð 4ð 4ð 4ð 4r   rÇ   c                 óh   — d„ t          | d d…         |d d…         |d d…         ¦  «        D ¦   «         S )Nc                 ó"   — g | ]\  }}}|||d œ‘ŒS )©rB   rn   ro   rr   )rt   ÚaÚbr…   s       r   rv   z!_set_aux_loss.<locals>.<listcomp>�  s8   € ð ð ð áˆAˆq�!ð  A°QÐ7Ð7ðð ð r   r.   )r¡   )Úoutputs_classÚoutputs_coordÚoutputs_maskss      r   Ú_set_aux_lossr  ‹  sM   € ðð å˜=¨¨"¨Ô-¨}¸S¸b¸SÔ/AÀ=ÐQTÐRTÐQTÔCUÑVÔVðñ ô ð r   c                 óP  ‡‡‡— t          |j        |j        |j        |j        |j        |j        ¬¦  «        }g d¢}t          ||j        |j	        ||j
        |j        ¬¦  «        }|                     |¦  «         i }d }| |d<   ||d<   ||d<   |	|
|dœ|d<   |j        rt          |||¦  «        }||d	<    |||¦  «        Š|j        |j        d
œŠ|j        ‰d<   |j        ‰d<   |j        ‰d<   |j        r•i }t#          |j        dz
  ¦  «        D ]5Š|                     ˆfd„‰                     ¦   «         D ¦   «         ¦  «         Œ6|                     d„ ‰                     ¦   «         D ¦   «         ¦  «         ‰                     |¦  «         t+          ˆˆfd„‰D ¦   «         ¦  «        }|‰|fS )N)r_   r`   ra   rb   rc   rd   )r   rx   Úcardinalityr{   )rÉ   rÊ   rË   rÌ   r¤   rb   rB   rn   ro   r   rÝ   rÞ   )Úloss_ceÚ	loss_bboxÚ	loss_giourÎ   rÏ   r   c                 ó(   •— i | ]\  }}|d ‰› �z   |“ŒS ræ   rr   rç   s      €r   rà   z-RfDetrForSegmentationLoss.<locals>.<dictcomp>Î  s)   ø€ Ð#SÐ#SÐ#S±t°q¸! A¨¨A¨¨¡K°Ð#SÐ#SÐ#Sr   c                 ó    — i | ]\  }}|d z   |“ŒS ré   rr   rß   s      r   rà   z-RfDetrForSegmentationLoss.<locals>.<dictcomp>Ï  s"   € ÐNÐNÐN±$°!°Q  F¡
¨AÐNÐNÐNr   c              3   óB   •K  — | ]}|‰v ¯‰|         ‰|         z  V — Œd S rf   rr   )rt   rK   Ú	loss_dictÚweight_dicts     €€r   râ   z,RfDetrForSegmentationLoss.<locals>.<genexpr>Ñ  s:   øè è € ÐTÐT°À1ÈÐCSÐCSˆy˜Œ|˜k¨!œnÑ,ÐCSÐCSÐCSÐCSÐTÐTr   )r]   r_   r`   ra   rb   Úmask_class_loss_coefficientÚmask_dice_loss_coefficientrÇ   Ú
num_labelsrË   r¤   r”   Úauxiliary_lossr  Úclass_loss_coefficientÚbbox_loss_coefficientÚgiou_loss_coefficientrŸ   Údecoder_layersrù   rì   r   )rB   r   rH   rn   ro   Úconfigr  r  r  Úenc_outputs_classÚenc_outputs_coordÚenc_outputs_masksr?   rÉ   rÌ   r   Úoutputs_lossrÞ   Úaux_weight_dictr   r„   r  r  s                       @@@r   ÚRfDetrForSegmentationLossr  “  s*  øøø€ õ  %ØÔ$ØÔ"ØÔ"Ø &Ô >Ø#Ô?Ø"Ô=ðñ ô €Gð 9Ð8Ð8€FÝØØÔ%ØÔ&ØØÔ$Ø &Ô >ðñ ô €Ið ‡L‚L�ÑÔÐà€LØÐØ#€L�ÑØ!+€L�ÑØ!+€L�Ñà#Ø'Ø'ð#ð #€L�Ñð
 Ôð >Ý)¨-¸ÈÑVÔVÐØ,=ˆÐ(Ñ)à�	˜,¨Ñ/Ô/€Ià$Ô;È&ÔJfÐgÐg€KØ%Ô;€K�ÑØ"(Ô"D€K�ÑØ$*Ô$E€KÐ Ñ!ØÔð ,ØˆÝ�vÔ,¨qÑ0Ñ1Ô1ð 	Uð 	UˆAØ×"Ò"Ð#SÐ#SÐ#SÐ#S¸{×?PÒ?PÑ?RÔ?RÐ#SÑ#SÔ#SÑTÔTÐTÐTØ×ÒÐNÐN¸+×:KÒ:KÑ:MÔ:MÐNÑNÔNÑOÔOÐOØ×Ò˜?Ñ+Ô+Ð+ÝÐTÐTÐTÐTÐT°iÐTÑTÔTÑTÔT€DØ�Ð-Ð-Ð-r   )F)NNNNNN)Únumpyr‡   r"   Útorch.distributedÚdistributedrð   r   r   Úimage_transformsr   Úutilsr   Úloss_for_object_detectionr	   r
   r   Úloss_lw_detrr   Úscipy.optimizer   rM   r   r,   r3   rA   rÂ   r[   r]   rÇ   r  r  rr   r   r   ú<module>r(     sŒ  ðð Ð Ð Ð Ø €€€Ø  Ð  Ð  Ð  Ð  Ð  Ø Ð Ð Ð Ð Ð Ð Ð à 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø &Ð &Ð &Ð &Ð &Ð &ðð ð ð ð ð ð ð ð ð ð
 *Ð )Ð )Ð )Ð )Ð )ð ÐÑÔð 5Ø4Ð4Ð4Ð4Ð4Ð4ð u¤|ð ¸U¼\ð ÐVYð Ð^cÔ^jð ð ð ð ð(°´ð ÀuÄ|ð ÐX]ÔXdð ð ð ð ð8 ð °ð ¸6ð ð ð ð ð0 LQðð Ø”LðØ5:´\ðà
„\ðð ð ð ðB.Øð.Ø #ð.Ø7:ð.ØUZð.àð.ð .ð .ð .ðbfð fð fð fð fÐ-ñ fô fð fðRcð cð cð cð c�oñ cô cð cðLð ð ð ØØØØØð?.ð ?.ð ?.ð ?.ð ?.ð ?.r   