Ë
    Fêñi]ã  ã                  óî  — d dl mZ d dl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mZ d dlmZmZmZ d dlmZmZmZmZmZ d dlmZ dd	lmZmZ dd
lmZmZ  G d„ dej@                  «      Z! G d„ dej@                  «      Z" G d„ dej@                  «      Z# G d„ dej@                  «      Z$ G d„ dej@                  «      Z% G d„ de$«      Z& G d„ dej@                  «      Z' G d„ dej@                  «      Z( G d„ dej@                  «      Z) G d„ d«      Z* G d„ d e*«      Z+ G d!„ d"e*«      Z, G d#„ d$e,«      Z- G d%„ d&«      Z. G d'„ d(e*«      Z/ G d)„ d*«      Z0 G d+„ d,«      Z1 G d-„ d.«      Z2 G d/„ d0e2«      Z3y)1é    )ÚannotationsN)ÚAny)Ú	OKS_SIGMAÚ
RLE_WEIGHT)Ú	crop_maskÚ	xywh2xyxyÚ	xyxy2xywh)ÚRotatedTaskAlignedAssignerÚTaskAlignedAssignerÚ	dist2bboxÚ	dist2rboxÚmake_anchors)Úautocasté   )Úbbox_iouÚprobiou)Ú	bbox2distÚ	rbox2distc                  ó.   ‡ — e Zd ZdZddˆ fd„Zdd„Zˆ xZS )ÚVarifocalLossaä  Varifocal loss by Zhang et al.

    Implements the Varifocal Loss function for addressing class imbalance in object detection by focusing on
    hard-to-classify examples and balancing positive/negative samples.

    Attributes:
        gamma (float): The focusing parameter that controls how much the loss focuses on hard-to-classify examples.
        alpha (float): The balancing factor used to address class imbalance.

    References:
        https://arxiv.org/abs/2008.13367
    c                ó>   •— t         ‰| �  «        || _        || _        y)zJInitialize the VarifocalLoss class with focusing and balancing parameters.N)ÚsuperÚ__init__ÚgammaÚalpha©Úselfr   r   Ú	__class__s      €úX/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/ultralytics/utils/loss.pyr   zVarifocalLoss.__init__#   s   ø€ ä‰ÑÔØˆŒ
Øˆ�
ó    c                óv  — | j                   |j                  «       j                  | j                  «      z  d|z
  z  ||z  z   }t	        d¬«      5  t        j                  |j                  «       |j                  «       d¬«      |z  j                  d«      j                  «       }ddd«       |S # 1 sw Y   S xY w)z<Compute varifocal loss between predictions and ground truth.r   F)ÚenabledÚnone©Ú	reductionN)
r   ÚsigmoidÚpowr   r   ÚFÚ binary_cross_entropy_with_logitsÚfloatÚmeanÚsum)r   Ú
pred_scoreÚgt_scoreÚlabelÚweightÚlosss         r   ÚforwardzVarifocalLoss.forward)   s¡   € à—‘˜j×0Ñ0Ó2×6Ñ6°t·z±zÓBÑBÀaÈ%ÁiÑPÐS[Ð^cÑScÑcˆÜ˜eÔ$ñ 	ä×3Ñ3°J×4DÑ4DÓ4FÈÏÉÓHXÐdjÔkÐntÑtß‘�a“ß‘“ð ÷	ð ˆ÷	ð ˆús   ÁAB.Â.B8)ç       @g      è?©r   r*   r   r*   )r-   útorch.Tensorr.   r5   r/   r5   Úreturnr5   ©Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r2   Ú__classcell__©r   s   @r   r   r      s   ø„ ñö÷	r    r   c                  ó.   ‡ — e Zd ZdZddˆ fd„Zdd„Zˆ xZS )Ú	FocalLossaã  Wraps focal loss around existing loss_fcn(), i.e. criteria = FocalLoss(nn.BCEWithLogitsLoss(), gamma=1.5).

    Implements the Focal Loss function for addressing class imbalance by down-weighting easy examples and focusing on
    hard negatives during training.

    Attributes:
        gamma (float): The focusing parameter that controls how much the loss focuses on hard-to-classify examples.
        alpha (torch.Tensor): The balancing factor used to address class imbalance.
    c                ód   •— t         ‰| �  «        || _        t        j                  |«      | _        y)zBInitialize FocalLoss class with focusing and balancing parameters.N)r   r   r   ÚtorchÚtensorr   r   s      €r   r   zFocalLoss.__init__@   s%   ø€ ä‰ÑÔØˆŒ
Ü—\‘\ %Ó(ˆ�
r    c                óÚ  — t        j                  ||d¬«      }|j                  «       }||z  d|z
  d|z
  z  z   }d|z
  | j                  z  }||z  }| j                  dkD  j                  «       r`| j                  j                  |j                  |j                  ¬«      | _        || j                  z  d|z
  d| j                  z
  z  z   }||z  }|j                  d«      j                  «       S )zACalculate focal loss with modulating factors for class imbalance.r#   r$   r   ç      ð?r   ©ÚdeviceÚdtype)r(   r)   r&   r   r   ÚanyÚtorF   rG   r+   r,   )r   Úpredr/   r1   Ú	pred_probÚp_tÚmodulating_factorÚalpha_factors           r   r2   zFocalLoss.forwardF   sÖ   € ä×1Ñ1°$¸ÈÔPˆð
 —L‘L“Nˆ	Ø�iÑ 1 u¡9°°Y±Ñ"?Ñ?ˆØ  3™Y¨4¯:©:Ñ5ÐØÐ!Ñ!ˆØ�J‰J˜‰N×ÑÔ!ØŸ™Ÿ™¨d¯k©kÀÇÁ˜ÓLˆDŒJØ  4§:¡:Ñ-°°U±¸qÀ4Ç:Á:¹~Ñ0NÑNˆLØ�LÑ ˆDØ�y‰y˜‹|×ÑÓ!Ð!r    )g      ø?g      Ð?r4   )rJ   r5   r/   r5   r6   r5   r7   r=   s   @r   r?   r?   5   s   ø„ ñö)÷"r    r?   c                  ó.   ‡ — e Zd ZdZddˆ fd„Zdd„Zˆ xZS )ÚDFLossz<Criterion class for computing Distribution Focal Loss (DFL).c                ó0   •— t         ‰| �  «        || _        y)z6Initialize the DFL module with regularization maximum.N)r   r   Úreg_max©r   rR   r   s     €r   r   zDFLoss.__init__[   s   ø€ ä‰ÑÔØˆ�r    c                ó´  — |j                  d| j                  dz
  dz
  «      }|j                  «       }|dz   }||z
  }d|z
  }t        j                  ||j                  d«      d¬«      j                  |j                  «      |z  t        j                  ||j                  d«      d¬«      j                  |j                  «      |z  z   j                  dd¬«      S )	zZReturn sum of left and right DFL losses from https://ieeexplore.ieee.org/document/9792391.r   r   g{®Gáz„?éÿÿÿÿr#   r$   T©Úkeepdim)Úclamp_rR   Úlongr(   Úcross_entropyÚviewÚshaper+   )r   Ú	pred_distÚtargetÚtlÚtrÚwlÚwrs          r   Ú__call__zDFLoss.__call__`   s¹   € à—‘˜q $§,¡,°Ñ"2°TÑ"9Ó:ˆØ�[‰[‹]ˆØ�!‰VˆØ�&‰[ˆØ�‰Vˆä�O‰O˜I r§w¡w¨r£{¸fÔE×JÑJÈ2Ï8É8ÓTÐWYÑYÜ�o‰o˜i¨¯©°«ÀÔG×LÑLÈRÏXÉXÓVÐY[Ñ[ñ\ç
‰$ˆr˜4ˆ$Ó
 ð	!r    ©é   )rR   Úintr6   ÚNone)r]   r5   r^   r5   r6   r5   )r8   r9   r:   r;   r   rc   r<   r=   s   @r   rP   rP   X   s   ø„ ÙFö÷

!r    rP   c                  óV   ‡ — e Zd ZdZddˆ fd„Z	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dd„Zˆ xZS )ÚBboxLosszACriterion class for computing training losses for bounding boxes.c                ó\   •— t         ‰| �  «        |dkD  rt        |«      | _        yd| _        y)zLInitialize the BboxLoss module with regularization maximum and DFL settings.r   N)r   r   rP   Údfl_lossrS   s     €r   r   zBboxLoss.__init__p   s%   ø€ ä‰ÑÔØ+2°Qª;œ˜w›ˆ�¸Dˆ�r    c
                ó  — |j                  d«      |   j                  d«      }
t        ||   ||   dd¬«      }d|z
  |
z  j                  «       |z  }| j                  rzt	        ||| j                  j
                  dz
  «      }| j                  ||   j                  d| j                  j
                  «      ||   «      |
z  }|j                  «       |z  }||fS t	        ||«      }||	z  }|ddd	d
…fxx   |d   z  cc<   |ddd	d
…fxx   |d   z  cc<   ||	z  }|ddd	d
…fxx   |d   z  cc<   |ddd	d
…fxx   |d   z  cc<   t        j                  ||   ||   d¬«      j                  dd¬«      |
z  }|j                  «       |z  }||fS )z.Compute IoU and DFL losses for bounding boxes.rU   FT)ÚxywhÚCIoUrD   r   .r   Né   r#   r$   rV   )
r,   Ú	unsqueezer   rk   r   rR   r[   r(   Úl1_lossr+   ©r   r]   Úpred_bboxesÚanchor_pointsÚtarget_bboxesÚtarget_scoresÚtarget_scores_sumÚfg_maskÚimgszÚstrider0   ÚiouÚloss_iouÚtarget_ltrbÚloss_dfls                  r   r2   zBboxLoss.forwardu   sÉ  € ð ×"Ñ" 2Ó& wÑ/×9Ñ9¸"Ó=ˆÜ�{ 7Ñ+¨]¸7Ñ-CÈ%ÐVZÔ[ˆØ˜3‘Y &Ñ(×-Ñ-Ó/Ð2CÑCˆð �=Š=Ü# M°=À$Ç-Á-×BWÑBWÐZ[ÑB[Ó\ˆKØ—}‘} Y¨wÑ%7×%<Ñ%<¸RÀÇÁ×AVÑAVÓ%WÐYdÐelÑYmÓnÐqwÑwˆHØ—|‘|“~Ð(9Ñ9ˆHð ˜Ð!Ð!ô $ M°=ÓAˆKà%¨Ñ.ˆKØ˜˜Q˜T ˜T˜	Ó" e¨A¡hÑ.Ó"Ø˜˜Q˜T ˜T˜	Ó" e¨A¡hÑ.Ó"Ø! FÑ*ˆIØ�c˜1˜4˜a˜4�iÓ  E¨!¡HÑ,Ó Ø�c˜1˜4˜a˜4�iÓ  E¨!¡HÑ,Ó ä—	‘	˜) GÑ,¨k¸'Ñ.BÈfÔU×ZÑZÐ[]ÐgkÐZÓlÐouÑuð ð  —|‘|“~Ð(9Ñ9ˆHà˜Ð!Ð!r    rd   ©rR   rf   ©r]   r5   rs   r5   rt   r5   ru   r5   rv   r5   rw   r5   rx   r5   ry   r5   rz   r5   r6   ú!tuple[torch.Tensor, torch.Tensor]r7   r=   s   @r   ri   ri   m   ss   ø„ ÙKöAð
$"àð$"ð "ð$"ð $ð	$"ð
 $ð$"ð $ð$"ð (ð$"ð ð$"ð ð$"ð ð$"ð 
+÷$"r    ri   c                  óD   ‡ — e Zd ZdZddˆ fd„Z	 d	 	 	 	 	 	 	 	 	 dd„Zˆ xZS )ÚRLELossaÈ  Residual Log-Likelihood Estimation Loss.

    Attributes:
        size_average (bool): Option to average the loss by the batch_size.
        use_target_weight (bool): Option to use weighted loss.
        residual (bool): Option to add L1 loss and let the flow learn the residual error distribution.

    References:
        https://arxiv.org/abs/2107.11291
        https://github.com/open-mmlab/mmpose/blob/main/mmpose/models/losses/regression_loss.py
    c                óL   •— t         ‰| �  «        || _        || _        || _        y)aG  Initialize RLELoss with target weight and residual options.

        Args:
            use_target_weight (bool): Whether to use target weights for loss calculation.
            size_average (bool): Whether to average the loss over elements.
            residual (bool): Whether to include residual log-likelihood term.
        N)r   r   Úsize_averageÚuse_target_weightÚresidual)r   r†   r…   r‡   r   s       €r   r   zRLELoss.__init__©   s'   ø€ ô 	‰ÑÔØ(ˆÔØ!2ˆÔØ ˆ�r    c                óž  — t        j                  |«      }||j                  d«      z
  }| j                  r1|t        j                  |dz  «      t        j                  |«      z   z  }| j
                  r2|€J d«       ‚|j                  «       dk(  r|j                  d«      }||z  }| j                  r|t        |«      z  }|j                  «       S )a&  
        Args:
            sigma (torch.Tensor): Output sigma, shape (N, D).
            log_phi (torch.Tensor): Output log_phi, shape (N).
            error (torch.Tensor): Error, shape (N, D).
            target_weight (torch.Tensor): Weights across different joint types, shape (N).
        r   ro   zD'target_weight' should not be None when 'use_target_weight' is True.)
