Ë
      çi¼  ã                   óÂ   — d dl mZ d dlmZmZmZ d dlZd dlmZ d dlm	Z	 d dl
mZ d dlmZmZ d dlmZ d d	lmZ d d
lmZ d dlmZ d dlmZ  G d„ de«      Zdedefd„Zy)é    )ÚAbstractContextManager)ÚAnyÚLiteralÚOptionalN)Úapply_to_collection)ÚTensor)ÚModule)ÚLBFGSÚ	Optimizer)Úoverride)Ú	Precision)Ú_convert_fp_tensor)Ú_TORCH_GREATER_EQUAL_2_4)ÚOptimizablec                   ó:  ‡ — e Zd ZdZ	 dded   deded   ddfd	„Zede	fd
„«       Z
ededefd„«       Zededefd„«       Zededee   dededdf
ˆ fd„«       Zedededefˆ fd„«       Zedeeef   fd„«       Zedeeef   ddfd„«       Zededdfd„«       Zˆ xZS )ÚMixedPrecisionaF  Plugin for Automatic Mixed Precision (AMP) training with ``torch.autocast``.

    Args:
        precision: Whether to use ``torch.float16`` (``'16-mixed'``) or ``torch.bfloat16`` (``'bf16-mixed'``).
        device: The device for ``torch.autocast``.
        scaler: An optional :class:`torch.cuda.amp.GradScaler` to use.

    NÚ	precision©ú16-mixedú
bf16-mixedÚdeviceÚscalerztorch.amp.GradScalerÚreturnc                 óò  — |dvr%t        dt        | «      j                  › d|›d�«      ‚|| _        |€]| j                  dk(  rNt        r t
        j                  j                  |¬«      n't
        j                  j                  j                  «       }|�| j                  dk(  rt        d|› d	�«      ‚|| _	        || _
        | j                  dk(  rt
        j                  | _        y t
        j                  | _        y )
Nr   zPassed `z(precision=z1)`. Precision must be '16-mixed' or 'bf16-mixed'.r   )r   r   z6`precision='bf16-mixed'` does not use a scaler, found Ú.)Ú
ValueErrorÚtypeÚ__name__r   r   ÚtorchÚampÚ
GradScalerÚcudar   r   Úbfloat16Úfloat16Ú_desired_input_dtype)Úselfr   r   r   s       ú{/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/lightning_fabric/plugins/precision/amp.pyÚ__init__zMixedPrecision.__init__(   så   € ð Ð6Ñ6ÜØœ4 ›:×.Ñ.Ð/¨{¸9¸-ð HAð Aóð ð
 #ˆŒØˆ>˜dŸn™n°
Ò:Ý<T”U—Y‘Y×)Ñ)°Ð)Ô8ÔZ_×ZdÑZd×ZhÑZh×ZsÑZsÓZuˆFØÐ $§.¡.°LÒ"@ÜÐUÐV\ÐU]Ð]^Ð_Ó`Ð`ØˆŒØˆŒà6:·n±nÈÒ6T¤E§N¡NˆÕ!ÔZ_×ZgÑZgˆÕ!ó    c                 óX   — t        j                  | j                  | j                  ¬«      S )N)Údtype)r   Úautocastr   r%   ©r&   s    r'   Úforward_contextzMixedPrecision.forward_context>   s   € ä�~‰~˜dŸk™k°×1JÑ1JÔKÐKr)   Údatac                 óD   — t        |t        t        | j                  ¬«      S ©N)Úfunctionr+   Údst_type)r   r   r   r%   ©r&   r/   s     r'   Úconvert_inputzMixedPrecision.convert_inputB   s   € ä" 4Ô2DÌFÐ]a×]vÑ]vÔwÐwr)   c                 óT   — t        |t        t        t        j                  «       ¬«      S r1   )r   r   r   r   Úget_default_dtyper4   s     r'   Úconvert_outputzMixedPrecision.convert_outputF   s   € ä" 4Ô2DÌFÔ]b×]tÑ]tÓ]vÔwÐwr)   ÚtensorÚmodelÚargsÚkwargsc                 ó|   •— | j                   �| j                   j                  |«      }t        ‰| �  ||g|¢­i |¤Ž y ©N)r   ÚscaleÚsuperÚbackward)r&   r9   r:   r;   r<   Ú	__class__s        €r'   rA   zMixedPrecision.backwardJ   s:   ø€ à�;‰;Ð"Ø—[‘[×&Ñ& vÓ.ˆFÜ‰Ñ˜ Ð8¨Ò8°Ó8r)   Ú	optimizerc                 óä   •— | j                   €t        ‰| �  |fi |¤ŽS t        |t        «      rt        d«      ‚ | j                   j                  |fi |¤Ž}| j                   j                  «        |S )Nz/AMP and the LBFGS optimizer are not compatible.)r   r@   Úoptimizer_stepÚ
isinstancer
   Ú	TypeErrorÚstepÚupdate)r&   rC   r<   Ústep_outputrB   s       €r'   rE   zMixedPrecision.optimizer_stepP   sk   ø€ ð �;‰;Ðä‘7Ñ)¨)Ñ>°vÑ>Ð>Ü�i¤Ô'ÜÐMÓNÐNà&�d—k‘k×&Ñ& yÑ;°FÑ;ˆØ�‰×ÑÔØÐr)   c                 óR   — | j                   �| j                   j                  «       S i S r>   )r   Ú
