Ë
    ÿÍ:jJ  ã                   óf   — d dl mZ d dlmZmZmZmZ d dlmZ dej                  deeef   defd„Z
y)é    )Úpartial)ÚCallableÚDictÚSetÚTextNÚtrunkÚbranchesÚreturnc                 ó¦  ‡ — ˆ fd„}t        ‰ d«      r |«        g ‰ _        ˆ fd„}‰ j                  |«      }‰ j                  j                  |«       ˆ fd„}t	        |t
        «      s|D �ci c]  }||“Œ }}t        «       }|j                  «       D ]*  \  }}	|	|vrt        «       ||	<   ||	   j                  |«       Œ, ‰ j                  «       D ]J  \  }	}
|	|vrŒ||	   D ]8  }|
j                  t        ||«      «      }‰ j                  j                  |«       Œ: ŒL ˆ fd„}‰ j                  |«      }‰ j                  j                  |«       |S c c}w )a.  Add probing branches to a trunk module

    Parameters
    ----------
    trunk : nn.Module
        Multi-layer trunk.
    branches : {branch_name: layer_name} dict or [layer_name] list
        Indicate where to plug a probing branch.

    Returns
    -------
    revert : Callable
        Callable that, when called, removes probing branches.

    Usage
    -----

    Define a trunk made out of three consecutive layers

    >>> import torch.nn as nn
    >>> class Trunk(nn.Module):
    ...
    ...     def __init__(self):
    ...         super().__init__()
    ...         self.layer1 = nn.Linear(1, 2)
    ...         self.layer2 = nn.Linear(2, 3)
    ...         self.layer3 = nn.Linear(3, 4)
    ...
    ...     def forward(self, x):
    ...         return self.layer3(self.layer2(self.layer1(x)))

    >>> trunk = Trunk()
    >>> x = torch.tensor((0.,))
    >>> trunk(x)
    # tensor([ 0.4548, -0.1814,  0.9494,  1.0445], grad_fn=<AddBackward0>)

    Add two probing branches:
    - first one is called "probe1" and probes the output of "layer1"
    - second one is called "probe2" and probes the output of "layer3"

    >>> revert = probe(trunk, {"probe1": "layer1", "probe2": "layer3"})
    >>> trunk(x)
    # {'probe1': tensor([ 0.5854, -0.9685], grad_fn=<AddBackward0>),
    #  'probe2': tensor([ 0.4548, -0.1814,  0.9494,  1.0445], grad_fn=<AddBackward0>)}

    Use callback returned by `probe` to revert its effect

    >>> revert()
    >>> trunk(x)
    # tensor([ 0.4548, -0.1814,  0.9494,  1.0445], grad_fn=<AddBackward0>)

    For convenience, one can also define probes as a list of layers:

    >>> revert = probe(trunk, ['layer1', 'layer3'])
    >>> trunk(x)
    # {'layer1': tensor([ 0.5854, -0.9685], grad_fn=<AddBackward0>),
    #  'layer3': tensor([ 0.4548, -0.1814,  0.9494,  1.0445], grad_fn=<AddBackward0>)}
    c                  óP   •— ‰` ‰j                  D ]  } | j                  «        Œ ‰`y ©N)Ú__probeÚ__probe_handlesÚremove)Úhandler   s    €úo/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/pyannote/audio/utils/probe.pyr   zprobe.<locals>.removeZ   s,   ø€ ØˆMØ×+Ñ+ò 	ˆFØ�M‰M�Oð	àÑ!ó    r   c                 ó$   •— t        «       ‰_        y r   )Údictr   )ÚmoduleÚinputr   s     €r   Ú__probe_initzprobe.<locals>.__probe_inite   s   ø€ Ü›ˆ�r   c                 ó$   •— |‰j                   | <   y r   ©r   )Úbranch_namer   r   Úoutputr   s       €r   Ú__probe_appendzprobe.<locals>.__probe_appendk   s   ø€ Ø%+ˆ�‰�kÒ"r   c                 ó   •— ‰j                   S r   r   )r   r   r   r   s      €r   Ú__probe_returnzprobe.<locals>.__probe_return~   s   ø€ Ø�}‰}Ðr   )Úhasattrr   Úregister_forward_pre_hookÚappendÚ
isinstancer   ÚitemsÚsetÚaddÚnamed_modulesÚregister_forward_hookr   )r   r	   r   r   r   r   ÚbÚsehcnarbr   Ú
layer_nameÚlayerr   s   `           r   Úprober-      s\  ø€ ôx"ô ˆu�iÔ ÙŒà€EÔôð ×,Ñ,¨\Ó:€FØ	×Ñ× Ñ  Ô(ô,ô �h¤Ô%Ø"*Ö+˜Q�A�q‘DÐ+ˆÐ+ä $£€HØ#+§>¡>Ó#3ò .Ñˆ�ZØ˜XÑ%Ü#&£5ˆH�ZÑ Ø�Ñ× Ñ  Õ-ð.ð
 #×0Ñ0Ó2ò 1Ñˆ
�EØ˜XÑ%ØØ# JÑ/ò 	1ˆKØ×0Ñ0´¸ÈÓ1UÓVˆFØ×!Ñ!×(Ñ(¨Õ0ñ	1ð1ôð ×(Ñ(¨Ó8€FØ	×Ñ× Ñ  Ô(à€Mùò+ ,s   Á+
E)Ú	functoolsr   Útypingr   r   r   r   Útorch.nnÚnnÚModuler-   © r   r   ú<module>r4      s>   ðõ0 ß ,Ó ,å ðf�—‘ð f d¨4°¨:Ñ&6ð f¸8ô fr   