rA   Úlogrp   r‡   Úabsr†   Údimr…   Úlenr,   )r   ÚsigmaÚlog_phiÚerrorÚtarget_weightÚ	log_sigmar1   s          r   r2   zRLELoss.forward¶   s¼   € ô —I‘I˜eÓ$ˆ	Ø˜7×,Ñ,¨QÓ/Ñ/ˆà�=Š=Ø”E—I‘I˜e a™iÓ(¬5¯9©9°UÓ+;Ñ;Ñ;ˆDà×!Ò!Ø Ð,ÐtÐ.tÓtÐ,Ø× Ñ Ó" aÒ'Ø -× 7Ñ 7¸Ó :�Ø�MÑ!ˆDà×ÒØ”C˜“IÑˆDà�x‰x‹zÐr    )TTT)r†   Úboolr…   r’   r‡   r’   )N)
r�   r5   rŽ   r5   r�   r5   r�   r5   r6   r5   r7   r=   s   @r   rƒ   rƒ   œ   sA   ø„ ñ
ö!ð nrðØ!ðØ,8ðØAMðØ^jðà	÷r    rƒ   c                  óT   ‡ — e Zd ZdZdˆ fd„Z	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dd„Zˆ xZS )ÚRotatedBboxLosszICriterion class for computing training losses for rotated bounding boxes.c                ó$   •— t         ‰| �  |«       y)zSInitialize the RotatedBboxLoss module with regularization maximum and DFL settings.N)r   r   rS   s     €r   r   zRotatedBboxLoss.__init__Õ   s   ø€ ä‰Ñ˜Õ!r    c
                óN  — |j                  d«      |   j                  d«      }
t        ||   ||   «      }d|z
  |
z  j                  «       |z  }| j                  rŠt	        |ddd…f   ||ddd…f   | j                  j
                  dz
  ¬«      }| j                  ||   j                  d| j                  j
                  «      ||   «      |
z  }|j                  «       |z  }||fS t	        |ddd…f   ||ddd…f   «      }||	z  }|dd	dd
…fxx   |d   z  cc<   |dddd
…fxx   |d	   z  cc<   ||	z  }|dd	dd
…fxx   |d   z  cc<   |dddd
…fxx   |d	   z  cc<   t        j                  ||   ||   d¬«      j                  dd¬«      |
z  }|j                  «       |z  }||fS )z6Compute IoU and DFL losses for rotated bounding boxes.rU   rD   .Né   é   r   )rR   r   ro   r#   r$   TrV   )
r,   rp   r   rk   r   rR   r[   r(   rq   r+   rr   s                  r   r2   zRotatedBboxLoss.forwardÙ   s	  € ð ×"Ñ" 2Ó& wÑ/×9Ñ9¸"Ó=ˆÜ�k 'Ñ*¨M¸'Ñ,BÓCˆØ˜3‘Y &Ñ(×-Ñ-Ó/Ð2CÑCˆð �=Š=Ü#Ø˜c 2 A 2˜gÑ&¨°}ÀSÈ!ÈAÈ#ÀXÑ7NÐX\×XeÑXe×XmÑXmÐpqÑXqôˆKð —}‘} Y¨wÑ%7×%<Ñ%<¸RÀÇÁ×AVÑAVÓ%WÐYdÐelÑYmÓnÐqwÑwˆHØ—|‘|“~Ð(9Ñ9ˆHð ˜Ð!Ð!ô $ M°#°r¸°r°'Ñ$:¸MÈ=ÐY\Ð^_Ð`aÐ^aÐYaÑKbÓcˆKØ%¨Ñ.ˆKØ˜˜Q˜T ˜T˜	Ó" e¨A¡hÑ.Ó"Ø˜˜Q˜T ˜T˜	Ó" e¨A¡hÑ.Ó"Ø! FÑ*ˆIØ�c˜1˜4˜a˜4�iÓ  E¨!¡HÑ,Ó Ø�c˜1˜4˜a˜4�iÓ  E¨!¡HÑ,Ó ä—	‘	˜) GÑ,¨k¸'Ñ.BÈfÔU×ZÑZÐ[]ÐgkÐZÓlÐouÑuð ð  —|‘|“~Ð(9Ñ9ˆHà˜Ð!Ð!r    r   r€   r7   r=   s   @r   r”   r”   Ò   sr   ø„ ÙSõ"ð%"àð%"ð "ð%"ð $ð	%"ð
 $ð%"ð $ð%"ð (ð%"ð ð%"ð ð%"ð ð%"ð 
+÷%"r    r”   c                  ó.   ‡ — e Zd ZdZddˆ fd„Zdd„Zˆ xZS )ÚMultiChannelDiceLossz8Criterion class for computing multi-channel Dice losses.c                ó>   •— t         ‰| �  «        || _        || _        y)zïInitialize MultiChannelDiceLoss with smoothing and reduction options.

        Args:
            smooth (float): Smoothing factor to avoid division by zero.
            reduction (str): Reduction method ('mean', 'sum', or 'none').
        N)r   r   Úsmoothr%   )r   rœ   r%   r   s      €r   r   zMultiChannelDiceLoss.__init__  s   ø€ ô 	‰ÑÔØˆŒØ"ˆ�r    c                óØ  — |j                  «       |j                  «       k(  sJ d«       ‚|j                  «       }||z  j                  d¬«      }|j                  d¬«      |j                  d¬«      z   }d|z  | j                  z   || j                  z   z  }d|z
  }|j	                  d¬«      }| j
                  dk(  r|j	                  «       S | j
                  dk(  r|j                  «       S |S )	zBCalculate multi-channel Dice loss between predictions and targets.z-the size of predict and target must be equal.)ro   é   ©r‹   r3   rD   r   r+   r,   )Úsizer&   r,   rœ   r+   r%   )r   rJ   r^   ÚintersectionÚunionÚdiceÚ	dice_losss          r   r2   zMultiChannelDiceLoss.forward  sÕ   € à�y‰y‹{˜fŸk™k›mÒ+Ð\Ð-\Ó\Ð+à�|‰|‹~ˆØ˜v™×*Ñ*¨vÐ*Ó6ˆØ—‘˜V�Ó$ v§z¡z°f zÓ'=Ñ=ˆØ�lÑ" T§[¡[Ñ0°U¸T¿[¹[Ñ5HÑIˆØ˜$‘Jˆ	Ø—N‘N q�NÓ)ˆ	à�>‰>˜VÒ#Ø—>‘>Ó#Ð#Ø�^‰^˜uÒ$Ø—=‘=“?Ð"àÐr    )g�íµ ÷Æ°>r+   )rœ   r*   r%   Ústr©rJ   r5   r^   r5   r6   r5   r7   r=   s   @r   rš   rš     s   ø„ ÙBö	#÷r    rš   c                  ó.   ‡ — e Zd ZdZddˆ fd„Zdd„Zˆ xZS )ÚBCEDiceLossz;Criterion class for computing combined BCE and Dice losses.c                ó’   •— t         ‰| �  «        || _        || _        t	        j
                  «       | _        t        d¬«      | _        y)zÞInitialize BCEDiceLoss with BCE and Dice weight factors.

        Args:
            weight_bce (float): Weight factor for BCE loss component.
            weight_dice (float): Weight factor for Dice loss component.
        r   )rœ   N)	r   r   Ú
weight_bceÚweight_diceÚnnÚBCEWithLogitsLossÚbcerš   r£   )r   rª   r«   r   s      €r   r   zBCEDiceLoss.__init__%  s;   ø€ ô 	‰ÑÔØ$ˆŒØ&ˆÔÜ×'Ñ'Ó)ˆŒÜ(°Ô2ˆ�	r    c                ó  — |j                   \  }}}}t        |j                   dd «      ||fk7  rt        j                  |||fd¬«      }| j                  | j                  ||«      z  | j                  | j                  ||«      z  z   S )zECalculate combined BCE and Dice loss between predictions and targets.éþÿÿÿNÚnearest)Úmode)r\   Útupler(   Úinterpolaterª   r®   r«   r£   )r   rJ   r^   Ú_Úmask_hÚmask_ws         r   r2   zBCEDiceLoss.forward2  s   € à#Ÿz™zÑˆˆ1ˆf�fÜ�—‘˜b˜cÐ"Ó#¨°Ð'7Ò7Ü—]‘] 6¨F°FÐ+;À)ÔLˆFØ�‰ §¡¨$°Ó!7Ñ7¸$×:JÑ:JÈTÏYÉYÐW[Ð]cÓMdÑ:dÑdÐdr    )ç      à?r¸   )rª   r*   r«   r*   r¦   r7   r=   s   @r   r¨   r¨   "  s   ø„ ÙEö3÷er    r¨   c                  ó@   ‡ — e Zd ZdZdˆ fd„Z	 	 	 	 	 	 	 	 	 	 dd„Zˆ xZS )ÚKeypointLossz.Criterion class for computing keypoint losses.c                ó0   •— t         ‰| �  «        || _        y)z7Initialize the KeypointLoss class with keypoint sigmas.N)r   r   Úsigmas)r   r¼   r   s     €r   r   zKeypointLoss.__init__=  s   ø€ ä‰ÑÔØˆ�r    c                ó”  — |d   |d   z
  j                  d«      |d   |d   z
  j                  d«      z   }|j                  d   t        j                  |dk7  d¬«      dz   z  }|d| j                  z  j                  d«      |dz   z  dz  z  }|j                  dd«      dt        j                  | «      z
  |z  z  j                  «       S )	zICalculate keypoint loss factor and Euclidean distance loss for keypoints.©.r   ro   ©.r   r   r   rŸ   ç•Ö&è.>rU   )r'   r\   rA   r,   r¼   r[   Úexpr+   )r   Ú	pred_kptsÚgt_kptsÚkpt_maskÚareaÚdÚkpt_loss_factorÚes           r   r2   zKeypointLoss.forwardB  sÌ   € ð �vÑ ¨¡Ñ0×5Ñ5°aÓ8¸IÀfÑ<MÐPWÐX^ÑP_Ñ<_×;dÑ;dÐefÓ;gÑgˆØ"Ÿ.™.¨Ñ+¬u¯y©y¸ÀQ¹ÈAÔ/NÐQUÑ/UÑVˆà�!�d—k‘k‘/×&Ñ& qÓ)¨T°D©[Ñ9¸AÑ=Ñ>ˆØ×$Ñ$ R¨Ó+°´E·I±I¸q¸b³MÑ0AÀXÑ/MÑN×TÑTÓVÐVr    )r¼   r5   r6   rg   )
