Ë
    ùÿæiä'  ã                   ór   — d dl Z d dlmZ d dlZd dlZd dlmZ g d¢Z G d„ d«      Zd„ Z	dd„Z
 G d	„ d
«      Zy)é    N)ÚOrderedDict)ÚAny)ÚRemovableHandleÚunserializable_hookÚwarn_if_has_hooksÚBackwardHookc                   óz   — e Zd ZU dZeed<   dZeed<   ddœdeded	dfd
„Zdd„Z	d„ Z
dd„Zdd„Zdededed	dfd„Zy)r   a]  
    A handle which provides the capability to remove a hook.

    Args:
        hooks_dict (dict): A dictionary of hooks, indexed by hook ``id``.
        extra_dict (Union[dict, List[dict]]): An additional dictionary or list of
            dictionaries whose keys will be deleted when the same keys are
            removed from ``hooks_dict``.
    Úidr   Únext_idN)Ú
extra_dictÚ
hooks_dictr   Úreturnc                óJ  — t        j                  |«      | _        t        j                  | _        t        xj                  dz  c_        d| _        t        |t        «      rt        j                  |«      f| _        y t        |t        «      rt        d„ |D «       «      | _        y y )Né   © c              3   óF   K  — | ]  }t        j                  |«      –— Œ y ­w©N©ÚweakrefÚref©Ú.0Úds     úf/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torch/utils/hooks.pyÚ	<genexpr>z+RemovableHandle.__init__.<locals>.<genexpr>!   s   è ø€ Ð'KÁ
¸1¬¯©°A¯Á
ùó   ‚!)r   r   Úhooks_dict_refr   r   r
   Úextra_dict_refÚ
isinstanceÚdictÚlistÚtuple)Úselfr   r   s      r   Ú__init__zRemovableHandle.__init__   sy   € Ü%Ÿk™k¨*Ó5ˆÔÜ!×)Ñ)ˆŒÜ×Ò 1Ñ$Õà%'ˆÔÜ�j¤$Ô'Ü#*§;¡;¨zÓ#:Ð"<ˆDÕÜ˜
¤DÔ)Ü"'Ñ'KÁ
Ó'KÓ"KˆDÕð *ó    c                 óÌ   — | j                  «       }|�| j                  |v r|| j                  = | j                  D ](  } |«       }|€Œ| j                  |v sŒ|| j                  = Œ* y r   )r   r
   r   )r#   r   r   r   s       r   ÚremovezRemovableHandle.remove#   sa   € Ø×(Ñ(Ó*ˆ
ØÐ! d§g¡g°Ñ&;Ø˜4Ÿ7™7Ð#à×&Ô&ˆCÙ›ˆJØÑ%¨$¯'©'°ZÒ*?Ø˜tŸw™wÑ'ñ 'r%   c                 óÀ   — | j                   €| j                  «       | j                  fS | j                  «       | j                  t        d„ | j                   D «       «      fS )Nc              3   ó*   K  — | ]  } |«       –— Œ y ­wr   r   )r   r   s     r   r   z/RemovableHandle.__getstate__.<locals>.<genexpr>1   s   è ø€ Ð9_ÑK^ÀC¹#¿%ÑK^ùs   ‚)r   r   r
   r"   ©r#   s    r   Ú__getstate__zRemovableHandle.__getstate__-   sQ   € Ø×ÑÐ&Ø×'Ñ'Ó)¨4¯7©7Ð3Ð3à×'Ñ'Ó)¨4¯7©7´EÑ9_È4×K^ÒK^Ó9_Ó4_Ð`Ð`r%   c                 ól  — |d   €#t        j                  t        «       «      | _        nt        j                  |d   «      | _        |d   | _        t        t        j                  | j                  dz   «      t        _        t        |«      dk  s|d   €d| _	        y t        d„ |d   D «       «      | _	        y )Nr   r   é   é   r   c              3   óF   K  — | ]  }t        j                  |«      –— Œ y ­wr   r   r   s     r   r   z/RemovableHandle.__setstate__.<locals>.<genexpr>?   s   è ø€ Ð'IÁ¸1¬¯©°A¯Áùr   )r   r   r   r   r
   Úmaxr   r   Úlenr   r"   )r#   Ústates     r   Ú__setstate__zRemovableHandle.__setstate__3   s�   € Ø�‰8Ðä")§+¡+¬k«mÓ"<ˆDÕä")§+¡+¨e°A©hÓ"7ˆDÔØ˜‘(ˆŒÜ"%¤o×&=Ñ&=¸t¿w¹wÈ¹{Ó"KŒÔäˆu‹:˜Š>˜U 1™XÐ-Ø"$ˆDÕä"'Ñ'IÀÀaÂÓ'IÓ"IˆDÕr%   c                 ó   — | S r   r   r*   s    r   Ú	__enter__zRemovableHandle.__enter__A   s   € Øˆr%   ÚtypeÚvalueÚtbc                 ó$   — | j                  «        y r   )r'   )r#   r6   r7   r8   s       r   Ú__exit__zRemovableHandle.__exit__D   s   € Ø�‰�r%   ©r   N)r   r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚintÚ__annotations__r   r   r$   r'   r+   r3   r5   r:   r   r%   r   r   r   
   sp   … ñð 	ƒGØ€GˆSÓà=Aò 	L 3ð 	L°sð 	LÀdó 	Ló(òaóJóð˜Sð ¨ð °#ð ¸$ô r%   r   c                 ó   — d| _         | S )z»
    Mark a function as an unserializable hook with this decorator.

    This suppresses warnings that would otherwise arise if you attempt
    to serialize a tensor that has a hook.
    T)Ú__torch_unserializable__)Úfs    r   r   r   H   s   € ð "&€AÔØ€Hr%   c                 óÀ   — | j                   rR| j                   D ]B  }| j                   |   }t        |d«      rŒt        j                  dt	        |«      › d�d¬«       ŒD y y )NrC   zbackward hook z› on tensor will not be serialized.  If this is expected, you can decorate the function with @torch.utils.hooks.unserializable_hook to suppress this warningr.   ©Ú
