Ë
    S^(hµ+  ã                   óè   — d dl Z d dlmZ ddlmZ ddlmZ ddlmZm	Z	m
Z
mZ  e«       rd dlmZ 	 	 dde j                  d	e j                  d
ededef
d„Z G d„ de«      Z G d„ de	«      Z	 	 	 	 dd„Zy)é    Né   )Úcenter_to_corners_format)Úis_scipy_availableé   )ÚHungarianMatcherÚ	ImageLossÚ_set_aux_lossÚgeneralized_box_iou)Úlinear_sum_assignmentÚinputsÚtargetsÚ	num_boxesÚalphaÚgammac                 óú   — | j                  «       }t        j                  j                  | |d¬«      }||z  d|z
  d|z
  z  z   }|d|z
  |z  z  }|dk\  r||z  d|z
  d|z
  z  z   }	|	|z  }|j	                  «       |z  S )aJ  
    Loss used in RetinaNet for dense detection: https://arxiv.org/abs/1708.02002.

    Args:
        inputs (`torch.FloatTensor` of arbitrary shape):
            The predictions for each example.
        targets (`torch.FloatTensor` with the same shape as `inputs`)
            A tensor storing the binary classification label for each element in the `inputs` (0 for the negative class
            and 1 for the positive class).
        num_boxes (`int`):
            The total number of boxes in the batch.
        alpha (`float`, *optional*, defaults to 0.25):
            Optional weighting factor in the range (0,1) to balance positive vs. negative examples.
        gamma (`int`, *optional*, defaults to 2):
            Exponent of the modulating factor (1 - p_t) to balance easy vs hard examples.

    Returns:
        Loss tensor
    Únone)Ú	reductionr   r   )ÚsigmoidÚnnÚ
functionalÚ binary_cross_entropy_with_logitsÚsum)
r   r   r   r   r   ÚprobÚce_lossÚp_tÚlossÚalpha_ts
             úc/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/loss/loss_grounding_dino.pyÚsigmoid_focal_lossr      s–   € ð4 �>‰>Ó€DÜ�m‰m×<Ñ<¸VÀWÐX^Ð<Ó_€Gà
�‰.˜A ™H¨¨W©Ñ5Ñ
5€CØ�q˜3‘w 5Ñ(Ñ)€Dà�‚zØ˜'‘/ Q¨¡Y°1°w±;Ñ$?Ñ?ˆØ˜‰~ˆà�8‰8‹:˜	Ñ!Ð!ó    c                   ó:   — e Zd Z ej                  «       d„ «       Zy)ÚGroundingDinoHungarianMatcherc           	      óð  — |d   j                   dd \  }}|d   j                  dd«      j                  «       }|d   j                  dd«      }|d   }t        j                  t        ||«      D ��	cg c]  \  }}	||	d      ‘Œ c}	}«      }||j                  d	d
¬«      z  }t        j                  |D �
cg c]  }
|
d   ‘Œ	 c}
«      }d}d}d|z
  ||z  z  d|z
  dz   j                  «        z  }|d|z
  |z  z  |dz   j                  «        z  }||z
  |j                  «       z  }t        j                  ||d¬«      }t        t        |«      t        |«      «       }| j                  |z  | j                  |z  z   | j                  |z  z   }|j                  ||d	«      j!                  «       }|D �
cg c]  }
t#        |
d   «      ‘Œ }}
t%        |j'                  |d	«      «      D ��cg c]  \  }}t)        ||   «      ‘Œ }}}|D ��cg c]O  \  }}t        j*                  |t        j,                  ¬«      t        j*                  |t        j,                  ¬«      f‘ŒQ c}}S c c}	}w c c}
w c c}
w c c}}w c c}}w )aô  
        Args:
            outputs (`dict`):
                A dictionary 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.
                * "label_maps": Tuple of tensors of dim [num_classes, hidden_dim].
            targets (`List[dict]`):
                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:
            `List[Tuple]`: 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)
        ÚlogitsNr   r   r   Ú