rÂ   r5   rÃ   r5   rÄ   r5   rÅ   r5   r6   r5   r7   r=   s   @r   rº   rº   :  s>   ø„ Ù8õð
WØ%ðWØ0<ðWØHTðWØ\hðWà	÷Wr    rº   c                  ó^   — e Zd ZdZd
dd„Zdd„Zdd„Zdd„Z	 	 	 	 dd„Z	 	 	 	 	 	 dd„Z	dd	„Z
y)Úv8DetectionLosszJCriterion class for computing training losses for YOLOv8 object detection.Nc                ól  — t        |j                  «       «      j                  }|j                  }|j                  d   }t        j                  d¬«      | _        || _        |j                  | _	        |j                  | _
        |j                  |j                  dz  z   | _        |j                  | _        || _        |j                  dkD  | _        t        |dd«      | _        | j                  �1| j                  j!                  |«      j#                  ddd«      | _        t%        || j                  dd	| j                  j'                  «       |¬
«      | _        t+        |j                  «      j!                  |«      | _        t/        j0                  |j                  t.        j2                  |¬«      | _        y)zVInitialize v8DetectionLoss with model parameters and task-aligned assignment settings.rU   r#   r$   r—   r   Úclass_weightsNr¸   ç      @©ÚtopkÚnum_classesr   Úbetarz   Útopk2©rG   rF   )ÚnextÚ
parametersrF   ÚargsÚmodelr¬   r­   r®   Úhyprz   ÚncrR   ÚnoÚuse_dflÚgetattrrÌ   rI   r[   r   ÚtolistÚassignerri   Ú	bbox_lossrA   Úaranger*   Úproj)r   r×   Útal_topkÚ	tal_topk2rF   ÚhÚms          r   r   zv8DetectionLoss.__init__P  sH  € ä�e×&Ñ&Ó(Ó)×0Ñ0ˆØ�J‰Jˆà�K‰K˜‰OˆÜ×'Ñ'°&Ô9ˆŒØˆŒØ—h‘hˆŒØ—$‘$ˆŒØ—$‘$˜Ÿ™ Q™Ñ&ˆŒØ—y‘yˆŒØˆŒà—y‘y 1‘}ˆŒô % U¨O¸TÓBˆÔØ×ÑÐ)Ø!%×!3Ñ!3×!6Ñ!6°vÓ!>×!CÑ!CÀAÀqÈ"Ó!MˆDÔä+ØØŸ™ØØØ—;‘;×%Ñ%Ó'Øô
ˆŒô " !§)¡)Ó,×/Ñ/°Ó7ˆŒÜ—L‘L §¡´%·+±+ÀfÔMˆ�	r    c                ó  — |j                   \  }}|dk(  r(t        j                  |d|dz
  | j                  ¬«      }|S |dd…df   j	                  «       }|j                  d¬«      \  }}	|	j                  t        j                  ¬«      }	t        j                  ||	j                  «       |dz
  | j                  ¬«      }t        j                  |dz   t        j                  | j                  ¬«      }
|
j                  d|dz   t        j                  |«      «       |
j                  d«      }
t        j                  || j                  ¬«      |
|   z
  }|dd…dd…f   |||f<   t        |d	dd
…f   j                  |«      «      |d	dd
…f<   |S )zJPreprocess targets by converting to tensor format and scaling coordinates.r   r   ©rF   NT©Úreturn_counts©rG   rÓ   .r˜   )r\   rA   ÚzerosrF   rY   ÚuniquerI   Úint32ÚmaxÚscatter_add_Ú	ones_likeÚcumsumrà   r   Úmul_)r   ÚtargetsÚ
batch_sizeÚscale_tensorÚnlÚneÚoutÚ	batch_idxrµ   ÚcountsÚoffsetsÚ
within_idxs               r   Ú
preprocesszv8DetectionLoss.preprocessp  sR  € à—‘‰ˆˆBØ�Š7Ü—+‘+˜j¨!¨R°!©V¸D¿K¹KÔHˆCð ˆ
ð  ¢ 1 ™×*Ñ*Ó,ˆIØ!×(Ñ(°tÐ(Ó<‰IˆAˆvØ—Y‘Y¤U§[¡[�YÓ1ˆFÜ—+‘+˜j¨&¯*©*«,¸¸Q¹ÀtÇ{Á{ÔSˆCÜ—k‘k *¨q¡.¼¿
¹
È4Ï;É;ÔWˆGØ× Ñ   I°¡M´5·?±?À9Ó3MÔNØ—n‘n QÓ'ˆGÜŸ™ b°·±Ô=ÀÈ	Ñ@RÑRˆJØ)0²°A±B°©ˆC�	˜:Ð%Ñ&Ü% c¨#¨q°¨s¨(¡m×&8Ñ&8¸Ó&FÓGˆC��Q�q�S�‰MØˆ
r    c                ó  — | j                   rh|j                  \  }}}|j                  ||d|dz  «      j                  d«      j	                  | j
                  j                  |j                  «      «      }t        ||d¬«      S )zUDecode predicted object bounding box coordinates from anchor points and distribution.r—   rž   F)rm   )	rÛ   r\   r[   ÚsoftmaxÚmatmulrá   ÚtyperG   r   )r   rt   r]   ÚbÚaÚcs         r   Úbbox_decodezv8DetectionLoss.bbox_decode‚  sk   € à�<Š<Ø—o‘o‰GˆAˆq�!Ø!Ÿ™ q¨!¨Q°°Q±Ó7×?Ñ?ÀÓB×IÑIÈ$Ï)É)Ï.É.ÐYb×YhÑYhÓJiÓjˆIô ˜ M¸Ô>Ð>r    c                óJ  — t        j                  d| j                  ¬«      }|d   j                  ddd«      j	                  «       |d   j                  ddd«      j	                  «       }}t        |d   | j                  d	«      \  }}|j                  }|j                  d   }	t        j                  |d   d   j                  dd
 | j                  |¬«      | j                  d   z  }
t        j                  |d   j                  dd«      |d   j                  dd«      |d   fd«      }| j                  |j                  | j                  «      |	|
g d¢   ¬«      }|j                  dd«      \  }}|j                  dd¬«      j!                  d«      }| j#                  ||«      }| j%                  |j'                  «       j)                  «       |j'                  «       |z  j+                  |j                  «      ||z  |||«      \  }}}}}t-        |j                  «       d«      }| j/                  ||j                  |«      «      }| j0                  �|| j0                  z  }|j                  «       |z  |d<   |j                  «       r%| j3                  |||||z  ||||
|«	      \  |d<   |d<   |dxx   | j4                  j6                  z  cc<   |dxx   | j4                  j8                  z  cc<   |dxx   | j4                  j:                  z  cc<   |||||f||j'                  «       fS )z‹Calculate the sum of the loss for box, cls and dfl multiplied by batch size and return foreground mask and
        target indices.
        rž   rç   Úboxesr   ro   r   ÚscoresÚfeatsr¸   NrE   rù   rU   ÚclsÚbboxes©r   r   r   r   ©rõ   )r   r—   TrV   ç        )rA   rë   rF   ÚpermuteÚ
contiguousr   rz   rG   r\   rB   Úcatr[   rý   rI   Úsplitr,   Úgt_r  rÞ   Údetachr&   r  rî   r®   rÌ   rß   rØ   Úboxr
  Údfl)r   ÚpredsÚbatchr1   Úpred_distriÚpred_scoresrt   Ústride_tensorrG   rô   ry   ró   Ú	gt_labelsÚ	gt_bboxesÚmask_gtrs   rµ   ru   rv   rx   Útarget_gt_idxrw   Úbce_losss                          r   Úget_assigned_targets_and_lossz-v8DetectionLoss.get_assigned_targets_and_loss‹  sü  € ô �{‰{˜1 T§[¡[Ô1ˆà�'‰N×"Ñ" 1 a¨Ó+×6Ñ6Ó8Ø�(‰O×#Ñ# A q¨!Ó,×7Ñ7Ó9ð !ˆô (4°E¸'±NÀDÇKÁKÐQTÓ'UÑ$ˆ�}à×!Ñ!ˆØ ×&Ñ& qÑ)ˆ
Ü—‘˜U 7™^¨AÑ.×4Ñ4°Q°RÐ8ÀÇÁÐTYÔZÐ]a×]hÑ]hÐijÑ]kÑkˆô —)‘)˜U ;Ñ/×4Ñ4°R¸Ó;¸UÀ5¹\×=NÑ=NÈrÐSTÓ=UÐW\Ð]eÑWfÐgÐijÓkˆØ—/‘/ '§*¡*¨T¯[©[Ó"9¸:ÐTYÒZfÑTg�/ÓhˆØ&Ÿ}™}¨V°QÓ7Ñˆ	�9Ø—-‘- ¨4�-Ó0×4Ñ4°SÓ9ˆð ×&Ñ& }°kÓBˆàBFÇ-Á-Ø×ÑÓ ×(Ñ(Ó*Ø×ÑÓ! MÑ1×7Ñ7¸	¿¹ÓHØ˜MÑ)ØØØóC
Ñ?ˆˆ=˜-¨°-ô   × 1Ñ 1Ó 3°QÓ7Ðð —8‘8˜K¨×)9Ñ)9¸%Ó)@ÓAˆØ×ÑÐ)Ø˜×*Ñ*Ñ*ˆHØ—,‘,“.Ð#4Ñ4ˆˆQ‰ð �;‰;Œ=Ø#Ÿ~™~ØØØØ Ñ-ØØ!ØØØó
 ÑˆD�‰G�T˜!‘Wð 	ˆQ‹�4—8‘8—<‘<Ñ‹ØˆQ‹�4—8‘8—<‘<Ñ‹ØˆQ‹�4—8‘8—<‘<Ñ‹à�m ]°MÀ=ÐQØØ�K‰K‹Mð
ð 	
r    c                ó0   — t        |t        «      r|d   S |S )ú,Parse model predictions to extract features.r   )Ú
isinstancer³   ©r   r  s     r   Úparse_outputzv8DetectionLoss.parse_outputË  s   € ô & e¬UÔ3ˆu�Q‰xÐ>¸Ð>r    c                óD   — | j                  | j                  |«      |«      S )úLCalculate the sum of the loss for box, cls and dfl multiplied by batch size.©r1   r&  ©r   r  r  s      r   rc   zv8DetectionLoss.__call__Ñ  s    € ð �y‰y˜×*Ñ*¨5Ó1°5Ó9Ð9r    c                ód   — |d   j                   d   }| j                  ||«      dd \  }}||z  |fS )z0Calculate detection loss using assigned targets.r  r   r   N)r\   r!  )r   r  r  rô   r1   Úloss_detachs         r   r1   zv8DetectionLoss.lossÙ  sD   € à˜7‘^×)Ñ)¨!Ñ,ˆ
Ø ×>Ñ>¸uÀeÓLÈQÈRÐPÑˆˆkØ�jÑ  +Ð-Ð-r    ©é
   N©râ   rf   rã   ú