stacklevel)Ú_backward_hooksÚhasattrÚwarningsÚwarnÚrepr)ÚtensorÚkÚhooks      r   r   r   S   sc   € Ø×ÒØ×'Ô'ˆAØ×)Ñ)¨!Ñ,ˆDÜ˜4Ð!;Õ<Ü—‘ ¬t°D«z¨lð ;9ð 9ð FGöHñ (ð r%   c                   ó>   — e Zd ZdZd
d„Zd„ Zd„ Zd
d„Zd„ Zd„ Z	d	„ Z
y)r   a¨  
    A wrapper class to implement nn.Module backward hooks.

    It handles:
      - Ignoring non-Tensor inputs and replacing them by None before calling the user hook
      - Generating the proper Node to capture a set of Tensor's gradients
      - Linking the gradients captures for the outputs with the gradients captured for the input
      - Calling the user hook once both output and input gradients are available
    Nc                 ót   — || _         || _        || _        d | _        d| _        d | _        d| _        d | _        y )Néÿÿÿÿ)Ú
user_hooksÚuser_pre_hooksÚmoduleÚgrad_outputsÚ	n_outputsÚoutput_tensors_indexÚn_inputsÚinput_tensors_index)r#   rU   rS   rT   s       r   r$   zBackwardHook.__init__h   s>   € Ø$ˆŒØ,ˆÔØˆŒà ˆÔØˆŒØ$(ˆÔ!ØˆŒØ#'ˆÕ r%   c                 óZ   — d g|z  }t        ||d¬«      D ]
  \  }}|||<   Œ t        |«      S )NT©Ústrict)Úzipr"   )r#   ÚindicesÚvaluesÚsizeÚresÚidxÚvals          r   Ú_pack_with_nonezBackwardHook._pack_with_nones   s9   € Øˆf�t‰mˆÜ˜G V°D×9‰HˆC�ØˆC�ŠHð :ô �S‹zÐr%   c                 óF   — |D �cg c]  }||   ‘Œ	 }}t        |«      S c c}w r   )r"   )r#   r_   r`   rc   rb   s        r   Ú_unpack_nonezBackwardHook._unpack_nonez   s)   € Ù&-Ó.¡g˜sˆv�c‹{ gˆÐ.ä�S‹zÐùò /s   …c                 ó2   ‡ — ˆ fd„}|j                  |«       y )Nc           	      óŽ  •— ‰j                   €y ‰j                  ‰j                  | ‰j                  «      }‰j                  D ]_  } |‰j
                  |‰j                   «      }|€Œ$t        |«      t        |«      k7  r#t        dt        |«      › dt        |«      › �«      ‚|}Œa d ‰_         ‰j                  ‰j                  |«      S )Nz<Backward hook returned an invalid number of grad_input, got ú, but expected )	rV   re   rZ   rY   rS   rU   r1   ÚRuntimeErrorrg   )Ú
