Ë
    ÿÍ:j¦  ã                   ó  — d dl mZ d dlZd dlmZ d dlmZ d dlmZ 	 ddej                  dej                  de
fd„Zddej                  de
fd	„Zedd
ede
fd„«       Z G d„ dej                  «      Zej                  dded
ede
fd„«       Zy)é    )ÚsingledispatchN)ÚBaseWaveformTransform)ÚModelÚaugmentationÚmoduleÚwhenc                 ój  ‡— t        | ||¬«      Št        |d«      s(t        j                  «       |_        t        «       |_        ||j                  v rt        ||¬«       ‰|j                  |<   |dk(  rˆfd„}|j                  |«      }n|dk(  rˆfd„}|j                  |«      }|j                  |<   y)uJ  Register augmentation

    Parameters
    ----------
    augmentation : nn.Module
        Augmentation module.
    module : nn.Module
        Module whose input or output should be augmented.
    when : {'input', 'output'}
        Whether to apply augmentation on the input or the output.
        Defaults to 'input'.

    Usage
    -----

    class Net(nn.Module):
        def __init__(self):
            super().__init__()
            self.spectogram = Spectrogram()
            self.other_layers = nn.Identity()

        def forward(self, waveforms):
            spectrogram = self.spectrogram(waveforms)
            return self.other_layers = other_layers

    net = Net()

    class AddNoise(nn.Module):
        def forward(self, waveforms):
            if not self.training:
                return waveforms

            augmented_waveforms = ...
            return augmented_waveforms

    # AddNoise will be automatically applied to `net` input
    register_augmentation(AddNoise(), net, when='input')

    class SpecAugment(nn.Module):
        def forward(self, spectrograms):
            if not self.training:
                return spectrograms

            augmented_spectrograms = ...
            return augmented_spectrograms

    # SpecAugment will be automatically applied to `net.spectrogram` output
    register_augmentation(SpecAugment(), net.spectrogram, when='output')

    #Â deactivate augmentations
    net.eval()  # or net.train(mode=False)

    #Â reactivate augmentations
    net.train()

    #Â unregister "AddNoise" augmentation
    unregister_augmentation(net, when='input')

    ©r   Ú__augmentationÚinputc                 ó   •—  ‰|Ž S ©N© )Úaugmented_moduler   Úwrapped_augmentations     €úy/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/pyannote/audio/augmentation/registry.pyÚ
input_hookz)register_augmentation.<locals>.input_hooko   s   ø€ Ù'¨Ð/Ð/ó    Úoutputc                 ó   •—  ‰|«      S r   r   )r   r   r   r   s      €r   Úoutput_hookz*register_augmentation.<locals>.output_hookv   s   ø€ Ù'¨Ó/Ð/r   N)
Úwrap_augmentationÚhasattrÚnnÚ
ModuleDictr   ÚdictÚ__augmentation_handleÚunregister_augmentationÚregister_forward_pre_hookÚregister_forward_hook)r   r   r   r   Úhandler   r   s         @r   Úregister_augmentationr"       s¯   ø€ ôB -¨\¸6ÈÔMÐä�6Ð+Ô,Ü "§¡£ˆÔÜ'+£vˆÔ$ð ˆv×$Ñ$Ñ$Ü ¨TÕ2à"6€F×Ñ˜$Ñàˆw‚ô	0ð ×1Ñ1°*Ó=‰à	�Ò	ô	0ð ×-Ñ-¨kÓ:ˆà)/€F× Ñ  Ò&r   c                 óÆ   — t        | d«      r|| j                  vrt        d|› d�«      ‚| j                  |= | j                  j	                  |«      }|j                  «        y)ad  Unregister augmentation

    Parameters
    ----------
    module : nn.Module
        Module whose augmentation should be removed.
    when : {'input', 'output'}
        Whether to remove augmentation of the input or the output.
        Defaults to 'input'.

    Raises
    ------
    ValueError if module has no corresponding registered augmentation.
    r   zModule has no registered u   Â augmentation.N)r   r   Ú
ValueErrorr   ÚpopÚremove)r   r   r!   s      r   r   r   ~   s`   € ô  �FÐ,Ô-°4¸v×?TÑ?TÑ3TÜÐ4°T°F¸/ÐJÓKÐKà×Ñ˜dÐ#ð ×)Ñ)×-Ñ-¨dÓ3€FØ
‡M�M…Or   Úmodelc                 ó   — | S r   r   ©r   r'   r   s      r   r   r   ˜   s   € àÐr   c                   óh   ‡ — e Zd Z	 ddededefˆ fd„Zdej                  dej                  fd„Z	ˆ xZ
S )	Ú,TorchAudiomentationsWaveformTransformWrapperr   r'   r   c                 óô   •— t         ‰| �  «        || _        t        |t        «      s#t        d|j                  j                  › d�«      ‚|dk7  rt        d|› d�«      ‚|j                  j                  | _        y )Nzttorch-audiomentations waveform transforms can only be applied to `pyannote.audio.Model` instances: you tried with a z
 instance.r   zetorch-audiomentations waveform transforms can only be applied to the model input: you tried with the ú.)ÚsuperÚ__init__r   Ú
isinstancer   Ú	TypeErrorÚ	__class__Ú__name__r$   ÚaudioÚsample_rateÚsample_rate_)Úselfr   r'   r   r2   s       €r   r/   z5TorchAudiomentationsWaveformTransformWrapper.__init__£   s‡   ø€ ô 	‰ÑÔà(ˆÔä˜%¤Ô'Üð$Ø$)§O¡O×$<Ñ$<Ð#=¸ZðIóð ð �7Š?Üð&Ø&* V¨1ð.óð ð
 "ŸK™K×3Ñ3ˆÕr   Ú	waveformsÚreturnc                 óP   — | j                  || j                  ¬«      j                  S )N)Úsamplesr5   )r   r6   r;   )r7   r8   s     r   Úforwardz4TorchAudiomentationsWaveformTransformWrapper.forward·   s*   € Ø× Ñ Ø¨4×+<Ñ+<ð !ó 
ç
‰'ð	r   ©r   )r3   Ú
__module__Ú__qualname__r   r   Ústrr/   ÚtorchÚTensorr<   Ú__classcell__)r2   s   @r   r+   r+   ¢   s>   ø„ àMTñ4Ø1ð4Ø:?ð4ØGJõ4ð( §¡ð °%·,±,÷ r   r+   c                 ó   — t        | ||¬«      S )Nr
   )r+   r)   s      r   Ú_rE   ½   s   € ä7¸ÀeÐRVÔWÐWr   r=   )Ú	functoolsr   rA   Útorch.nnr   Ú/torch_audiomentations.core.transforms_interfacer   Úpyannote.audio.core.modelr   ÚModuler@   r"   r   r   r+   ÚregisterrE   r   r   r   ú<module>rL      sÅ   ðõ. %ã Ý Ý Qå +ð ñ[0Ø—)‘)ð[0à�I‰Ið[0ð ó[0ñ| B§I¡Ið °Só ð4 ñ¨5ð ¸ò ó ðô°2·9±9ô ð6 ×ÑñXÐ)ð X°%ð X¸sò Xó ñXr   