Ë
    ÿÍ:jt  ã                   ó°   — d dl mZ d dlZd dlmZ d dlmc mZ dej                  dej                  dej                  fd„Z	 G d„ dej                  «      Zy)	é    )ÚOptionalNÚ	sequencesÚweightsÚreturnc                 óÊ  — |j                  d¬«      }|j                  d¬«      dz   }t        j                  | |z  d¬«      |z  }t        j                  | |j                  d«      z
  «      }t        j                  |«      j                  d¬«      }t        j                  ||z  d¬«      |||z  z
  dz   z  }t        j                  |«      }t        j
                  ||gd¬«      S )a  Helper function to compute statistics pooling

    Assumes that weights are already interpolated to match the number of frames
    in sequences and that they encode the activation of only one speaker.

    Parameters
    ----------
    sequences : (batch, features, frames) torch.Tensor
        Sequences of features.
    weights : (batch, frames) torch.Tensor
        (Already interpolated) weights.

    Returns
    -------
    output : (batch, 2 * features) torch.Tensor
        Concatenation of mean and (unbiased) standard deviation.
    é   ©Údimé   g:Œ0âŽyE>)Ú	unsqueezeÚsumÚtorchÚsquareÚsqrtÚcat)r   r   Úv1ÚmeanÚdx2Úv2ÚvarÚstds           úy/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/pyannote/audio/models/blocks/pooling.pyÚ_poolr      sÈ   € ð& ×Ñ AÐÓ&€Gð 
�‰˜ˆÓ	˜dÑ	"€BÜ�9‰9�Y Ñ(¨aÔ0°2Ñ5€Dä
�,‰,�y 4§>¡>°!Ó#4Ñ4Ó
5€CÜ	�‰�gÓ	×	"Ñ	" qÐ	"Ó	)€Bä
�)‰)�C˜'‘M qÔ
)¨R°"°r±'©\¸DÑ-@Ñ
A€CÜ
�*‰*�S‹/€Cä�9‰9�d˜C�[ aÔ(Ð(ó    c                   ój   — e Zd ZdZ	 ddej
                  deej
                     dej
                  fd„Zy)Ú	StatsPoolzÒStatistics pooling

    Compute temporal mean and (unbiased) standard deviation
    and returns their concatenation.

    Reference
    ---------
    https://en.wikipedia.org/wiki/Weighted_arithmetic_mean

    Nr   r   r   c                 ó  — |€>|j                  d¬«      }|j                  dd¬«      }t        j                  ||gd¬«      S |j	                  «       dk(  rd}|j                  d¬«      }nd}|j                  «       \  }}}|j                  «       \  }}}	||	k7  rt        j                  ||d	¬
«      }t        j                  t        |«      D �
cg c]  }
t        ||dd…|
dd…f   «      ‘Œ c}
d¬«      }|s|j                  d¬«      S |S c c}
w )a�  Forward pass

        Parameters
        ----------
        sequences : (batch, features, frames) torch.Tensor
            Sequences of features.
        weights : (batch, frames) or (batch, speakers, frames) torch.Tensor, optional
            Compute weighted mean and standard deviation, using provided `weights`.

        Note
        ----
        `sequences` and `weights` might use a different number of frames, in which case `weights`
        are interpolated linearly to reach the number of frames in `sequences`.

        Returns
        -------
        output : (batch, 2 * features) or (batch, speakers, 2 * features) torch.Tensor
            Concatenation of mean and (unbiased) standard deviation. When `weights` are
            provided with the `speakers` dimension, `output` is computed for each speaker
            separately and returned as (batch, speakers, 2 * channel)-shaped tensor.
        Néÿÿÿÿr	   r   )r
   Ú
correctionr   FTÚnearest)ÚsizeÚmode)r   r   r   r   r
   r   r!   ÚFÚinterpolateÚstackÚranger   Úsqueeze)Úselfr   r   r   r   Úhas_speaker_dimensionÚ_Ú
num_framesÚnum_speakersÚnum_weightsÚspeakerÚoutputs               r   ÚforwardzStatsPool.forwardL   s  € ð2 ˆ?Ø—>‘> b�>Ó)ˆDØ—-‘- B°1�-Ó5ˆCÜ—9‘9˜d C˜[¨bÔ1Ð1à�;‰;‹=˜AÒØ$)Ð!Ø×'Ñ'¨AÐ'Ó.‰Gð %)Ð!ð %Ÿ>™>Ó+Ñˆˆ1ˆjØ'.§|¡|£~Ñ$ˆˆ<˜Ø˜Ò$Ü—m‘m G°*À9ÔMˆGä—‘ô  % \Ó2öàô �i ª¨G²Q¨Ñ!7Õ8òð ô
ˆñ %Ø—>‘> a�>Ó(Ð(àˆùòs   ÃD)N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚTensorr   r0   © r   r   r   r   @   s<   „ ñ	ð JNñ6ØŸ™ð6Ø08¸¿¹Ñ0Fð6à	�‰ô6r   r   )Útypingr   r   Útorch.nnÚnnÚtorch.nn.functionalÚ
functionalr#   r5   r   ÚModuler   r6   r   r   ú<module>r=      sO   ðõ. ã Ý ß Ð ð)�U—\‘\ð )¨E¯L©Lð )¸U¿\¹\ó )ôDB�—	‘	õ Br   