grad_inputÚ_rb   rO   Úoutr#   s        €r   rO   z)BackwardHook._set_user_hook.<locals>.hook€   sÇ   ø€ Ø× Ñ Ð(ð Ø×&Ñ& t×'?Ñ'?ÀÈTÏ]É]Ó[ˆCàŸœ�Ù˜4Ÿ;™;¨¨T×->Ñ->Ó?�à�;Øä�s“8œs 3›xÒ'Ü&ð (.Ü.1°#«h¨Z°ÄsÈ3ÃxÀjð(Ró Sð Sð ‘ð (ð !%ˆDÔà×$Ñ$ T×%=Ñ%=¸sÓCÐCr%   ©Úregister_hook)r#   Úgrad_fnrO   s   `  r   Ú_set_user_hookzBackwardHook._set_user_hook   s   ø€ ô	Dð2 	×Ñ˜dÕ#r%   c                 ó0  — g }g }d}t        |«      D ]Q  \  }}t        |t        j                  «      sŒ!|j	                  |«       |j	                  |«       ||j
                  z  }ŒS |rt        j                  «       s|d fS t        j                  j                  j                  j                  j                  |Ž }t        |«      dk(  rt        d«      ‚|D �	cg c]9  }	|	j                  €Œ|	j                  j                  «       dk(  sŒ.|	j                  ‘Œ; }
}	t        |
«      dk(  rt        d«      ‚ ||
d   «       t!        |«      }t#        ||d¬«      D ]
  \  }}|||<   Œ t%        |«      t&        u rt'        |«      }||fS  t%        |«      |Ž }||fS c c}	w )NFr   zCCannot set Module backward hook for a Module with no input Tensors.ÚBackwardHookFunctionBackwardzaError while setting up backward hooks. Please open an issue with a code sample to reproduce this.Tr\   )Ú	enumerater   ÚtorchÚTensorÚappendÚrequires_gradÚis_grad_enabledÚnnÚmodulesÚ
_functionsÚBackwardHookFunctionÚapplyr1   rk   rq   Únamer!   r^   r6   r"   )r#   ÚfnÚargsÚtensors_idxÚtensorsry   ÚiÚargÚnew_tensorsÚtÚgrad_fnsÚarg_listrc   rd   rn   s                  r   Ú_apply_on_tensorszBackwardHook._apply_on_tensors›   sœ  € ð ˆØˆàˆÜ –o‰FˆAˆsÜ˜#œuŸ|™|Õ,Ø×"Ñ" 1Ô%Ø—‘˜sÔ#Ø ×!2Ñ!2Ñ2‘ð	 &ñ ¤%×"7Ñ"7Ô"9Ø˜�:Ðä—h‘h×&Ñ&×1Ñ1×FÑF×LÑLÈgÐVˆÜˆ{Ó˜qÒ ÜÐdÓeÐeá'2ó  D¡{ !°a·i±iÑ6KÐPQ×PYÑPY×P^ÑP^ÓP`ð  eCó  QC�A—I“I {ˆð  DÜˆx‹=˜AÒÜð  Pó Qð Qñ 	ˆ8�A‰;Œä˜“:ˆÜ˜K¨¸T×B‰HˆC�ØˆH�SŠMð Cô �‹:œÑÜ˜“/ˆCð �KÐÐð ”$�t“*˜hÐ'ˆCØ�KÐÐùò Ds   ÃFÃ)FÄFc                 ól   ‡ — dˆ fd„}‰ j                  ||«      \  }}t        |«      ‰ _        |‰ _        |S )Nc                 ó(   •— ‰j                  | «       y r   )rr   )rq   r#   s    €r   r�   z)BackwardHook.setup_input_hook.<locals>.fnÁ   s   ø€ Ø×Ñ Õ(r%   r;   )r‹   r1   rY   rZ   )r#   r‚   r�   rb   Ú	input_idxs   `    r   Úsetup_input_hookzBackwardHook.setup_input_hookÀ   s8   ø€ õ	)ð ×/Ñ/°°DÓ9‰ˆˆYÜ˜D›	ˆŒØ#,ˆÔ Øˆ
r%   c                 ó¨   ‡ — dˆ fd„}d}t        |t        «      s|f}d}‰ j                  ||«      \  }}t        |«      ‰ _        |‰ _        |s|d   }|S )Nc                 ó2   •— ˆfd„}| j                  |«       y )Nc                 ó&  •‡	— ‰
j                  ‰
j                  |‰
j                  «      ‰
_        ‰
j                  rnt        ‰
j                  «      }‰
j                  D ]J  } |‰
j                  ‰
j                  «      }|€Œ#t        |«      }||k7  rt        d|› d|› �«      ‚|‰
_        ŒL ‰
j                  Š	‰
j                  €št        j                  dd¬«       ‰
j                  g g ‰
j                  «      }‰
j                  D ]P  } |‰
j                  |‰
j                  «      }|€Œ$t        |t        «      rt        d„ |D «       «      rŒGt        d«      ‚ d ‰
_        ‰	�5‰
j                  €t!        d«      ‚t        ˆ	fd	„‰
j                  D «       «      S y )
NzABackward pre hook returned an invalid number of grad_output, got rj   zþFull backward hook is firing when gradients are computed with respect to module outputs since no inputs require gradients. See https://docs.pytorch.org/docs/main/generated/torch.nn.Module.html#torch.nn.Module.register_full_backward_hook for more details.é   rF   c              3   ó$   K  — | ]  }|d u –— Œ
 y ­wr   r   )r   Úels     r   r   zKBackwardHook.setup_output_hook.<locals>.fn.<locals>.hook.<locals>.<genexpr>ë   s   è ø€ ÐRlÑhkÐbdÐSUÐY]ÔS]Ñhkùs   ‚zoBackward hook for Modules where no input requires gradient should always return None or None for all gradients.zEoutput_tensors_index should not be None when grad_outputs is not Nonec              3   ó(   •K  — | ]	  }‰|   –— Œ y ­wr   r   )r   r…   Úlocal_grad_outputss     €r   r   zKBackwardHook.setup_output_hook.<locals>.fn.<locals>.hook.<locals>.<genexpr>ó   s   øè ø€ Ð ZÑ@Y¸1Ð!3°AÕ!6Ñ@Yùs   ƒ)re   rX   rW   rV   rT   r1   rU   rk   rZ   rJ   rK   rY   rS   r   r"   ÚallÚAssertionError)rm   Úgrad_outputÚexpected_lenÚuser_pre_hookÚhook_grad_outputsÚ
