
      i                     
   d dl Z 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
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ddfdZde
deddfdZde
dej2                  ddfdZddee
ef   dedefdZy)    N)Sequence)AnyCallableOptionalUnion)Module	Parameter)	Optimizer)TorchFunctionMode)override)rank_zero_warn)_DEVICEc                   f     e Zd ZdZddeddf fdZe	 	 ddededee	   d	e
e   de	f
d
       Z xZS )
_EmptyInitzInitialize `nn.Module` with empty tensors, i.e., uninitialized memory.

    Example::

        with _EmptyInit():
            model = BigModel()
        model.load_state_dict(torch.load("checkpoint.pt"))

    enabledreturnNc                 0    t         |           || _        y N)super__init__r   )selfr   	__class__s     t/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/lightning_fabric/utilities/init.pyr   z_EmptyInit.__init__(   s        functypesargskwargsc                     |xs i }| j                   s ||i |S t        |dd       dk(  rd|v r|d   S |d   S  ||i |S )N
__module__ztorch.nn.inittensorr   )r   getattr)r   r   r   r   r   s        r   __torch_function__z_EmptyInit.__torch_function__,   sa     2||(((4t,?6!h''7NT$V$$r   T) N)__name__r    __qualname____doc__boolr   r   r   r   r   r   dictr#   __classcell__)r   s   @r   r   r      so       
 !!%%% % sm	%
 % 
% %r   r   moduledevicer   c                     | j                  |d       t        | d      s"t        dt        |       j                   d      t        | j                        r| j                          yy)zMaterialize a module.F)r-   recursereset_parametersz#Materialization requires that the `z~.reset_parameters` method is implemented. This method is used to initialize any children parameters or buffers in this module.N)to_emptyhasattr	TypeErrortyper&   callabler0   r,   r-   s     r   _materializer7   >   sj    
OO65O16-.1$v,2G2G1H Id d
 	
 ''(! )r   c                 b    | j                         D ]  } t        | d      st        | |        y)z*Materialize all tensors in a given module.Fr/   N)modules&_has_meta_device_parameters_or_buffersr7   r6   s     r   _materialize_meta_tensorsr<   J   s'    .."1&%H( #r   c           
         t        |       sy | j                  |       t               }| j                         D ]  }t	        d t        j                  |j                  d      |j                  d            D              rJt        t        |dd       x}      r |        j|j                  t        |      j                          |rt        ddj                  |              y y )N)r-   c              3       K   | ]  }d   yw)FNr%   ).0_s     r   	<genexpr>z2_materialize_distributed_module.<locals>.<genexpr>\   s     ututs   Fr9   r0   zParameter initialization incomplete. The following modules have parameters or buffers with uninitialized memory because they don't define a `reset_parameters()` method for re-initialization: z, )r;   r1   setr:   all	itertoolschain
parametersbuffersr5   r"   addr4   r&   r   join)r,   r-   uninitialized_modules	submodulereset_methods        r   _materialize_distributed_modulerM   Q   s     2&9
OO6O"E^^%	uiooi.B.B5.B.QS\SdSdmrSdSstuuGI7I4$PPLQN!%%d9o&>&>? & 		/013	
 r   objr/   c           	      H   t        | t              rt        d | j                  D              S t        | t              rFt        d t        j                  | j                  |      | j                  |            D              S t        dt        |       j                         )Nc              3   j   K   | ]+  }|d    D ]!  }t        |t              s|j                   # - yw)paramsN)
isinstancer	   is_meta)r?   param_groupts      r   rA   z9_has_meta_device_parameters_or_buffers.<locals>.<genexpr>n   s3      
)9+;xCXa\fghjs\tAIICXI)9s   33c              3   4   K   | ]  }|j                     y wr   )rS   )r?   rU   s     r   rA   z9_has_meta_device_parameters_or_buffers.<locals>.<genexpr>r   s     u&t199&ts   r9   z<Expected `torch.nn.Module` or `torch.optim.Optimizer`, got: )rR   r
   anyparam_groupsr   rD   rE   rF   rG   r3   r4   r&   )rN   r/   s     r   r;   r;   l   s    #y! 
),)9)9
 
 	
 #vuioocnnWn6UWZWbWbkrWbWs&tuuu
RSWX[S\SeSeRfg
hhr   r$   )rD   collections.abcr   typingr   r   r   r   torchtorch.nnr   r	   torch.optimr
   torch.overridesr   typing_extensionsr   $lightning_fabric.utilities.rank_zeror    lightning_fabric.utilities.typesr   r   r7   r<   r-   rM   r)   r;   r%   r   r   <module>rb      s     $ 1 1  & ! - & ? 4%" %B	" 	" 	"T 	")f )g )$ )
F 
ELL 
T 
6ifi6G0H iSW icg ir   