Ë
    ÿÍ:j¿  ã                   óä  — d Z ddlmZ ddlZddlmc mZ ddej                  deej                     fd„Z		 ddej                  dej                  deej                     dej                  fd	„Z
	 ddej                  dej                  deej                     dej                  fd
„Z	 	 ddej                  dej                  deej                     deej                     dej                  f
d„Zy)z0Frame-weighted versions of common loss functionsé    )ÚOptionalNÚtargetÚweightc                 ó¾   — | j                   d   }|�K|j                   d   |k7  r9t        j                  |j                  dd«      |dd¬«      j                  dd«      }|S )a²  Interpolate weight to match target frame resolution

    Parameters
    ----------
    target : torch.Tensor
        Target with shape (batch_size, num_frames) or (batch_size, num_frames, num_classes)
    weight : torch.Tensor, optional
        Frame weight with shape (batch_size, num_frames_weight, 1).

    Returns
    -------
    weight : torch.Tensor
        Interpolated frame weight with shape (batch_size, num_frames, 1).
    é   é   ÚlinearF)ÚsizeÚmodeÚalign_corners)ÚshapeÚFÚinterpolateÚ	transpose)r   r   Ú
num_framess      ún/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/pyannote/audio/utils/loss.pyr   r       sd   € ð  —‘˜a‘€JØÐ˜fŸl™l¨1™o°Ò;Ü—‘Ø×Ñ˜Q Ó"ØØØô	
÷
 ‰)�A�q‹/ð 	ð €Mó    Ú
predictionÚreturnc                 ó:  — t        |j                  «      dk(  r|j                  d¬«      }|€$t        j                  | |j                  «       «      S t        ||¬«      }t        j                  | |j                  «       |j                  |j                  «      ¬«      S )a   Frame-weighted binary cross entropy

    Parameters
    ----------
    prediction : torch.Tensor
        Prediction with shape (batch_size, num_frames, num_classes).
    target : torch.Tensor
        Target with shape (batch_size, num_frames) for binary or multi-class classification,
        or (batch_size, num_frames, num_classes) for multi-label classification.
    weight : (batch_size, num_frames, 1) torch.Tensor, optional
        Frame weight with shape (batch_size, num_frames, 1).

    Returns
    -------
    loss : torch.Tensor
    r   ©Údim©r   )Úlenr   Ú	unsqueezer   Úbinary_cross_entropyÚfloatr   Úexpand)r   r   r   s      r   r   r   ;   s�   € ô. ˆ6�<‰<Ó˜AÒØ×!Ñ! aÐ!Ó(ˆà€~Ü×%Ñ% j°&·,±,³.ÓAÐAô ˜V¨FÔ3ˆä×%Ñ%Ø˜Ÿ™›¨v¯}©}¸V¿\¹\Ó/Jô
ð 	
r   c                 óx  — t        |j                  «      dk(  r|j                  d¬«      }t        j                  | |j                  «       d¬«      }|€t        j                  |«      S t        ||¬«      j                  |j                  «      }t        j                  ||z  «      t        j                  |«      z  S )a#  Frame-weighted mean-squared error loss

    Parameters
    ----------
    prediction : torch.Tensor
        Prediction with shape (batch_size, num_frames, num_classes).
    target : torch.Tensor
        Target with shape (batch_size, num_frames) for binary or multi-class classification,
        or (batch_size, num_frames, num_classes) for multi-label classification.
    weight : (batch_size, num_frames, 1) torch.Tensor, optional
        Frame weight with shape (batch_size, num_frames, 1).

    Returns
    -------
    loss : torch.Tensor
    r   r   Únone)Ú	reductionr   )r   r   r   r   Úmse_lossr   ÚtorchÚmeanr   r   Úsum)r   r   r   Úlossess       r   r"   r"   a   s–   € ô. ˆ6�<‰<Ó˜AÒØ×!Ñ! aÐ!Ó(ˆä�Z‰Z˜
 F§L¡L£N¸fÔE€Fð €~Ü�z‰z˜&Ó!Ð!ô ˜V¨FÔ3×:Ñ:¸6¿<¹<ÓHˆô �y‰y˜ &™Ó)¬E¯I©I°fÓ,=Ñ=Ð=r   Úclass_weightc                 ó†  — | j                   d   }t        j                  | j                  d|«      |j                  d«      |d¬«      j                  |j                   «      }|€t	        j
                  |«      S t        ||¬«      j                  d¬«      }t	        j                  ||z  «      t	        j                  |«      z  S )a  Frame-weighted negative log-likelihood loss

    Parameters
    ----------
    prediction : torch.Tensor
        Prediction with shape (batch_size, num_frames, num_classes).
    target : torch.Tensor
        Target with shape (batch_size, num_frames)
    class_weight : (num_classes, ) torch.Tensor, optional
        Class weight with shape (num_classes,  )
    weight : (batch_size, num_frames, 1) torch.Tensor, optional
        Frame weight with shape (batch_size, num_frames, 1).

    Returns
    -------
    loss : torch.Tensor
    r   éÿÿÿÿr    )r   r!   r   r   )	r   r   Únll_lossÚviewr#   r$   r   Úsqueezer%   )r   r   r'   r   Únum_classesr&   s         r   r*   r*   ‰   s©   € ð0 ×"Ñ" 1Ñ%€Kä�Z‰ZØ�‰˜˜KÓ(à�‰�B‹ààô÷ �dˆ6�<‰<Óð ð €~Ü�z‰z˜&Ó!Ð!ô ˜V¨FÔ3×;Ñ;ÀÐ;ÓBˆô �y‰y˜ &™Ó)¬E¯I©I°fÓ,=Ñ=Ð=r   )N)NN)Ú__doc__Útypingr   r#   Útorch.nn.functionalÚnnÚ
functionalr   ÚTensorr   r   r"   r*   © r   r   ú<module>r5      s  ðñ0 7å ã ß Ð ñ˜Ÿ™ð ¨h°u·|±|Ñ.Dó ð< &*ñ#
Ø—‘ð#
à�L‰Lð#
ð �U—\‘\Ñ"ð#
ð ‡\�\ó	#
ðR &*ñ%>Ø—‘ð%>à�L‰Lð%>ð �U—\‘\Ñ"ð%>ð ‡\�\ó	%>ðV ,0Ø%)ñ	->Ø—‘ð->à�L‰Lð->ð ˜5Ÿ<™<Ñ(ð->ð �U—\‘\Ñ"ð	->ð
 ‡\�\ô->r   