int | None©ró   r5   rô   rf   rõ   r5   r6   r5   )rt   r5   r]   r5   r6   r5   )r  údict[str, torch.Tensor]r  zdict[str, Any]r6   r³   )r  úFdict[str, torch.Tensor] | tuple[torch.Tensor, dict[str, torch.Tensor]]r6   r5   )r  r3  r  r2  r6   r�   ©r  r2  r  r2  r6   r�   )r8   r9   r:   r;   r   rý   r  r!  r&  rc   r1   © r    r   rÊ   rÊ   M  sW   „ ÙTôNó@ó$?ó>
ð@?Ø[ð?à	ó?ð:àUð:ð 'ð:ð 
+ó	:ô.r    rÊ   c                  ó„   ‡ — e Zd ZdZddˆ fd„Zdd„Ze	 	 	 	 	 	 	 	 	 	 	 	 d	d„«       Z	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d
d„Zˆ xZ	S )Úv8SegmentationLosszFCriterion class for computing training losses for YOLOv8 segmentation.c                ó‚   •— t         ‰| �  |||«       |j                  j                  | _        t        dd¬«      | _        y)zWInitialize the v8SegmentationLoss class with model parameters and mask overlap setting.r¸   )rª   r«   N)r   r   rÖ   Úoverlap_maskÚoverlapr¨   Úbcedice_loss©r   r×   râ   rã   r   s       €r   r   zv8SegmentationLoss.__init__ã  s4   ø€ ä‰Ñ˜ ¨)Ô4Ø—z‘z×.Ñ.ˆŒÜ'°3ÀCÔHˆÕr    c           
     óà  — |d   j                  ddd«      j                  «       |d   }}t        j                  d| j                  ¬«      }t        |t        «      rt        |«      dk(  r|\  }}nd}| j                  ||«      \  \  }}}	}
}
}}
|d   |d   |d   c|d<   |d<   |d	<   |j                  \  }}
}}|j                  «       �r |d
   j                  | j                  «      j                  «       }t        |j                  dd «      ||fk7  r&t        j                  ||j                  dd dd¬«      }t        j                  |d   d   j                  dd | j                  |j                   ¬«      | j"                  d   z  }| j%                  ||||	|d   j'                  dd«      |||«      |d<   |��ƒ|d   j                  | j                  «      }t        j(                  |j+                  «       | j,                  ¬«      j                  dd	dd«      j                  «       }| j.                  r)|dk(  }d||j1                  d«      j3                  |«      <   nX|d   j'                  d«      }t5        |«      D ]6  }|||k(     }t        |«      dk(  rŒd||dd…|j                  d¬«      dk(  f<   Œ8 | j7                  ||«      |d<   |dxx   | j8                  j:                  z  cc<   nR|dxx   |dz  j                  «       |dz  j                  «       z   z  cc<   |�|dxx   |dz  j                  «       z  cc<   |dxx   | j8                  j:                  z  cc<   ||z  |j=                  «       fS )zFCalculate and return the combined loss for detection and segmentation.Úmask_coefficientr   ro   r   Úprotor˜   rç   Nrž   Úmasksr°   ÚbilinearF)r²   Úalign_cornersr	  rE   rù   rU   Ú	sem_masks)rÐ   rŸ   r—   )r  r  rA   rë   rF   r$  r³   rŒ   r!  r\   r,   rI   r*   r(   r´   rB   rG   rz   Úcalculate_segmentation_lossr[   Úone_hotrY   rÙ   r:  rp   Ú	expand_asÚranger;  rØ   r  r  )r   r  r  Ú
pred_masksr?  r1   Úpred_semsegrx   r  ru   rµ   Údet_lossrô   r¶   r·   r@  ry   rC  Ú	mask_zerorù   ÚiÚinstance_mask_is                         r   r1   zv8SegmentationLoss.lossé  sW  € à!Ð"4Ñ5×=Ñ=¸aÀÀAÓF×QÑQÓSÐUZÐ[bÑUc�Eˆ
Ü�{‰{˜1 T§[¡[Ô1ˆÜ�eœUÔ#¬¨E«
°aªØ!&ÑˆE‘;àˆKØEI×EgÑEgÐhmÐotÓEuÑBÑ5ˆ�- °°1°xÀà$,¨Q¡K°¸!±¸hÀq¹kÐ!ˆˆQ‰��a‘˜$˜q™'à(-¯©Ñ%ˆ
�A�v˜vØ�;‰;�=à˜'‘N×%Ñ% d§k¡kÓ2×8Ñ8Ó:ˆEÜ�U—[‘[  Ð%Ó&¨6°6Ð*:Ò:äŸ™ e¨U¯[©[¸¸Ð-=ÀJÐ^cÔd�ô —‘˜U 7™^¨AÑ.×4Ñ4°Q°RÐ8ÀÇÁÐT^×TdÑTdÔeÐhl×hsÑhsÐtuÑhvÑvð ð ×6Ñ6ØØØØØ�kÑ"×'Ñ'¨¨AÓ.ØØØó	ˆD�‰Gð Ñ&Ø! +Ñ.×1Ñ1°$·+±+Ó>�	ÜŸI™I i§n¡nÓ&6ÀDÇGÁGÔL×TÑTÐUVÐXYÐ[\Ð^_Ó`×fÑfÓh�	à—<’<Ø %¨¡
�IØMN�I˜i×1Ñ1°!Ó4×>Ñ>¸yÓIÒJà % kÑ 2× 7Ñ 7¸Ó ;�IÜ" :Ó.ò M˜Ø*/°	¸Q±Ñ*?˜Ü˜Ó/°1Ò4Ø$ØKL˜	 !¢Q¨×(;Ñ(;ÀÐ(;Ó(BÀaÑ(GÐ"GÒHð	Mð ×+Ñ+¨K¸ÓC��Q‘Ø�Q“˜4Ÿ8™8Ÿ<™<Ñ'”ð �‹G˜ ™	—‘Ó(¨J¸©N×+?Ñ+?Ó+AÑAÑA‹GØÐ&Ø�Q“˜K¨!™O×0Ñ0Ó2Ñ2“àˆQ‹�4—8‘8—<‘<Ñ‹Ø�jÑ  $§+¡+£-Ð/Ð/r    c                óº   — t        j                  d||«      }t        j                  || d¬«      }t	        ||«      j                  d¬«      |z  j                  «       S )aO  Compute the instance segmentation loss for a single image.

        Args:
            gt_mask (torch.Tensor): Ground truth mask of shape (N, H, W), where N is the number of objects.
            pred (torch.Tensor): Predicted mask coefficients of shape (N, 32).
            proto (torch.Tensor): Prototype masks of shape (32, H, W).
            xyxy (torch.Tensor): Ground truth bounding boxes in xyxy format, normalized to [0, 1], of shape (N, 4).
            area (torch.Tensor): Area of each ground truth bounding box of shape (N,).

        Returns:
            (torch.Tensor): The calculated mask loss for a single image.

        Notes:
            The function uses the equation pred_mask = torch.einsum('in,nhw->ihw', pred, proto) to produce the
            predicted masks from the prototype masks and predicted mask coefficients.
        zin,nhw->ihwr#   r$   )r   ro   rŸ   )rA   Úeinsumr(   r)   r   r+   r,   )Úgt_maskrJ   r?  ÚxyxyrÅ   Ú	pred_maskr1   s          r   Úsingle_mask_lossz#v8SegmentationLoss.single_mask_loss%  sT   € ô( —L‘L °°eÓ<ˆ	Ü×1Ñ1°)¸WÐPVÔWˆÜ˜$ Ó%×*Ñ*¨vÐ*Ó6¸Ñ=×BÑBÓDÐDr    c	                ó®  — |j                   \  }	}	}
}d}||g d¢   z  }t        |«      ddd…f   j                  d«      }|t        j                  ||
||
g|j
                  ¬«      z  }t        t        |||||||«      «      D ]À  \  }}|\  }}}}}}}|j                  «       rw||   }| j                  r*||dz   j                  ddd«      k(  }|j                  «       }n||j                  d«      |k(     |   }|| j                  |||   |||   ||   «      z  }Œ—||dz  j                  «       |dz  j                  «       z   z  }ŒÂ ||j                  «       z  S )	aô  Calculate the loss for instance segmentation.

        Args:
            fg_mask (torch.Tensor): A binary tensor of shape (BS, N_anchors) indicating which anchors are positive.
            masks (torch.Tensor): Ground truth masks of shape (BS, H, W) if `overlap` is False, otherwise (BS, ?, H, W).
            target_gt_idx (torch.Tensor): Indexes of ground truth objects for each anchor of shape (BS, N_anchors).
            target_bboxes (torch.Tensor): Ground truth bounding boxes for each anchor of shape (BS, N_anchors, 4).
            batch_idx (torch.Tensor): Batch indices of shape (N_labels_in_batch, 1).
            proto (torch.Tensor): Prototype masks of shape (BS, 32, H, W).
            pred_masks (torch.Tensor): Predicted masks for each anchor of shape (BS, N_anchors, 32).
            imgsz (torch.Tensor): Size of the input image as a tensor of shape (2), i.e., (H, W).

        Returns:
            (torch.Tensor): The calculated loss for instance segmentation.

        Notes:
            The batch loss can be computed for improved speed at higher memory usage.
            For example, pred_mask can be computed as follows:
                pred_mask = torch.einsum('in,nhw->ihw', pred, proto)  # (i, 32) @ (32, 160, 160) -> (i, 160, 160)
        r   r  .ro   Nrç   r   rU   )r\   r	   ÚprodrA   rB   rF   Ú	enumerateÚziprH   r:  r[   r*   rS  r,   )r   rx   r@  r  ru   rù   r?  rH  ry   rµ   r¶   r·   r1   Útarget_bboxes_normalizedÚmareaÚmxyxyrL  Úsingle_iÚ	fg_mask_iÚtarget_gt_idx_iÚpred_masks_iÚproto_iÚmxyxy_iÚmarea_iÚmasks_iÚmask_idxrP  s                              r   rD  z.v8SegmentationLoss.calculate_segmentation_loss=  s‰  € ð>  %Ÿ{™{Ñˆˆ1ˆf�fØˆð $1°5ºÑ3FÑ#FÐ ô Ð2Ó3°C¸¹°GÑ<×AÑAÀ!ÓDˆð )¬5¯<©<¸ÀÈÐQWÐ8XÐaf×amÑamÔ+nÑnˆä$¤S¨°-ÀÈUÐTYÐ[`ÐbgÓ%hÓiò 	C‰KˆAˆxØ[cÑXˆI�¨°g¸wÈÐQXØ�}‰}ŒØ*¨9Ñ5�Ø—<’<Ø%¨(°Q©,×)<Ñ)<¸RÀÀAÓ)FÑF�GØ%Ÿm™m›o‘Gà# I§N¡N°2Ó$6¸!Ñ$;Ñ<¸XÑF�Gà˜×-Ñ-Ø˜\¨)Ñ4°g¸wÀyÑ?QÐSZÐ[dÑSeóñ ‘ð ˜ ™Ÿ™Ó)¨Z¸!©^×,@Ñ,@Ó,BÑBÑB‘ð!	Cð$ �g—k‘k“mÑ#Ð#r    r-  r/  r4  )rP  r5   rJ   r5   r?  r5   rQ  r5   rÅ   r5   r6   r5   )rx   r5   r@  r5   r  r5   ru   r5   rù   r5   r?  r5   rH  r5   ry   r5   r6   r5   )
r8   r9   r:   r;   r   r1   ÚstaticmethodrS  rD  r<   r=   s   @r   r7  r7  à  s»   ø„ ÙPöIó:0ðx ðEØðEØ%1ðEØ:FðEØNZðEØbnðEà	òEó ðEð.=$àð=$ð ð=$ð $ð	=$ð
 $ð=$ð  ð=$ð ð=$ð !ð=$ð ð=$ð 
÷=$r    r7  c                  ó„   ‡ — e Zd ZdZddˆ fd„Zd	d„Zed
d„«       Z	 	 	 	 	 	 	 	 	 	 dd„Z	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dd„Z	ˆ xZ
S )Ú
v8PoseLosszICriterion class for computing training losses for YOLOv8 pose estimation.c                ó¨  •— t         ‰| �  |||«       |j                  d   j                  | _        t	        j
                  «       | _        | j                  ddgk(  }| j                  d   }|r2t        j                  t        «      j                  | j                  «      n#t        j                  || j                  ¬«      |z  }t        |¬«      | _        y)zQInitialize v8PoseLoss with model parameters and keypoint-specific loss functions.rU   é   rž   r   rç   )r¼   N)r   r   r×   Ú	kpt_shaper¬   r­   Úbce_poserA   Ú