state_dictr-   s    r'   rL   zMixedPrecision.state_dict`   s$   € à�;‰;Ð"Ø—;‘;×)Ñ)Ó+Ð+Øˆ	r)   rL   c                 óT   — | j                   �| j                   j                  |«       y y r>   )r   Úload_state_dict)r&   rL   s     r'   rN   zMixedPrecision.load_state_dictf   s#   € à�;‰;Ð"Ø�K‰K×'Ñ'¨
Õ3ð #r)   c                 óp   — | j                   }|�(t        |«      rt        d«      ‚|j                  |«       y y )NzKGradient clipping is not implemented for optimizers handling the unscaling.)r   Ú_optimizer_handles_unscalingÚNotImplementedErrorÚunscale_)r&   rC   r   s      r'   Úunscale_gradientsz MixedPrecision.unscale_gradientsk   s6   € à—‘ˆØÐÜ+¨IÔ6Ü)Ð*wÓxÐxØ�O‰O˜IÕ&ð r)   r>   )r   Ú
__module__Ú__qualname__Ú__doc__r   Ústrr   r(   r   r   r.   r   r5   r8   r   r	   rA   r   rE   ÚdictrL   rN   r   rS   Ú__classcell__)rB   s   @r'   r   r      s‘  ø„ ñð 48ñ	hàÐ3Ñ4ðhð ðhð Ð/Ñ0ð	hð
 
óhð, ðLÐ!7ò Ló ðLð ðx #ð x¨#ò xó ðxð ðx 3ð x¨3ò xó ðxð ð9˜vð 9¨h°vÑ.>ð 9Àsð 9ÐVYð 9Ð^bô 9ó ð9ð
 ðàðð ðð 
ô	ó ðð ð˜D  c ™Nò ó ðð
 ð4¨$¨s°C¨x©.ð 4¸Tò 4ó ð4ð ð'¨9ð '¸ò 'ó ô'r)   r   rC   r   c                 ó   — t        | dd«      S )aT  Determines whether a PyTorch optimizer handles unscaling gradients in the step method rather than through the
    :class:`torch.cuda.amp.GradScaler`.

    Since, the current implementation of this function checks a PyTorch internal variable on the optimizer, the return
    value will only be reliable for built-in PyTorch optimizers.

    Ú_step_supports_amp_scalingF)Úgetattr)rC   s    r'   rP   rP   t   s   € ô �9Ð:¸EÓBÐBr)   )Ú
contextlibr   Útypingr   r   r   r   Ú#lightning_utilities.core.apply_funcr   r   Útorch.nnr	   Útorch.optimr
   r   Útyping_extensionsr   Ú,lightning_fabric.plugins.precision.precisionr   Ú(lightning_fabric.plugins.precision.utilsr   Ú"lightning_fabric.utilities.importsr   Ú lightning_fabric.utilities.typesr   r   ÚboolrP   © r)   r'   Ú<module>ri      sS   ðõ .ß )Ñ )ã Ý CÝ Ý ß (Ý &å BÝ GÝ GÝ 8ôS'�Yô S'ðlC¨Cð C°Dô Cr)   