
      i                         d dl mZ d dlmZ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 d dlmZ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 erd dlmZ d dl m!Z! ed   Z" G d de      Z#y)    )AbstractContextManager)TYPE_CHECKINGAnyLiteralOptionalN)apply_to_collection)Tensor)Module)	Optimizer)get_argsoverride)_optimizer_handles_unscaling)	Precision)_convert_fp_tensor_DtypeContextManager)rank_zero_warn)OptimizableMixedPrecisionShardedGradScaler)32-true16-true	bf16-true16-mixed
bf16-mixedc                       e Zd ZdZddeded   ddfdZededefd	       Z	e
dd
       Zedefd       Zede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ddfd       Zedeeef   fd       Zedeeef   ddfd       Z xZS )FSDPPrecisiona  Precision plugin for training with Fully Sharded Data Parallel (FSDP).

    .. warning::  This is an :ref:`experimental <versioning:Experimental API>` feature.

    Args:
        precision: Full precision (32-true), half precision (16-true, bf16-true) or
            mixed precision (16-mixed, bf16-mixed).
        scaler: An optional :class:`torch.distributed.fsdp.sharded_grad_scaler.ShardedGradScaler` to use.

    Raises:
        ValueError:
            If unsupported ``precision`` is provided.

    N	precisionscalerr   returnc                    t        t              }||vrt        d|d| d      ddlm} |!| j
                  dk7  rt        d|d| d      ||dk(  r |       nd | _        || _        t        j                  t        j                  t        j                  t        j                  t        j                  d}|| j
                     | _        y )	Nz`precision=z9)` is not supported in FSDP. `precision` must be one of: .r   r   r   z` does not use a scaler, found )r   r   r   r   r   )r   _PRECISION_INPUT
ValueError*torch.distributed.fsdp.sharded_grad_scalerr   r   r    torchfloat32bfloat16float16_desired_input_dtype)selfr   r    supported_precisionr   precision_to_types         |/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/lightning/fabric/plugins/precision/fsdp.py__init__zFSDPPrecision.__init__5   s    &'78//i] +00C/DAG 
 	Q$..J">{9-7VW]V^^_`aa-3^	Z@W')]a"  --}}}}
 %6dnn$E!    modulec                 Z    d| j                   v r|j                  | j                        S |S )Ntruedtype)r   tor+   )r,   r2   s     r/   convert_modulezFSDPPrecision.convert_moduleN   s*    T^^#994#<#<9==r1   c                 x   ddl m} | j                  dv rt        d| j                   d       | j                  dv rt        j
                  x}x}}n`| j                  dv rt        j                  x}x}}n=| j                  dk(  rt        j                  x}x}}nt        d	| j                  d
       ||||      S )Nr   r   )r   r   zFSDP with `z}` enables computation in lower precision. FSDP will always retain a full-precision copy of the model parameters for sharding.)r   r   )r   r   r   z-Was unable to infer precision type, received r#   )param_dtypereduce_dtypebuffer_dtype)	2torch.distributed.fsdp.fully_sharded_data_parallelr   r   r   r'   r*   r)   r(   r%   )r,   TorchMixedPrecisionr:   r;   r<   s        r/   mixed_precision_configz$FSDPPrecision.mixed_precision_configT   s    l>>55dnn- .f f
 >>448=EKE,^^::8=FKF,^^y(8=EKE,LT^^L^^_`aa"#%%
 	
r1   c                 ,    t        | j                        S N)r   r+   r,   s    r/   tensor_init_contextz!FSDPPrecision.tensor_init_contextm   s    #D$=$=>>r1   c                 d    t        | j                  j                  xs t        j                        S rA   )r   r?   r:   r'   r(   rB   s    r/   module_init_contextz!FSDPPrecision.module_init_contextq   s"    #D$?$?$K$K$\u}}]]r1   c                     d| j                   v rIt        j                  d| j                   dk(  rt        j                        S t        j                        S | j                         S )Nmixedcudar   r5   )r   r'   autocastr)   r*   rC   rB   s    r/   forward_contextzFSDPPrecision.forward_contextu   sM    dnn$>>&4>>UaCavvglgtgtvv''))r1   datac                 D    t        |t        t        | j                        S N)functionr6   dst_type)r   r   r	   r+   r,   rK   s     r/   convert_inputzFSDPPrecision.convert_input{   s    "42DF]a]v]vwwr1   c                 T    t        |t        t        t        j                               S rM   )r   r   r	   r'   get_default_dtyperP   s     r/   convert_outputzFSDPPrecision.convert_output   s    "42DF]b]t]t]vwwr1   tensormodelargskwargsc                 |    | j                   | j                   j                  |      }t        |   ||g|i | y rA   )r    scalesuperbackward)r,   rU   rV   rW   rX   	__class__s        r/   r\   zFSDPPrecision.backward   s:    ;;"[[&&v.F888r1   	optimizerc                     | j                   t        |   |fi |S  | j                   j                  |fi |}| j                   j	                          |S rA   )r    r[   optimizer_stepstepupdate)r,   r^   rX   step_outputr]   s       r/   r`   zFSDPPrecision.optimizer_step   sU     ;;7))>v>>&dkk&&y;F;r1   c                 p    | j                   }|(t        |      rt        d      |j                  |       y y )NzKGradient clipping is not implemented for optimizers handling the unscaling.)r    r   NotImplementedErrorunscale_)r,   r^   r    s      r/   unscale_gradientszFSDPPrecision.unscale_gradients   s6    +I6)*wxxOOI& r1   c                 R    | j                   | j                   j                         S i S rA   )r    
state_dictrB   s    r/   ri   zFSDPPrecision.state_dict   s$    ;;";;))++	r1   ri   c                 T    | j                   | j                   j                  |       y y rA   )r    load_state_dict)r,   ri   s     r/   rk   zFSDPPrecision.load_state_dict   s#    ;;"KK''
3 #r1   rA   )r!   r>   )__name__
__module____qualname____doc__r$   r   r0   r   r
   r8   propertyr?   r   rC   rE   rJ   r   rQ   rT   r	   r\   r   r`   r   rg   dictstrri   rk   __classcell__)r]   s   @r/   r   r   %   s   F"2 FHEX<Y Fei F2 V   
 
 
0 ?%; ? ? ^%; ^ ^ *!7 * *
 x# x# x x x3 x3 x x 9v 9hv.> 9s 9VY 9^b 9 9
   
	  '9 ' ' ' DcN  
 4$sCx. 4T 4 4r1   r   )$
contextlibr   typingr   r   r   r   r'   lightning_utilitiesr   r	   torch.nnr
   torch.optimr   typing_extensionsr   r   &lightning.fabric.plugins.precision.ampr   ,lightning.fabric.plugins.precision.precisionr   (lightning.fabric.plugins.precision.utilsr   r   lightning.fabric.utilitiesr    lightning.fabric.utilities.typesr   r=   r   r>   r&   r   r$   r    r1   r/   <module>r      sT    . 8 8  3   ! 0 O B ] 5 8hLVW C4I C4r1   