from_numpyr   rI   rF   Úonesrº   Úkeypoint_loss)r   r×   râ   rã   Úis_poseÚnkptr¼   r   s          €r   r   zv8PoseLoss.__init__€  s¢   ø€ ä‰Ñ˜ ¨)Ô4ØŸ™ R™×2Ñ2ˆŒÜ×,Ñ,Ó.ˆŒØ—.‘. R¨ GÑ+ˆØ�~‰~˜aÑ ˆÙ@G”×!Ñ!¤)Ó,×/Ñ/°·±Ô<ÌUÏZÉZÐX\Ðei×epÑepÔMqÐtxÑMxˆÜ)°Ô8ˆÕr    c           	     óö  — |d   j                  ddd«      j                  «       }t        j                  d| j                  ¬«      }| j                  ||«      \  \  }}}}}	}
}|
d   |
d   |
d   c|d<   |d<   |d<   |j                  d   }t        j                  |d	   d   j                  dd
 | j                  |j                  ¬«      | j                  d   z  }| j                  | |j                  |dg| j                  ¢­Ž «      }|j                  «       r�|d   j                  | j                  «      j                  «       j!                  «       }|dxx   |d   z  cc<   |dxx   |d   z  cc<   | j#                  ||||d   j                  dd«      |	||«      \  |d<   |d<   |dxx   | j$                  j&                  z  cc<   |dxx   | j$                  j(                  z  cc<   ||z  |j+                  «       fS )ú;Calculate the total loss and detach it for pose estimation.Úkptsr   ro   r   r˜   rç   rž   r—   r	  NrE   rU   Ú	keypointsr¾   r¿   rù   )r  r  rA   rë   rF   r!  r\   rB   rG   rz   Úkpts_decoder[   ri  r,   rI   r*   ÚcloneÚcalculate_keypoints_lossrØ   ÚposeÚkobjr  )r   r  r  rÂ   r1   rx   r  ru   rt   r  rJ  rµ   rô   ry   rs  s                  r   r1   zv8PoseLoss.lossŠ  sï  € à˜&‘M×)Ñ)¨!¨Q°Ó2×=Ñ=Ó?ˆ	Ü�{‰{˜1 T§[¡[Ô1ˆà×.Ñ.¨u°eÓ<ñ 	[ÑMˆ�- °¸}ÈxÐYZð %-¨Q¡K°¸!±¸hÀq¹kÐ!ˆˆQ‰��a‘˜$˜q™'à—_‘_ QÑ'ˆ
Ü—‘˜U 7™^¨AÑ.×4Ñ4°Q°RÐ8ÀÇÁÐT]×TcÑTcÔdÐgk×grÑgrÐstÑguÑuˆð ×$Ñ$ ]°N°I·N±NÀ:ÈrÐ4cÐTX×TbÑTbÒ4cÓdˆ	ð �;‰;Œ=Ø˜kÑ*×-Ñ-¨d¯k©kÓ:×@Ñ@ÓB×HÑHÓJˆIØ�fÓ  q¡Ñ)ÓØ�fÓ  q¡Ñ)Óà#×<Ñ<ØØØØ�kÑ"×'Ñ'¨¨AÓ.ØØØó ÑˆD�‰G�T˜!‘Wð 	ˆQ‹�4—8‘8—=‘=Ñ ‹ØˆQ‹�4—8‘8—=‘=Ñ ‹à�jÑ  $§+¡+£-Ð/Ð/r    c                ó¨   — |j                  «       }|ddd…fxx   dz  cc<   |dxx   | dd…dgf   dz
  z  cc<   |dxx   | dd…d	gf   dz
  z  cc<   |S )
ú0Decode predicted keypoints to image coordinates..Nro   r3   r¾   r   r¸   r¿   r   ©ru  ©rt   rÂ   Úys      r   rt  zv8PoseLoss.kpts_decode¯  sg   € ð �O‰OÓˆØ	ˆ#ˆr�ˆrˆ'‹
�cÑ‹
Ø	ˆ&‹	�]¢1 q c 6Ñ*¨SÑ0Ñ0‹	Ø	ˆ&‹	�]¢1 q c 6Ñ*¨SÑ0Ñ0‹	Øˆr    c           
     ó.  — |j                  «       }t        |«      }t        j                  |d¬«      d   j	                  «       }t        j
                  |||j                  d   |j                  d   f|j                  ¬«      }|j                  «       }t        j
                  |dz   t        j                  |j                  ¬«      }	|	j                  d|dz   t        j                  |«      «       |	j                  d«      }	t        j                  t        |«      |j                  ¬«      |	|   z
  }
||||
f<   |j                  d«      j                  d«      }|j                  d|j                  dd|j                  d   |j                  d   «      «      }|S )	a§  Select target keypoints for each anchor based on batch index and target ground truth index.

        Args:
            keypoints (torch.Tensor): Ground truth keypoints, shape (N_kpts_in_batch, N_kpts_per_object, kpts_dim).
            batch_idx (torch.Tensor): Batch index tensor for keypoints, shape (N_kpts_in_batch, 1).
            target_gt_idx (torch.Tensor): Index tensor mapping anchors to ground truth objects, shape (BS, N_anchors).
            masks (torch.Tensor): Binary mask tensor indicating object presence, shape (BS, N_anchors).

        Returns:
            (torch.Tensor): Selected keypoints tensor, shape (BS, N_anchors, N_kpts_per_object, kpts_dim).
        Trè   r   ro   rç   rÓ   r   rU   )ÚflattenrŒ   rA   rì   rî   rë   r\   rF   rY   rï   rð   rñ   rà   rp   ÚgatherÚexpand)r   rs  rù   r  r@  rô   Úmax_kptsÚbatched_keypointsÚbatch_idx_longrû   rü   Útarget_gt_idx_expandedÚselected_keypointss                r   Ú_select_target_keypointsz#v8PoseLoss._select_target_keypoints¸  sg  € ð$ ×%Ñ%Ó'ˆ	Ü˜“Zˆ
ô —<‘< 	¸Ô>¸qÑA×EÑEÓGˆô "ŸK™KØ˜ 9§?¡?°1Ñ#5°y·±ÀqÑ7IÐJÐS\×ScÑScô
Ðð
 #Ÿ™Ó)ˆÜ—+‘+˜j¨1™n´E·J±JÀy×GWÑGWÔXˆØ×Ñ˜Q °Ñ 2´E·O±OÀNÓ4SÔTØ—.‘. Ó#ˆÜ—\‘\¤# i£.¸×9IÑ9IÔJÈWÐUcÑMdÑdˆ
Ø8AÐ˜.¨*Ð4Ñ5ð "/×!8Ñ!8¸Ó!<×!FÑ!FÀrÓ!JÐð /×5Ñ5ØÐ%×,Ñ,¨R°°Y·_±_ÀQÑ5GÈÏÉÐYZÑI[Ó\ó
Ðð "Ð!r    c           	     ó  — | j                  ||||«      }|ddd…fxx   |j                  dddd«      z  cc<   d}	d}
|j                  «       r³||z  }||   }t        ||   «      dd…dd…f   j	                  dd¬«      }||   }|j
                  d   d	k(  r|d
   dk7  nt        j                  |d   d«      }| j                  ||||«      }	|j
                  d   d	k(  r#| j                  |d
   |j                  «       «      }
|	|
fS )a  Calculate the keypoints loss for the model.

        This function calculates the keypoints loss and keypoints object loss for a given batch. The keypoints loss is
        based on the difference between the predicted keypoints and ground truth keypoints. The keypoints object loss is
        a binary classification loss that classifies whether a keypoint is present or not.

        Args:
            masks (torch.Tensor): Binary mask tensor indicating object presence, shape (BS, N_anchors).
            target_gt_idx (torch.Tensor): Index tensor mapping anchors to ground truth objects, shape (BS, N_anchors).
            keypoints (torch.Tensor): Ground truth keypoints, shape (N_kpts_in_batch, N_kpts_per_object, kpts_dim).
            batch_idx (torch.Tensor): Batch index tensor for keypoints, shape (N_kpts_in_batch, 1).
            stride_tensor (torch.Tensor): Stride tensor for anchors, shape (N_anchors, 1).
            target_bboxes (torch.Tensor): Ground truth boxes in (x1, y1, x2, y2) format, shape (BS, N_anchors, 4).
            pred_kpts (torch.Tensor): Predicted keypoints, shape (BS, N_anchors, N_kpts_per_object, kpts_dim).

        Returns:
            kpts_loss (torch.Tensor): The keypoints loss.
            kpts_obj_loss (torch.Tensor): The keypoints object loss.
        .Nro   r   rU   r   TrV   rž   ©.ro   r¾   )r‡  r[   rH   r	   rU  r\   rA   Ú	full_likerm  rj  r*   )r   r@  r  rs  rù   r  ru   rÂ   r†  Ú	kpts_lossÚkpts_obj_lossÚgt_kptrÅ   Úpred_kptrÄ   s                  r   rv  z#v8PoseLoss.calculate_keypoints_lossç  s$  € ð< "×:Ñ:¸9ÀiÐQ^Ð`eÓfÐð 	˜3   ˜7Ó# }×'9Ñ'9¸!¸RÀÀAÓ'FÑFÓ#àˆ	Øˆà�9‰9Œ;Ø˜]Ñ*ˆMØ'¨Ñ.ˆFÜ˜]¨5Ñ1Ó2²1°a±b°5Ñ9×>Ñ>¸qÈ$Ð>ÓOˆDØ  Ñ'ˆHØ.4¯l©l¸2Ñ.>À!Ò.C�v˜f‘~¨Ò*ÌÏÉÐY_Ð`fÑYgÐimÓInˆHØ×*Ñ*¨8°V¸XÀtÓLˆIà�~‰~˜bÑ! QÒ&Ø $§¡¨h°vÑ.>ÀÇÁÓ@PÓ Q�à˜-Ð'Ð'r    )r.  r.  )râ   rf   rã   rf   r4  ©rt   r5   rÂ   r5   r6   r5   )
rs  r5   rù   r5   r  r5   r@  r5   r6   r5   )r@  r5   r  r5   rs  r5   rù   r5   r  r5   ru   r5   rÂ   r5   r6   r�   )r8   r9   r:   r;   r   r1   rd  rt  r‡  rv  r<   r=   s   @r   rf  rf  }  s®   ø„ ÙSö9ó#0ðJ òó ðð-"àð-"ð  ð-"ð $ð	-"ð
 ð-"ð 