pred_boxesÚ
label_mapsÚclass_labelséÿÿÿÿT)ÚdimÚkeepdimÚboxesç      Ð?ç       @g:Œ0âŽyE>)Úp)Údtype)ÚshapeÚflattenr   ÚtorchÚcatÚzipr   ÚlogÚtÚcdistr
   r   Ú	bbox_costÚ
class_costÚ	giou_costÚviewÚcpuÚlenÚ	enumerateÚsplitr   Ú	as_tensorÚint64)ÚselfÚoutputsr   Ú
batch_sizeÚnum_queriesÚout_probÚout_bboxr&   Ú	label_mapÚtargetÚvÚtarget_bboxr   r   Úneg_cost_classÚpos_cost_classr9   r8   r:   Úcost_matrixÚsizesÚiÚcÚindicesÚjs                            r   Úforwardz%GroundingDinoHungarianMatcher.forwardD   sf  € ð, #*¨(Ñ"3×"9Ñ"9¸"¸1Ð"=Ñˆ
�Kð ˜8Ñ$×,Ñ,¨Q°Ó2×:Ñ:Ó<ˆØ˜<Ñ(×0Ñ0°°AÓ6ˆØ˜\Ñ*ˆ
ô —Y‘YÔ[^Ð_iÐkrÓ[s×tÑFWÀiÐQW 	¨&°Ñ*@Ó AÓtÓuˆ
à *§.¡.°RÀ .Ó"FÑFˆ
ô —i‘i°WÖ =°  7£Ò =Ó>ˆð ˆØˆØ˜e™)¨°%©Ñ8¸aÀ(¹lÈTÑ>Q×=VÑ=VÓ=XÐ<XÑYˆØ 1 x¡<°EÑ"9Ñ:ÀÈ4Á×?TÑ?TÓ?VÐ>VÑWˆà$ ~Ñ5¸¿¹»ÑGˆ
ô —K‘K ¨+¸Ô;ˆ	ô )Ô)AÀ(Ó)KÔMeÐfqÓMrÓsÐsˆ	ð —n‘n yÑ0°4·?±?ÀZÑ3OÑOÐRV×R`ÑR`ÐclÑRlÑlˆØ!×&Ñ& z°;ÀÓC×GÑGÓIˆà*1Ö2 Q”�Q�w‘Z•Ð2ˆÐ2Ü;DÀ[×EVÑEVÐW\Ð^`ÓEaÓ;b×c±4°1°aÔ(¨¨1©Õ.ÐcˆÑcØkr×sÑcgÐcdÐfg”—‘ ¬%¯+©+Ô6¼¿¹ÈÔQV×Q\ÑQ\Ô8]Ò^ÓsÐsùó7  uùò
 !>ùò( 3ùÛcùÛss   Á1I
Â4I"Æ/I'Ç$I,ÈAI2N)Ú__name__Ú
__module__Ú__qualname__r2   Úno_gradrT   © r    r   r"   r"   C   s   „ Ø€U‡]�]ƒ_ñ8tó ñ8tr    r"   c                   ó"   — e Zd ZdZd„ Zd„ Zd„ Zy)ÚGroundingDinoImageLossa†  
    This class computes the losses for `GroundingDinoForObjectDetection`. 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 (`GroundingDinoHungarianMatcher`):
            Module able to compute a matching between targets and proposals.
        focal_alpha (`float`):
            Alpha parameter in focal loss.
        losses (`List[str]`):
            List of all the losses to be applied. See `get_loss` for a list of all available losses.
    c                 ól   — t         j                  j                  | «       || _        || _        || _        y ©N)r   ÚModuleÚ__init__ÚmatcherÚfocal_alphaÚlosses)rB   r`   ra   rb   s       r   r_   zGroundingDinoImageLoss.__init__�   s*   € Ü
�	‰	×Ñ˜4Ô ØˆŒØ&ˆÔØˆ�r    c                 óô  — |d   }t        j                  t        t        ||«      «      D ����cg c]2  \  }\  }\  }}|dkD  r|d   |   t	        |d   |   «      z   n|d   |   ‘Œ4 c}}}}«      }	t        j                  |d   d¬«      }
| j                  |«      }t        j                  ||j                  t         j                  ¬«      }|
|	   j                  t         j                  «      ||<   |S c c}}}}w )z>
        Create one_hot based on the matching indices
        r$   r   r'   r&   )r)   )Údevicer/   )
