
      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 
 #>dnn
:<TUYY)))8Z_ZdZdZhZhZsZsZuF$..L"@UV\U]]^_``6:nn6TENN!Z_ZgZg!    c                 X    t        j                  | j                  | j                        S )N)dtype)r   autocastr   r%   r&   s    r'   forward_contextzMixedPrecision.forward_context>   s    ~~dkk1J1J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    "42DF]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    "42DF]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&dkk&&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#    ;;"KK''
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OOI& 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 9v 9hv.> 9s 9VY 9^b 9 9
   
	  DcN  
 4$sCx. 4T 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)   