ó-"ð^1(àð1(ð $ð1(ð  ð	1(ð
  ð1(ð $ð1(ð $ð1(ð  ð1(ð 
+÷1(r    rf  c                  óp   ‡ — e Zd ZdZddˆ fd„Zd	d„Zed
d„«       Zdd„Z	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dd„Z	ˆ xZ
S )Ú
PoseLoss26z_Criterion class for computing training losses for YOLOv8 pose estimation with RLE loss support.c                ó
  •— t         ‰| �  |||«       | j                  ddgk(  }| j                  d   }d| _        t	        |j
                  d   d«      r|j
                  d   j                  nd| _        | j                  �…t        d¬«      j                  | j                  «      | _        |r2t        j                  t        «      j                  | j                  «      n t        j                  || j                  ¬	«      | _        yy)
zdInitialize PoseLoss26 with model parameters and keypoint-specific loss functions including RLE loss.rh  rž   r   NrU   Ú
flow_modelT)r†   rç   )r   r   ri  Úrle_lossÚhasattrr×   r“  rƒ   rI   rF   rA   rk  r   rl  Útarget_weights)r   r×   râ   rã   rn  ro  r   s         €r   r   zPoseLoss26.__init__  sË   ø€ ä‰Ñ˜ ¨)Ô4Ø—.‘. R¨ GÑ+ˆØ�~‰~˜aÑ ˆØˆŒÜ8?ÀÇÁÈBÁÐQ]Ô8^˜%Ÿ+™+ b™/×4Ò4ÐdhˆŒØ�?‰?Ð&Ü#°dÔ;×>Ñ>¸t¿{¹{ÓKˆDŒMá@G”× Ñ ¤Ó,×/Ñ/°·±Ô<ÌUÏZÉZÐX\Ðei×epÑepÔMqð Õð 'r    c           	     óž  — |d   j                  ddd«      j                  «       }t        j                  | j                  rdnd| j
                  ¬«      }| j                  ||«      \  \  }}}}}	}
}|
d   |
d   |
d   c|d<   |d<   |d	<   |j                  d   }t        j                  |d
   d   j                  dd | j
                  |j                  ¬«      | j                  d   z  } |j                  |dg| j                  ¢­Ž }| j                  rp|j                  dd«      �^|d   j                  ddd«      j                  «       }|j                  |d| j                  d   d«      }t        j                  ||gd¬«      }| j                  ||«      }|j!                  «       r¨|d   j#                  | j
                  «      j%                  «       j'                  «       }|dxx   |d   z  cc<   |dxx   |d   z  cc<   | j)                  ||||d   j                  dd«      |	||«      }|d   |d<   |d   |d<   | j                  �|d   |d<   |dxx   | j*                  j,                  z  cc<   |dxx   | j*                  j.                  z  cc<   | j                  �!|dxx   | j*                  j0                  z  cc<   ||z  |j3                  «       fS )rq  rr  r   ro   r   é   r˜   rç   rž   r—   r	  NrE   rU   Ú
kpts_sigmarŸ   rs  r¾   r¿   rù   )r  r  rA   rë   r”  rF   r!  r\   rB   rG   rz   r[   ri  Úgetr  rt  r,   rI   r*   ru  rv  rØ   rw  rx  Úrler  )r   r  r  rÂ   r1   rx   r  ru   rt   r  rJ  rµ   rô   ry   Ú
pred_sigmars  Úkeypoints_losss                    r   r1   zPoseLoss26.loss+  s®  € à˜&‘M×)Ñ)¨!¨Q°Ó2×=Ñ=Ó?ˆ	Ü�{‰{Ø—’‰A A¨d¯k©kô
ˆð ×.Ñ.¨u°eÓ<ñ 	[ÑMˆ�- °¸}ÈxÐYZð %-¨Q¡K°¸!±¸hÀq¹kÐ!ˆˆQ‰��a‘˜$˜q™'à—_‘_ QÑ'ˆ
Ü—‘˜U 7™^¨AÑ.×4Ñ4°Q°RÐ8ÀÇÁÐT]×TcÑTcÔdÐgk×grÑgrÐstÑguÑuˆà"�I—N‘N :¨rÐC°D·N±NÒCˆ	à�=Š=˜UŸY™Y |°TÓ:ÐFØ˜|Ñ,×4Ñ4°Q¸¸1Ó=×HÑHÓJˆJØ#Ÿ™¨°R¸¿¹ÈÑ9JÈAÓNˆJÜŸ	™	 9¨jÐ"9¸rÔBˆIà×$Ñ$ ]°IÓ>ˆ	ð �;‰;Œ=Ø˜kÑ*×-Ñ-¨d¯k©kÓ:×@Ñ@ÓB×HÑHÓJˆIØ�fÓ  q¡Ñ)ÓØ�fÓ  q¡Ñ)Óà!×:Ñ:ØØØØ�kÑ"×'Ñ'¨¨AÓ.ØØØóˆNð % QÑ'ˆD�‰GØ$ QÑ'ˆD�‰GØ�}‰}Ð(Ø(¨Ñ+��Q‘àˆQ‹�4—8‘8—=‘=Ñ ‹ØˆQ‹�4—8‘8—=‘=Ñ ‹Ø�=‰=Ð$Ø�‹G�t—x‘x—|‘|Ñ#‹Gà�jÑ  $§+¡+£-Ð/Ð/r    c                óz   — |j                  «       }|dxx   | dd…dgf   z  cc<   |dxx   | dd…dgf   z  cc<   |S )rz  r¾   Nr   r¿   r   r{  r|  s      r   rt  zPoseLoss26.kpts_decode^  sG   € ð �O‰OÓˆØ	ˆ&‹	�]¢1 q c 6Ñ*Ñ*‹	Ø	ˆ&‹	�]¢1 q c 6Ñ*Ñ*‹	Øˆr    c                ó‚  — ||   }||   }|dd…dd…f   }|dd…dd…f   }|dd…dd…f   }| j                   j                  d«      j                  |j                  d   d«      }	|	|   }	|j	                  «       }||z
  |dz   z  }
t        j                  |
«      t        j                  |
«      z  j                  d¬«       }|j                  «       s!t        j                  d	|j                  ¬
«      S |
|   }
|
j                  dd«      }
||   }|	|   }	| j                  j                  |
«      }| j                  |||
|	«      S )aÎ  Calculate the RLE (Residual Log-likelihood Estimation) loss for keypoints.

        Args:
            pred_kpt (torch.Tensor): Predicted kpts with sigma, shape (N, num_keypoints, kpts_dim) where kpts_dim >= 4.
            gt_kpt (torch.Tensor): Ground truth keypoints, shape (N, num_keypoints, kpts_dim).
            kpt_mask (torch.Tensor): Mask for valid keypoints, shape (N, num_keypoints).

        Returns:
            (torch.Tensor): The RLE loss.
        Nr   ro   r°   r   rÀ   rU   rŸ   r  rç   iœÿÿÿéd   )r–  rp   Úrepeatr\   r&   rA   ÚisnanÚisinfrH   rB   rF   Úclampr“  Úlog_probr”  )r   rŽ  r�  rÄ   Úpred_kpt_visibleÚgt_kpt_visibleÚpred_coordsrœ  Ú	gt_coordsr–  r�   Ú
valid_maskrŽ   s                r   Úcalculate_rle_losszPoseLoss26.calculate_rle_lossf  sL  € ð $ HÑ-ÐØ Ñ)ˆØ&¢q¨!¨A¨# vÑ.ˆØ%¢a¨© fÑ-ˆ
Ø"¢1 a¨ c 6Ñ*ˆ	à×,Ñ,×6Ñ6°qÓ9×@Ñ@ÀÇÁÐPQÑARÐTUÓVˆØ'¨Ñ1ˆà×'Ñ'Ó)ˆ
Ø˜yÑ(¨Z¸$Ñ->Ñ?ˆô —{‘{ 5Ó)¬E¯K©K¸Ó,>Ñ>×CÑCÈÐCÓKÐKˆ
Ø�~‰~ÔÜ—<‘< ¨H¯O©OÔ<Ð<à�jÑ!ˆØ—‘˜D #Ó&ˆØ 
Ñ+ˆ
Ø'¨
Ñ3ˆà—/‘/×*Ñ*¨5Ó1ˆà�}‰}˜Z¨°%¸ÓHÐHr    c           	     óÜ  — | j                  ||||«      }|ddd…fxx   |j                  dddd«      z  cc<   d}	d}
d}|j                  «       �r||z  }||   }t        ||   «      dd…dd…f   j	                  dd¬«      }||   }|j
                  d   d	k(  r|d
   dk7  nt        j                  |d   d«      }| j                  ||||«      }	| j                  �I|j
                  d   dk(  s|j
                  d   dk(  r%| j                  |||«      }|j                  d¬«      }|j
                  d   d	k(  s|j
                  d   dk(  r#| j                  |d
   |j                  «       «      }
|	|
|fS )aG  Calculate the keypoints loss for the model.

        This function calculates the keypoints loss and keypoints object loss for a given batch. The keypoints loss is
        based on the difference between the predicted keypoints and ground truth keypoints. The keypoints object loss is
        a binary classification loss that classifies whether a keypoint is present or not.

        Args:
            masks (torch.Tensor): Binary mask tensor indicating object presence, shape (BS, N_anchors).
            target_gt_idx (torch.Tensor): Index tensor mapping anchors to ground truth objects, shape (BS, N_anchors).
            keypoints (torch.Tensor): Ground truth keypoints, shape (N_kpts_in_batch, N_kpts_per_object, kpts_dim).
            batch_idx (torch.Tensor): Batch index tensor for keypoints, shape (N_kpts_in_batch, 1).
            stride_tensor (torch.Tensor): Stride tensor for anchors, shape (N_anchors, 1).
            target_bboxes (torch.Tensor): Ground truth boxes in (x1, y1, x2, y2) format, shape (BS, N_anchors, 4).
            pred_kpts (torch.Tensor): Predicted keypoints, shape (BS, N_anchors, N_kpts_per_object, kpts_dim).

        Returns:
            kpts_loss (torch.Tensor): The keypoints loss.
            kpts_obj_loss (torch.Tensor): The keypoints object loss.
            rle_loss (torch.Tensor): The RLE loss.
        .Nro   r   rU   r   TrV   rž   r‰  r¾   r—   r˜   )Úmin)r‡  r[   rH   r	   rU  r\   rA   rŠ  rm  r”  r«  r¤  rj  r*   )r   r@  r  rs  rù   r  ru   rÂ   r†  r‹  rŒ  r”  r�  rÅ   rŽ  rÄ   s                   r   rv  z#PoseLoss26.calculate_keypoints_loss‹  s‹  € ð> "×:Ñ:¸9ÀiÐQ^Ð`eÓfÐð 	˜3   ˜7Ó# }×'9Ñ'9¸!¸RÀÀAÓ'FÑFÓ#àˆ	ØˆØˆà�9‰9�;Ø˜]Ñ*ˆMØ'¨Ñ.ˆFÜ˜]¨5Ñ1Ó2²1°a±b°5Ñ9×>Ñ>¸qÈ$Ð>ÓOˆDØ  Ñ'ˆHØ.4¯l©l¸2Ñ.>À!Ò.C�v˜f‘~¨Ò*ÌÏÉÐY_Ð`fÑYgÐimÓInˆHØ×*Ñ*¨8°V¸XÀtÓLˆIà�}‰}Ð(¨h¯n©n¸RÑ.@ÀAÒ.EÈÏÉÐXZÑI[Ð_`ÒI`Ø×2Ñ2°8¸VÀXÓN�Ø#Ÿ>™>¨a˜>Ó0�Ø�~‰~˜bÑ! QÒ&¨(¯.©.¸Ñ*<ÀÒ*AØ $§¡¨h°vÑ.>ÀÇÁÓ@PÓ Q�à˜-¨Ð1Ð1r    r-  r/  r4  r�  )rŽ  r5   r�  r5   rÄ   r5   r6   r5   )r@  r5   r  r5   rs  r5   rù   r5   r  r5   ru   r5   rÂ   r5   r6   z/tuple[torch.Tensor, torch.Tensor, torch.Tensor])r8   r9   r:   r;   r   r1   rd  rt  r«  rv  r<   r=   s   @r   r‘  r‘    s   ø„ Ùiöó10ðf òó ðó#IðJ62àð62ð $ð62ð  ð	62ð
  ð62ð $ð62ð $ð62ð  ð62ð 