actual_lenÚgrad_inputsÚ	user_hookrb   r—   r#   s            @€r   rO   z8BackwardHook.setup_output_hook.<locals>.fn.<locals>.hookË   s�  ù€ Ø$(×$8Ñ$8¸×9RÑ9RØ9DØ9=¿¹ó%I�Ô!ð ×&Ò&Ü#& t×'8Ñ'8Ó#9�LØ)-×)<Ô)<˜Ù,9¸$¿+¹+Àt×GXÑGXÓ,YÐ)Ø,Ð4Ø$ä%(Ð):Ó%;˜
Ø%¨Ò5Ü".ð 06Ø6@°\ÀÐQ]ÐP^ð0`ó #að aà,=˜Õ)ð *=ð &*×%6Ñ%6Ð"ð ×+Ñ+Ð3Ü—M‘Mð #6ð ./õ	0ð
 #'×"6Ñ"6°r¸2¸t¿}¹}Ó"M�KØ%)§_¤_˜	Ù'¨¯©°[À$×BSÑBSÓT˜Ø™?´J¸sÄEÔ4JÌsÑRlÑhkÓRlÕOlÜ".ð 0oó #pð pð &5ð
 )-�DÔ%à%Ð1Ø×0Ñ0Ð8Ü,Ð-tÓuÐuÜ Ó ZÀ×@YÒ@YÓ ZÓZÐZð 2r%   ro   )rq   rO   r#   s     €r   r�   z*BackwardHook.setup_output_hook.<locals>.fnÊ   s   ø€ ô([ðT ×!Ñ! $Õ'r%   TFr   r;   )r   r"   r‹   r1   rW   rX   )r#   r‚   r�   Úis_tuplerb   Ú
output_idxs   `     r   Úsetup_output_hookzBackwardHook.setup_output_hookÉ   s`   ø€ õ+	(ðZ ˆÜ˜$¤Ô&Ø�7ˆDØˆHà×0Ñ0°°TÓ:‰ˆˆZÜ˜T›ˆŒØ$.ˆÔ!áØ�a‘&ˆCØˆ
r%   r;   )r<   r=   r>   r?   r$   re   rg   rr   r‹   r�   r£   r   r%   r   r   r   ]   s+   „ ñó	(òòó
$ò8# òJó9r%   r   r;   )rv   Úcollectionsr   r   rJ   Útypingr   Ú__all__r   r   r   r   r   r%   r   Ú<module>r§      s;   ðã Ý #Û Û Ý â
Y€÷;ñ ;ò|óH÷eò er%   