r2   r3   r>   r4   r=   Ú_get_source_permutation_idxÚ
zeros_likerd   ÚlongÚto)rB   rC   r   rR   r$   rP   rI   Ú_ÚJr'   r&   ÚidxÚtarget_classes_onehots                r   Ú_get_target_classes_one_hotz2GroundingDinoImageLoss._get_target_classes_one_hot•   sü   € ð ˜Ñ"ˆä—y‘yô ,5´S¸À'Ó5JÓ+K÷ñ á'�AÑ'˜¡  Að NOÐQRÊU��~Ñ& qÑ)¬C°¸Ñ0EÀaÑ0HÓ,IÒIÐX^Ð_mÑXnÐopÑXqÑqõó
ˆô —Y‘Y˜w |Ñ4¸!Ô<ˆ
à×.Ñ.¨wÓ7ˆÜ %× 0Ñ 0°ÀÇÁÔUZ×U_ÑU_Ô `ÐØ%/°Ñ%=×%@Ñ%@ÄÇÁÓ%LÐ˜cÑ"à$Ð$ùõs   ¯7C2c                 ó0  — d|vrt        d«      ‚d|vrt        d«      ‚| j                  |||«      }|d   }|d   }t        j                  ||«      }t        j                  ||«      }|j	                  «       }t        |||| j                  d¬«      }d|i}	|	S )z 
        Classification loss (Binary focal loss) targets dicts must contain the key "class_labels" containing a tensor
        of dim [nb_target_boxes]
        r$   z#No logits were found in the outputsÚ	text_maskz&No text_mask were found in the outputsr   )r   r   r   r   r   Úloss_ce)ÚKeyErrorrm   r2   Úmasked_selectÚfloatr   ra   )
rB   rC   r   rR   r   rl   Úsource_logitsro   rp   rb   s
             r   Úloss_labelsz"GroundingDinoImageLoss.loss_labels©   s½   € ð
 ˜7Ñ"ÜÐ@ÓAÐAØ˜gÑ%ÜÐCÓDÐDà $× @Ñ @ÀÈ'ÐSZÓ [ÐØ Ñ)ˆØ˜KÑ(ˆ	ô ×+Ñ+¨M¸9ÓEˆÜ %× 3Ñ 3Ð4IÈ9Ó UÐà 5× ;Ñ ;Ó =ÐÜ$Ø Ø)ØØ×"Ñ"Øô
ˆð ˜WÐ%ˆàˆr    N)rU   rV   rW   Ú__doc__r_   rm   ru   rY   r    r   r[   r[   €   s   „ ñòò%ó(r    r[   c           
      ó  ‡‡— t        |j                  |j                  |j                  ¬«      }g d¢}t	        ||j
                  |¬«      }|j                  |«       i }| |d<   ||d<   ||d<   ||d<   d }|j                  r"t        ||«      }|D ]  }||d<   ||d<   Œ ||d<    |||«      Š|j                  rG|	|
||d	œ} |||«      }|j                  «       D ��ci c]  \  }}|d
z   |“Œ }}}‰j                  |«       d|j                  |j                  dœŠ|j                  r7‰j                  «       D ��ci c]  \  }}|d
z   |“Œ }}}‰j                  |«       |j                  rii }t        |j                  dz
  «      D ];  }|j                  ‰j                  «       D ��ci c]  \  }}|d|› �z   |“Œ c}}«       Œ= ‰j                  |«       t!        ˆˆfd„‰j#                  «       D «       «      }|‰|fS c c}}w c c}}w c c}}w )N)r9   r8   r:   )Úlabelsr+   Úcardinality)r`   ra   rb   r$   r%   r&   ro   Úauxiliary_outputs)r$   r%   r&   ro   Ú_encr-   )rp   Ú	loss_bboxÚ	loss_giour   ri   c              3   ó>   •K  — | ]  }|‰v sŒ‰|   ‰|   z  –— Œ y ­wr]   rY   )Ú.0ÚkÚ	loss_dictÚweight_dicts     €€r   ú	<genexpr>z6GroundingDinoForObjectDetectionLoss.<locals>.<genexpr>  s%   øè ø€ Ò[°È!È{ÒJZˆy˜‰|˜k¨!™nÕ,Ñ[ùs   ƒ	�)r"   r9   r8   r:   r[   ra   rh   Úauxiliary_lossr	   Ú	two_stageÚitemsÚupdateÚbbox_loss_coefficientÚgiou_loss_coefficientÚrangeÚdecoder_layersr   Úkeys)r$   rx   rd   r%   Úconfigr&   ro   Úoutputs_classÚoutputs_coordÚencoder_logitsÚencoder_pred_boxesr`   rb   Ú	criterionÚoutputs_lossrz   Ú
aux_outputÚencoder_outputs_lossÚencoder_loss_dictr€   rJ   Úenc_weight_dictÚaux_weight_dictrP   r   r�   r‚   s                            @@r   Ú#GroundingDinoForObjectDetectionLossr™   É   sG  ù€ ô ,Ø×$Ñ$°×0@Ñ0@ÈF×L\ÑL\ô€Gò 0€FÜ&ØØ×&Ñ&Øô€Ið
 ‡L�L�Ôà€LØ#€L�ÑØ!+€L�ÑØ!+€L�ÑØ )€L�ÑàÐØ×ÒÜ)¨-¸ÓGÐØ+ò 	0ˆJØ'1ˆJ�|Ñ$Ø&/ˆJ�{Ò#ð	0ð ->ˆÐ(Ñ)á˜,¨Ó/€Ià×Òà$Ø,Ø$Ø"ñ	 
Ðñ &Ð&:¸FÓCÐØ7H×7NÑ7NÓ7P×Q©t¨q°!˜Q ™Z¨™]ÐQÐÑQØ×ÑÐ*Ô+ð Ø×1Ñ1Ø×1Ñ1ñ€Kð ×ÒØ5@×5FÑ5FÓ5H×I©T¨Q°˜1˜v™: q™=ÐIˆÑIØ×Ñ˜?Ô+à×ÒØˆÜ�v×,Ñ,¨qÑ0Ó1ò 	UˆAØ×"Ñ"¸{×?PÑ?PÓ?R×#S±t°q¸! A¨!¨A¨3¨¡K°¡NÓ#SÕTð	Uà×Ñ˜?Ô+äÔ[°i·n±nÓ6FÔ[Ó[€DØ�Ð-Ð-Ð-ùó) Rùó Jùó $Ts   ÃG8Ä7G>Æ"H)r,   r   )NNNN)r2   Útorch.nnr   Úimage_transformsr   Úutilsr   Úloss_for_object_detectionr   r   r	   r
   Úscipy.optimizer   ÚTensorÚintrs   r   r"   r[   r™   rY   r    r   ú<module>r¡      s›   ðó Ý å 7Ý &ß fÓ fñ ÔÝ4ð Øñ$"Ø�L‰Lð$"à�\‰\ð$"ð ð$"ð ð	$"ð
 ó$"ôN:tÐ$4ô :tôzF˜Yô Fðb ØØØôF.r    