9÷62r    r‘  c                  ó   — e Zd ZdZdd„Zy)Úv8ClassificationLosszACriterion class for computing training losses for classification.c                ó–   — t        |t        t        f«      r|d   n|}t        j                  ||d   d¬«      }||j                  «       fS )zDCompute the classification loss between predictions and true labels.r   r
  r+   r$   )r$  Úlistr³   r(   rZ   r  )r   r  r  r1   s       r   rc   zv8ClassificationLoss.__call__Ç  sA   € ä& u¬t´U¨mÔ<��a’À%ˆÜ�‰˜u e¨E¡l¸fÔEˆØ�T—[‘[“]Ð"Ð"r    N©r  r   r  r2  r6   r�   )r8   r9   r:   r;   rc   r5  r    r   r¯  r¯  Ä  s
   „ ÙKô#r    r¯  c                  óV   ‡ — e Zd ZdZddˆ fd„Zd	d„Zd
d„Z	 	 	 	 	 	 	 	 dd„Zdd„Zˆ xZ	S )Ú	v8OBBLosszdCalculates losses for object detection, classification, and box distribution in rotated YOLO models.c                óþ   •— t         ‰| �  ||¬«       t        || j                  dd| j                  j                  «       |¬«      | _        t        | j                  «      j                  | j                  «      | _        y)z^Initialize v8OBBLoss with model, assigner, and rotated bbox loss; model must be de-paralleled.©râ   r¸   rÍ   rÎ   N)r   r   r
   rÙ   rz   rÝ   rÞ   r”   rR   rI   rF   rß   r<  s       €r   r   zv8OBBLoss.__init__Ñ  se   ø€ ä‰Ñ˜¨ÐÔ2Ü2ØØŸ™ØØØ—;‘;×%Ñ%Ó'Øô
ˆŒô )¨¯©Ó6×9Ñ9¸$¿+¹+ÓFˆ�r    c                ó  — |j                   d   dk(  r%t        j                  |dd| j                  ¬«      }|S |dd…df   j	                  «       }|j                  d¬«      \  }}|j                  t        j                  ¬«      }t        j                  ||j                  «       d| j                  ¬«      }|dd…dd…f   j                  «       }|dd…dd	…f   j                  |«       t        j                  |dz   t        j                  | j                  ¬
«      }	|	j                  d|dz   t        j                  |«      «       |	j                  d«      }	t        j                  t        |«      | j                  ¬«      |	|   z
  }
||||
f<   |S )z7Preprocess targets for oriented bounding box detection.r   r˜  rç   NTrè   rê   r   r˜   rÓ   )r\   rA   rë   rF   rY   rì   rI   rí   rî   ru  rò   rï   rð   rñ   rà   rŒ   )r   ró   rô   rõ   rø   rù   rµ   rú   Úpacked_targetsrû   rü   s              r   rý   zv8OBBLoss.preprocessÞ  sH  € à�=‰=˜Ñ˜qÒ Ü—+‘+˜j¨!¨Q°t·{±{ÔCˆCð ˆ
ð  ¢ 1 ™×*Ñ*Ó,ˆIØ!×(Ñ(°tÐ(Ó<‰IˆAˆvØ—Y‘Y¤U§[¡[�YÓ1ˆFÜ—+‘+˜j¨&¯*©*«,¸À$Ç+Á+ÔNˆCØ$¢Q¨© U™^×1Ñ1Ó3ˆNØš1˜a ˜c˜6Ñ"×'Ñ'¨Ô5Ü—k‘k *¨q¡.¼¿
¹
È4Ï;É;ÔWˆGØ× Ñ   I°¡M´5·?±?À9Ó3MÔNØ—n‘n QÓ'ˆGÜŸ™¤c¨'£l¸4¿;¹;ÔGÈ'ÐR[ÑJ\Ñ\ˆJØ)7ˆC�	˜:Ð%Ñ&Øˆ
r    c                ó`  — t        j                  d| j                  ¬«      }|d   j                  ddd«      j	                  «       |d   j                  ddd«      j	                  «       |d   j                  ddd«      j	                  «       }}}t        |d	   | j                  d
«      \  }}|j                  d   }	|j                  }
t        j                  |d	   d   j                  dd | j                  |
¬«      | j                  d   z  }	 |d   j                  dd«      }t        j                  ||d   j                  dd«      |d   j                  dd«      fd«      }|dd…df   t        |d   «      z  |dd…df   t        |d   «      z  }}||dk\  |dk\  z     }| j                  |j                  | j                  «      |	|g d¢   ¬«      }|j                  dd«      \  }}|j!                  dd¬«      j#                  d«      }| j)                  |||«      }|j+                  «       j-                  «       }|ddd…fxx   |z  cc<   | j/                  |j-                  «       j1                  «       |j3                  |j                  «      ||z  |||«      \  }}}}}t5        |j!                  «       d«      }| j7                  ||j                  |
«      «      j!                  «       |z  |d<   |j!                  «       r`|ddd…fxx   |z  cc<   | j9                  |||||||||«	      \  |d<   |d<   |j!                  d«      |   }| j;                  |||||«      |d<   n|dxx   |dz  j!                  «       z  cc<   |dxx   | j<                  j>                  z  cc<   |dxx   | j<                  j@                  z  cc<   |dxx   | j<                  jB                  z  cc<   |dxx   | j<                  jD                  z  cc<   ||	z  |j-                  «       fS # t$        $ r}t'        d«      |‚d}~ww xY w)zBCalculate and return the loss for oriented bounding box detection.r—   rç   r  r   ro   r   r  Úangler	  r¸   NrE   rù   rU   r
  r  r˜   r  r  )r   r˜   TrV   r  uh  ERROR â�Œ OBB dataset incorrectly formatted or not a OBB dataset.
This error can occur when incorrectly training a 'OBB' model on a 'detect' dataset, i.e. 'yolo train model=yolo26n-obb.pt data=dota8.yaml'.
Verify your dataset is a correctly formatted 'OBB' dataset using 'data=dota8.yaml' as an example.
See https://docs.ultralytics.com/datasets/obb/ for help..rž   )#rA   rë   rF   r  r  r   rz   r\   rG   rB   r[   r  r*   rý   rI   r  r,   r  ÚRuntimeErrorÚ	TypeErrorr  ru  r  rÞ   r&   r  rî   r®   rß   Úcalculate_angle_lossrØ   r  r
  r  rº  )r   r  r  r1   r  r  Ú
pred_anglert   r  rô   rG   ry   rù   ró   ÚrwÚrhr  r  r  rÈ   rs   Úbboxes_for_assignerrµ   ru   rv   rx   rw   r0   s                               r   r1   zv8OBBLoss.lossð  s  € ä�{‰{˜1 T§[¡[Ô1ˆà�'‰N×"Ñ" 1 a¨Ó+×6Ñ6Ó8Ø�(‰O×#Ñ# A q¨!Ó,×7Ñ7Ó9Ø�'‰N×"Ñ" 1 a¨Ó+×6Ñ6Ó8ð #-�[ˆô
 (4°E¸'±NÀDÇKÁKÐQTÓ'UÑ$ˆ�}Ø×%Ñ% aÑ(ˆ
à×!Ñ!ˆÜ—‘˜U 7™^¨AÑ.×4Ñ4°Q°RÐ8ÀÇÁÐTYÔZÐ]a×]hÑ]hÐijÑ]kÑkˆð	Ø˜kÑ*×/Ñ/°°AÓ6ˆIÜ—i‘i ¨E°%©L×,=Ñ,=¸bÀ!Ó,DÀeÈHÁo×FZÑFZÐ[]Ð_`ÓFaÐ bÐdeÓfˆGØšQ ˜T‘]¤U¨5°©8£_Ñ4°gºaÀ¸d±mÄeÈEÐRSÉHÃoÑ6U�ˆBØ˜r Q™w¨2°©7Ñ3Ñ4ˆGØ—o‘o g§j¡j°·±Ó&=¸zÐX]Ò^jÑXk�oÓlˆGØ#*§=¡=°¸Ó#;Ñ ˆI�yØ—m‘m A¨t�mÓ4×8Ñ8¸Ó=ˆGð ×&Ñ& }°kÀ:ÓNˆà)×/Ñ/Ó1×8Ñ8Ó:Ðà˜C  ! ˜GÓ$¨Ñ5Ó$Ø6:·m±mØ×ÑÓ ×(Ñ(Ó*Ø×$Ñ$ Y§_¡_Ó5Ø˜MÑ)ØØØó7
Ñ3ˆˆ=˜-¨°!ô   × 1Ñ 1Ó 3°QÓ7Ðð —(‘(˜;¨×(8Ñ(8¸Ó(?Ó@×DÑDÓFÐIZÑZˆˆQ‰ð �;‰;Œ=Ø˜#˜r ˜r˜'Ó" mÑ3Ó"Ø#Ÿ~™~ØØØØØØ!ØØØó
 ÑˆD�‰G�T˜!‘Wð #×&Ñ& rÓ*¨7Ñ3ˆFØ×/Ñ/Ø˜]¨G°VÐ=NóˆD�ŠGð �‹G˜
 Q™×+Ñ+Ó-Ñ-‹GàˆQ‹�4—8‘8—<‘<Ñ‹ØˆQ‹�4—8‘8—<‘<Ñ‹ØˆQ‹�4—8‘8—<‘<Ñ‹ØˆQ‹�4—8‘8—>‘>Ñ!‹à�jÑ  $§+¡+£-Ð/Ð/øôq ò 	Üð[óð ðûð	ús   ÄC;P Ð	P-ÐP(Ð(P-c                ó2  — | j                   rh|j                  \  }}}|j                  ||d|dz  «      j                  d«      j	                  | j
                  j                  |j                  «      «      }t        j                  t        |||«      |fd¬«      S )a°  Decode predicted object bounding box coordinates from anchor points and distribution.

        Args:
            anchor_points (torch.Tensor): Anchor points, (h*w, 2).
            pred_dist (torch.Tensor): Predicted rotated distance, (bs, h*w, 4).
            pred_angle (torch.Tensor): Predicted angle, (bs, h*w, 1).

        Returns:
            (torch.Tensor): Predicted rotated bounding boxes with angles, (bs, h*w, 5).
        r—   rž   rU   rŸ   )rÛ   r\   r[   rÿ   r   rá   r  rG   rA   r  r   )r   rt   r]   r¾  r  r  r  s          r   r  zv8OBBLoss.bbox_decodeA  s|   € ð �<Š<Ø—o‘o‰GˆAˆq�!Ø!Ÿ™ q¨!¨Q°°Q±Ó7×?Ñ?ÀÓB×IÑIÈ$Ï)É)Ï.É.ÐYb×YhÑYhÓJiÓjˆIÜ�y‰yœ) I¨z¸=ÓIÈ:ÐVÐ\^Ô_Ð_r    c                óž  — |d   }|d   }|d   }	|d   }
t        j                  |dz   |dz   z  «      }t        j                  |dz   |dz  z  «      }|	|
z
  }|t        j                  |t        j
                  z  «      t        j
                  z  z
  }t        j                  d||   z  «      dz  }||   |z  }||z  }|j                  «       |z  S )aƒ  Calculate oriented angle loss.

        Args:
            pred_bboxes (torch.Tensor): Predicted bounding boxes with shape [N, 5] (x, y, w, h, theta).
            target_bboxes (torch.Tensor): Target bounding boxes with shape [N, 5] (x, y, w, h, theta).
            fg_mask (torch.Tensor): Foreground mask indicating valid predictions.
            weight (torch.Tensor): Loss weights for each prediction.
            target_scores_sum (torch.Tensor): Sum of target scores for normalization.
            lambda_val (int): Controls the sensitivity to aspect ratio.

        Returns:
            (torch.Tensor): The calculated angle loss.
        r‰  ).rž   ).r—   rÀ   ro   )rA   r‰   rÁ   ÚroundÚmathÚpiÚsinr,   )r   rs   ru   rx   r0   rw   Ú
lambda_valÚw_gtÚh_gtÚ
pred_thetaÚtarget_thetaÚlog_arÚscale_weightÚdelta_thetaÚdelta_theta_wrappedÚang_losss                   r   r½  zv8OBBLoss.calculate_angle_lossS  sà   € ð ˜VÑ$ˆØ˜VÑ$ˆØ  Ñ(ˆ
Ø$ VÑ,ˆä—‘˜D 4™K¨D°4©KÑ8Ó9ˆÜ—y‘y 6¨1¡9 °¸Q±Ñ!?Ó@ˆà  <Ñ/ˆØ)¬E¯K©K¸ÄdÇgÁgÑ8MÓ,NÔQU×QXÑQXÑ,XÑXÐÜ—9‘9˜QÐ!4°WÑ!=Ñ=Ó>À!ÑCˆà Ñ(¨8Ñ3ˆØ˜fÑ$ˆà�|‰|‹~Ð 1Ñ1Ð1r    r-  ©rã   r0  r1  r4  )rt   r5   r]   r5   r¾  r5   r6   r5   )rž   )
r8   r9   r:   r;   r   rý   r1   r  r½  r<   r=   s   @r   r´  r´  Î  sG   ø„ ÙnöGóó$O0ðb`Ø)ð`Ø6Bð`ØP\ð`à	ó`÷$2r    r´  c                  ó   — e Zd ZdZd„ Zdd„Zy)ÚE2EDetectLossúGCriterion class for computing training losses for end-to-end detection.c                óL   — t        |d¬«      | _        t        |d¬«      | _        y)zcInitialize E2EDetectLoss with one-to-many and one-to-one detection losses using the provided model.r.  r¶  r   N)rÊ   Úone2manyÚone2one)r   r×   s     r   r   zE2EDetectLoss.__init__v  s   € ä'¨¸Ô;ˆŒÜ& u°qÔ9ˆ�r    c                ó¸   — t        |t        «      r|d   n|}|d   }| j                  ||«      }|d   }| j                  ||«      }|d   |d   z   |d   |d   z   fS )r(  r   r×  rØ  r   )r$  r³   r×  rØ  )r   r  r  r×  Úloss_one2manyrØ  Úloss_one2ones          r   rc   zE2EDetectLoss.__call__{  sr   € ä& u¬eÔ4��a’¸%ˆØ˜Ñ$ˆØŸ™ h°Ó6ˆØ˜	Ñ"ˆØ—|‘| G¨UÓ3ˆØ˜QÑ ,¨q¡/Ñ1°=ÀÑ3CÀlÐSTÁoÑ3UÐUÐUr    Nr²  )r8   r9   r:   r;   r   rc   r5  r    r   rÔ  rÔ  s  s   „ ÙQò:ô
Vr    rÔ  c                  ó2   — e Zd ZdZefd„Zdd„Zdd„Zd	d„Zy)
ÚE2ELossrÕ  c                óØ   —  ||d¬«      | _          ||dd¬«      | _        d| _        d| _        d| _        | j                  | j                  z
  | _        | j                  | _        d	| _        y
)z]Initialize E2ELoss with one-to-many and one-to-one detection losses using the provided model.r.  r¶  é   r   )râ   rã   r   rD   gš™™™™™é?gš™™™™™¹?N)r×  rØ  ÚupdatesÚtotalÚo2mÚo2oÚo2m_copyÚ	final_o2m)r   r×   Úloss_fns      r   r   zE2ELoss.__init__ˆ  s[   € á °Ô3ˆŒÙ˜u¨q¸AÔ>ˆŒØˆŒØˆŒ
àˆŒØ—:‘: §¡Ñ(ˆŒØŸ™ˆŒàˆ�r    c                ó  — | j                   j                  |«      }|d   |d   }}| j                   j                  ||«      }| j                  j                  ||«      }|d   | j                  z  |d   | j
                  z  z   |d   fS )r(  r×  rØ  r   r   )r×  r&  r1   rØ  râ  rã  )r   r  r  r×  rØ  rÚ  rÛ  s          r   rc   zE2ELoss.__call__•  s…   € à—‘×*Ñ*¨5Ó1ˆØ! *Ñ-¨u°YÑ/?�'ˆØŸ™×*Ñ*¨8°UÓ;ˆØ—|‘|×(Ñ(¨°%Ó8ˆØ˜QÑ $§(¡(Ñ*¨\¸!©_¸t¿x¹xÑ-GÑGÈÐVWÉÐXÐXr    c                ó¾   — | xj                   dz  c_         | j                  | j                   «      | _        t        | j                  | j                  z
  d«      | _        y)zUUpdate the weights for one-to-many and one-to-one losses based on the decay schedule.r   r   N)rà  Údecayrâ  rî   rá  rã  )r   s    r   ÚupdatezE2ELoss.update�  s?   € à�Š˜Ñ�Ø—:‘:˜dŸl™lÓ+ˆŒÜ�t—z‘z D§H¡HÑ,¨aÓ0ˆ�r    c                óÊ   — t        d|t        | j                  j                  j                  dz
  d«      z  z
  d«      | j                  | j
                  z
  z  | j
                  z   S )zSCalculate the decayed weight for one-to-many loss based on the current update step.r   r   )rî   rØ  rØ   Úepochsrä  rå  )r   Úxs     r   ré  zE2ELoss.decay£  sV   € ä�1�qœ3˜tŸ|™|×/Ñ/×6Ñ6¸Ñ:¸AÓ>Ñ>Ñ>ÀÓBÀdÇmÁmÐVZ×VdÑVdÑFdÑeÐhl×hvÑhvÑvÐvr    Nr²  )r6   rg   )r6   r*   )	r8   r9   r:   r;   rÊ   r   rc   rê  ré  r5  r    r   rÝ  rÝ  …  s   „ ÙQà&5ó óYó1ôwr    rÝ  c                  ó:   — e Zd ZdZdd	d„Zd
d„Zdd„Zdd„Zdd„Zy)ÚTVPDetectLosszOCriterion class for computing training losses for text-visual prompt detection.Nc                ó   — t        |||«      | _        | j                  j                  | _        | j                  j                  | _        | j                  j
                  | _        | j                  j                  | _        y)z^Initialize TVPDetectLoss with task-prompt and visual-prompt criteria using the provided model.N)	rÊ   Úvp_criterionrØ   rÙ   Úori_ncrÚ   Úori_norR   Úori_reg_max)r   r×   râ   rã   s       r   r   zTVPDetectLoss.__init__«  s`   € ä+¨E°8¸YÓGˆÔà×$Ñ$×(Ñ(ˆŒØ×'Ñ'×*Ñ*ˆŒØ×'Ñ'×*Ñ*ˆŒØ×,Ñ,×4Ñ4ˆÕr    c                ó8   — | j                   j                  |«      S )r#  )rñ  r&  r%  s     r   r&  zTVPDetectLoss.parse_output´  s   € à× Ñ ×-Ñ-¨eÓ4Ð4r    c                óD   — | j                  | j                  |«      |«      S )ú4Calculate the loss for text-visual prompt detection.r)  r*  s      r   rc   zTVPDetectLoss.__call__¸  ó   € à�y‰y˜×*Ñ*¨5Ó1°5Ó9Ð9r    c                ó&  — | j                   |d   j                  d   k(  r>t        j                  d| j                  j
                  d¬«      }||j                  «       fS | j                  |«      |d<   | j	                  ||«      }|d   d   }||d   fS )r÷  r  r   rž   T©rF   Úrequires_gradr   ©rò  r\   rA   rë   rñ  rF   r  Ú_get_vp_features)r   r  r  r1   Úvp_lossÚbox_losss         r   r1   zTVPDetectLoss.loss¼  ó�   € à�;‰;˜% ™/×/Ñ/°Ñ2Ò2Ü—;‘;˜q¨×):Ñ):×)AÑ)AÐQUÔVˆDØ˜Ÿ™›Ð&Ð&à×/Ñ/°Ó6ˆˆh‰Ø×#Ñ# E¨5Ó1ˆØ˜1‘:˜a‘=ˆØ˜ ™Ð#Ð#r    c                óÜ   — |d   }|j                   d   }|| j                  _        || j                  j                  dz  z   | j                  _        || j                  j
                  _        |S )z5Extract visual-prompt features from the model output.r  r   r—   )r\   rñ  rÙ   rR   rÚ   rÞ   rÐ   )r   r  r  Úvncs       r   rý  zTVPDetectLoss._get_vp_featuresÇ  sc   € à�x‘ˆØ�l‰l˜1‰oˆà"ˆ×ÑÔØ" T×%6Ñ%6×%>Ñ%>ÀÑ%BÑBˆ×ÑÔØ14ˆ×Ñ×"Ñ"Ô.Øˆr    r-  rÒ  )r6   r2  r²  r4  )r  r2  r6   zlist[torch.Tensor])	r8   r9   r:   r;   r   r&  rc   r1   rý  r5  r    r   rï  rï  ¨  s   „ ÙYô5ó5ó:ó	$ôr    rï  c                  ó4   ‡ — e Zd ZdZdˆ fd„	Zdd„Zdd„Zˆ xZS )ÚTVPSegmentLosszRCriterion class for computing training losses for text-visual prompt segmentation.c                ó|   •— t         ‰| �  |«       t        ||«      | _        | j                  j                  | _        y)z_Initialize TVPSegmentLoss with task-prompt and visual-prompt criteria using the provided model.N)r   r   r7  rñ  rØ   )r   r×   râ   r   s      €r   r   zTVPSegmentLoss.__init__Õ  s2   ø€ ä‰Ñ˜ÔÜ.¨u°hÓ?ˆÔØ×$Ñ$×(Ñ(ˆ�r    c                óD   — | j                  | j                  |«      |«      S )ú7Calculate the loss for text-visual prompt segmentation.r)  r*  s      r   rc   zTVPSegmentLoss.__call__Û  rø  r    c                ó&  — | j                   |d   j                  d   k(  r>t        j                  d| j                  j
                  d¬«      }||j                  «       fS | j                  |«      |d<   | j	                  ||«      }|d   d   }||d   fS )r  r  r   r—   Trú  r   ro   rü  )r   r  r  r1   rþ  Úcls_losss         r   r1   zTVPSegmentLoss.lossß  r   r    )r.  r²  )r8   r9   r:   r;   r   rc   r1   r<   r=   s   @r   r  r  Ò  s   ø„ Ù\õ)ó:÷	$r    r  )4Ú
__future__r   rÅ  Útypingr   rA   Útorch.nnr¬   Útorch.nn.functionalÚ
functionalr(   Úultralytics.utils.metricsr   r   Úultralytics.utils.opsr   r   r	   Úultralytics.utils.talr
   r   r   r   r   Úultralytics.utils.torch_utilsr   Úmetricsr   r   Útalr   r   ÚModuler   r?   rP   ri   rƒ   r”   rš   r¨   rº   rÊ   r7  rf  r‘  r¯  r´  rÔ  rÝ  rï  r  r5  r    r   ú<module>r     sM  ðõ #ã Ý ã Ý ß Ð ç ;ß AÑ Aß uÕ uÝ 2ç &ß %ô�B—I‘Iô ô@ "�—	‘	ô  "ôF!ˆR�Y‰Yô !ô*,"ˆr�y‰yô ,"ô^3ˆb�i‰iô 3ôl,"�hô ,"ô^˜2Ÿ9™9ô ôBe�"—)‘)ô eô0W�2—9‘9ô W÷&P.ñ P.ôfZ$˜ô Z$ôz[(�ô [(ô|f2�ô f2÷R#ñ #ôb2�ô b2÷JVñ V÷$ wñ  w÷F'ñ 'ôT$�]õ $r    