§
    ŠŠtjo7  ã                   óü   — U d dl Z d dlZd dlmZmZmZ d dlmZ d dl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 g Zee         ed<    e j        e¦  «        Z G d	„ d
ej        ¦  «        Zdee         defd„ZdS )é    N)ÚCallableÚ
CollectionÚMapping)Údeepcopy)ÚAnyÚoverload)Úoptim)ÚShardedTensor)ÚFullyShardedDataParallelÚ__all__c                   ó2  — e Zd ZdZ	 	 ddeeej        ez  f         de	j
        deeeef                  dz  dej        dz  deedf         d	eeef         d
dfd„Zdd„Zd
eeef         fd„Zeddd„¦   «         Zedeg ef         d
efd„¦   «         Zddeg ef         dz  d
edz  fd„Zed
eej        ef         fd„¦   «         Zdeeef         d
dfd„Zdeeef         d
dfd„Zdd„Zdeeef         d
eeef         fd„Zdeeef         d
eeef         fd„ZdS )Ú_NamedOptimizerað  
    ``_NamedOptimizer`` takes a dict of parameters and exposes ``state_dict`` by parameter key.

    We replace the original key (number) in an optim to the
    fully qualified name (FQN) string. User can initialize the optim as they
    initialize a PyTorch optim, the only difference is that they also need to
    pass in the FQN of each parameters.

    Args:
        named_parameters (Mapping[str, Union[torch.Tensor, ShardedTensor]]):
            Mapping from FQN to parameter.
        optimizer_class (optim.Optimizer):
            The class of optimizer to instantiate.
        param_groups (Collection[Mapping[str, Any]]):
            `param_groups` to pass to optimizer if specified.
            The key of the inner map needs to be FQNs.
            Default: None
        module (nn.Module): the module whose parameters to be updated
            by the optimizer.
        args: arguments to pass to the optimizer constructor.
        kwargs: arguments to pass to the optimizer constructor.

    Example::
        >>> # xdoctest: +SKIP("distributed")
        >>> from torch import optim
        >>> from torch.distributed.optim import _NamedOptimizer
        >>>
        >>> # Define the named optimizer.
        >>> m = Model(...)
        >>> named_optim = _NamedOptimizer(m.named_parameters(), optim.SGD)
        >>> # Forward pass + backward pass.
        >>> named_optim.step()
        >>> ...
        >>> # Call state_dict for the named optimizer returns an FQN state_dict.
        >>> named_optim.state_dict()

    Warning: This API is still in development and subject to change.

    TODO: Add tutorial for _NamedOptimizer.
    TODO: Add documentation in the docstring for the public attributes
          like self.param_groups and self.named_parameters.
    NÚnamed_parametersÚoptimizer_classÚparam_groupsÚmoduleÚargs.ÚkwargsÚreturnc                 ó’  — t           j                             d¦  «         || _        |                      ¦   «          t          |¦  «        | _        |€| j                             ¦   «         n|} ||g|¢R i |¤Ž| _        || _	        |€,t          | j                             ¦   «         ¦  «        | _        n„t          j        dd¬¦  «         d„ | j                             ¦   «         D ¦   «         }g }	|D ]?}
|
d         D ]4}||vrt!          d|› d�¦  «        ‚|	                     ||         ¦  «         Œ5Œ@|	| _        | j        j        | _        d S )	Nz'torch.distributed.optim._NamedOptimizerzvSince we pass in param_groups, we will use param_groups to initialize the optimizer, not all parameters of the module.é   )Ú
stacklevelc                 ó   — i | ]\  }}||“Œ	S © r   ©Ú.0ÚkeyÚparams      úe/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/distributed/optim/named_optimizer.pyú
<dictcomp>z,_NamedOptimizer.__init__.<locals>.<dictcomp>]   s   € ÐWÐWÐW©:¨3°˜E 3ÐWÐWÐWó    ÚparamszExpect param name z% found in param group but is missing.)ÚtorchÚ_CÚ_log_api_usage_oncer   Ú_param_groups_checkÚdictr   ÚvaluesÚ
_optimizerr   ÚlistÚkeysÚordered_param_keysÚwarningsÚwarnÚitemsÚ
ValueErrorÚappend)Úselfr   r   r   r   r   r   Úparams_for_optimizerÚparam_to_keyr,   Úgroupr   s               r   Ú__init__z_NamedOptimizer.__init__?   s«  € õ 	Œ×$Ò$Ð%NÑOÔOÐOØ;GˆÔØ× Ò Ñ"Ô"Ð"Ý $Ð%5Ñ 6Ô 6ˆÔà.:Ð.BˆDÔ!×(Ò(Ñ*Ô*Ð*Èð 	ð *˜/Ø ð
àð
ð 
ð 
ð ð
ð 
ˆŒð
 ˆŒØÐÝ&*¨4Ô+@×+EÒ+EÑ+GÔ+GÑ&HÔ&HˆDÔ#Ð#åŒMðNàðñ ô ð ð
 XÐW¸Ô9N×9TÒ9TÑ9VÔ9VÐWÑWÔWˆLØ!#ÐØ%ð Cð C�Ø" 8œ_ð Cð C�EØ LÐ0Ð0Ý(Ø]°Ð]Ð]Ð]ñô ð ð '×-Ò-¨l¸5Ô.AÑBÔBÐBÐBðCð '9ˆDÔ#à œOÔ8ˆÔÐÐr!   c                 ó’  — | j         �½| j         D ]·}t          |t          ¦  «        st          d¦  «        ‚d|vrt          d¦  «        ‚|d         }t          |t          j        ¦  «        r|g}t          |¦  «        }|D ]@}t          |t          j        ¦  «        s$t          dt	          j        |¦  «        z   ¦  «        ‚ŒA||d<   Œ¶d S d S )Núparam group must be a dictr"   z#param group must contain key paramsz>optimizer can only optimize Tensors, but one of the params is )	r   Ú
isinstancer'   ÚAssertionErrorr#   ÚTensorr*   Ú	TypeErrorÚtypename)r2   Úparam_groupr"   r   s       r   r&   z#_NamedOptimizer._param_groups_checkj   sü   € ØÔÐ(Ø#Ô0ð /ð /�Ý! +­tÑ4Ô4ð GÝ(Ð)EÑFÔFÐFØ ;Ð.Ð.Ý(Ð)NÑOÔOÐOØ$ XÔ.�Ý˜f¥e¤lÑ3Ô3ð &Ø$˜X�FÝ˜f™œ�Ø#ð ð �EÝ% e­U¬\Ñ:Ô:ð Ý'ð8Ý:?¼.ÈÑ:OÔ:OñPñô ð ðð
 )/�˜HÑ%Ð%ð! )Ð(ð/ð /r!   c                 ó¨  ‡ — ‰ j                              ¦   «         }|d         }ˆ fd„|d                              ¦   «         D ¦   «         }g }|D ]n}ˆ fd„|d         D ¦   «         }dt          |¦  «        i}|                     ¦   «         D ]\  }}	|dk    rt	          |	¦  «        ||<   Œ|                     |¦  «         Œo‰                      ||dœ¦  «        S )zµ
        Return the ``state_dict`` of the optimizer.

        Instead of using number to index
        parameters, we will use module fully qualified name (FQN) as the key.
        r   c                 ó2   •— i | ]\  }}‰j         |         |“ŒS r   ©r,   )r   Úst_keyÚ	state_valr2   s      €r   r    z._NamedOptimizer.state_dict.<locals>.<dictcomp>‡   s7   ø€ ð 
ð 
ð 
á!�˜	ð Ô# FÔ+¨Yð
ð 
ð 
r!   Ústatec                 ó*   •— g | ]}‰j         |         ‘ŒS r   rA   )r   r   r2   s     €r   ú
<listcomp>z._NamedOptimizer.state_dict.<locals>.<listcomp>Ž   s!   ø€ ÐVÐVÐV¸U˜$Ô1°%Ô8ÐVÐVÐVr!   r"   )rD   r   )r)   Ú
state_dictr/   Úsortedr   r1   Ú_post_state_dict)
r2   rG   r   Ú	ret_stateÚ
ret_groupsr5   Ú
param_keysÚ	ret_groupÚkÚvs
   `         r   rG   z_NamedOptimizer.state_dict}   s  ø€ ð ”_×/Ò/Ñ1Ô1ˆ
Ø! .Ô1ˆð
ð 
ð 
ð 
à%/°Ô%8×%>Ò%>Ñ%@Ô%@ð
ñ 
ô 
ˆ	ð
 ˆ
Ø!ð 	)ð 	)ˆEØVÐVÐVÐVÀeÈHÄoÐVÑVÔVˆJØ!¥6¨*Ñ#5Ô#5Ð6ˆIØŸš™œð /ð /‘��1Ø˜’=�=Ý#+¨A¡;¤;�I˜a‘LøØ×Ò˜iÑ(Ô(Ð(Ð(à×$Ò$¨yÈ*Ð%UÐ%UÑVÔVÐVr!   Úclosurec                 ó   — d S ©Nr   ©r2   rP   s     r   Ústepz_NamedOptimizer.step—   s   € Ø25°#r!   c                 ó   — d S rR   r   rS   s     r   rT   z_NamedOptimizer.stepš   s   € Ø;>¸3r!   c                 ó8   — | j                              |¬¦  «        S )z’
        Perform a single optimization step.

        This will call :meth:`torch.optim.Optimizer.step` on the wrapped
        optimizer.
        ©rP   )r)   rT   rS   s     r   rT   z_NamedOptimizer.step�   s   € ð Œ×#Ò#¨GÐ#Ñ4Ô4Ð4r!   c                 ó   — | j         j        S rR   )r)   rD   )r2   s    r   rD   z_NamedOptimizer.state¦   s   € àŒÔ$Ð$r!   rG   c                 ó   — | j                              ¦   «         }|                      |¦  «        }|d         }|d         }t          |¦  «        dk    rt	          d¦  «        ‚t          | j        ¦  «        D �]B\  }}||vrŒt          ||         ¦  «        t          ||         ¦  «        k    r>t	          dt          ||         ¦  «        › d|› dt          ||         ¦  «        › �¦  «        ‚||                              ¦   «         D �]±\  }}|||         vrt	          d|› d|› d�¦  «        ‚||         |         }	t          |t          ¦  «        rìt          |	t          ¦  «        st          ‚t          |                     ¦   «         ¦  «        }
t          |	                     ¦   «         ¦  «        }|
|k    rt	          d	|› d
|
› d|› d|› �¦  «        ‚t          |                     ¦   «         |	                     ¦   «         ¦  «        D ]6\  }}|j                             ¦   «                              |j        ¦  «         Œ7�Œ5t          |t           j        ¦  «        rJt          |	t           j        ¦  «        st          ‚|                     ¦   «                              |	¦  «         �Œ™t%          |	¦  «        ||         |<   �Œ³�ŒD|d         }|d         }i }|D ])}t'          |d         ¦  «        }||t)          |¦  «        <   Œ*i }|D ]A}g }|d         D ]"}|                     | j        |         ¦  «         Œ#||t)          |¦  «        <   ŒB|                     ¦   «         D ]¢\  }}||vrŒ
||         }t          |¦  «        t          |¦  «        k    r3t	          dt          |¦  «        › d|› d
t          |¦  «        › d�¦  «        ‚|D ]:}||vrt	          d|› d|› d�¦  «        ‚|dk    rt%          ||         ¦  «        ||<   Œ;Œ£| j                              |¦  «         dS )aè  
        Define the default behavior to load a state_dict for ``_NamedOptimizer``.

        Sample Code
        ```
            my_model = MyModule()
            optimizer = _NamedOptimizer(my_model.named_parameters(), Adagrad)
            ...

            optim_state_dict = optimizer.state_dict()
            ...
            ...

            optimizer.load_state_dict(optim_state_dict)
            ...
        ```
        Args:
            state_dict (dict[str, Any]) : A ``state_dict`` to load into the optimizer.
                Note that this state dict update is performed in place.

        .. note:: PyTorch is using lazy init to initialize the optim states.
            So it is possible that there is no optim state when user call
            ``load_state_dict`` and for ``_NamedOptimizer`` we make it stricter
            that users can only call ``load_state_dict`` after the state is initialized.
            By doing this, we can validate the optim ``state_dict`` to be loaded.
        rD   r   zJExpects the optim to be initialized before load but found not initialized.zExpects equal length as z for parameter z but found: zExpects state z but not found.z"Expects equal number of shards as z but found z for ú/r   r"   z"Expects equal param_group size as z for group ú.zExpects group key z to be in group z  in `state_dict` but is missing.N)r)   rG   Ú_pre_load_state_dictÚlenr0   Ú	enumerater,   r/   r9   r
   r:   Úlocal_shardsÚzipÚtensorÚdetachÚcopy_r#   r;   r   r*   Ú_gen_param_group_keyr1   Úload_state_dict)r2   rG   Únew_state_dictrD   Ú	new_stateÚidxÚ	param_keyÚ	state_keyrC   Úsrc_state_valÚ
num_shardsÚnum_new_shardsÚshardÚ	src_shardÚsrc_param_groupsÚnew_param_groupsÚsrc_group_mapr5   rL   Únew_group_mapÚ	new_groupÚ	group_keyÚ	src_grouprN   s                           r   re   z_NamedOptimizer.load_state_dictª   sþ  € ð6 œ×3Ò3Ñ5Ô5ˆØ×.Ò.¨zÑ:Ô:ˆ
Ø˜7Ô#ˆØ" 7Ô+ˆ	Ýˆy‰>Œ>˜QÒÐÝØ\ñô ð õ (¨Ô(?Ñ@Ô@ð "	Hñ "	H‰NˆC�à Ð%Ð%ØÝ�5˜Ô#Ñ$Ô$­¨I°c¬NÑ(;Ô(;Ò;Ð;Ý ð B­s°9¸S´>Ñ/BÔ/Bð  Bð  BÐS\ð  Bð  BÕjmÐnsÐt}Ôn~ÑjÔjð  Bð  Bñô ð ð )2°#¬×(<Ò(<Ñ(>Ô(>ð Hñ HÑ$�	˜9Ø E¨)Ô$4Ð4Ð4Ý$Ø]¨Ð]Ð]À9Ð]Ð]Ð]ñô ð ð !& iÔ 0°Ô ;�Ý˜i­Ñ7Ô7ð HÝ% mµ]ÑCÔCð -Ý,Ð,Ý!$ Y×%;Ò%;Ñ%=Ô%=Ñ!>Ô!>�JÝ%(¨×)CÒ)CÑ)EÔ)EÑ%FÔ%F�NØ! ^Ò3Ð3Ý(ð EÀð  Eð  EÐ\fð  Eð  EÐmvð  Eð  Eð  zCð  Eð  Eñô ð õ -0Ø!×.Ò.Ñ0Ô0°-×2LÒ2LÑ2NÔ2Nñ-ô -ð Fð FÑ(˜˜yð œ×+Ò+Ñ-Ô-×3Ò3°IÔ4DÑEÔEÐEÐEñFõ   	­5¬<Ñ8Ô8ð HÝ% mµU´\ÑBÔBð -Ý,Ð,Ø×$Ò$Ñ&Ô&×,Ò,¨]Ñ;Ô;Ð;Ñ;å08¸Ñ0GÔ0G�I˜c”N 9Ñ-Ñ-ñ3Hð8 & nÔ5ÐØ)¨.Ô9ÐàˆØ%ð 	Dð 	DˆEÝ˜e HœoÑ.Ô.ˆJØ>CˆMÕ.¨zÑ:Ô:Ñ;Ð;ØˆØ)ð 	Hð 	HˆIØˆJØ& xÔ0ð Fð F�	Ø×!Ò! $Ô"9¸)Ô"DÑEÔEÐEÐEØ>GˆMÕ.¨zÑ:Ô:Ñ;Ð;Ø$1×$7Ò$7Ñ$9Ô$9ð 	:ð 	:Ñ ˆI�yð  Ð-Ð-ØØ% iÔ0ˆIÝ�9‰~Œ~¥ Y¡¤Ò/Ð/Ý Ø{½¸Y¹¼Ð{Ð{ÐT]Ð{Ð{ÕjmÐnwÑjxÔjxÐ{Ð{Ð{ñô ð ð ð :ð :�Ø˜IÐ%Ð%Ý$Øk¨QÐkÐkÀ	ÐkÐkÐkñô ð ð ˜’=�=Ý#+¨I°a¬LÑ#9Ô#9�I˜a‘Løð:ð 	Œ×'Ò'¨Ñ7Ô7Ð7Ð7Ð7r!   r>   c                 óÜ  — t          |t          ¦  «        st          d¦  «        ‚|d         }t          |t          j        ¦  «        r|g|d<   nt          |¦  «        |d<   d„ | j                             ¦   «         D ¦   «         }|d         D ]5}||vrt          d¦  «        ‚| j	         
                    ||         ¦  «         Œ6| j                             |¦  «         | j        j        | _        dS )zŸ
        Add a param group to the :class:`_NamedOptimizer` s `param_groups`.

        Warning: This API is still in development and subject to change.
        r8   r"   c                 ó   — i | ]\  }}||“Œ	S r   r   r   s      r   r    z3_NamedOptimizer.add_param_group.<locals>.<dictcomp>#  s   € ÐSÐSÐS¡z s¨E˜˜sÐSÐSÐSr!   z%some parameters are not in the moduleN)r9   r'   r:   r#   r;   r*   r   r/   r0   r,   r1   r)   Úadd_param_groupr   )r2   r>   r"   r4   r   s        r   ry   z_NamedOptimizer.add_param_group  sü   € õ ˜+¥tÑ,Ô,ð 	?Ý Ð!=Ñ>Ô>Ð>à˜XÔ&ˆÝ�f�eœlÑ+Ô+ð 	1Ø%+ HˆK˜Ñ!Ð!å$(¨¡L¤LˆK˜Ñ!àSÐS°TÔ5J×5PÒ5PÑ5RÔ5RÐSÑSÔSˆØ  Ô*ð 	@ð 	@ˆEØ˜LÐ(Ð(Ý Ð!HÑIÔIÐIØÔ#×*Ò*¨<¸Ô+>Ñ?Ô?Ð?Ð?àŒ×'Ò'¨Ñ4Ô4Ð4à œOÔ8ˆÔÐÐr!   c                 óè   — | j                              ¦   «         D ]A}|j        r8t          j        |¦  «        }t          j                             |¦  «        |_        ŒB|                      d¬¦  «         dS )zÚ
        Run a dummy optimizer step, which allows us to initialize optimizer state because we do lazy init for most optimizers.

        This allows doing in-place loading of optimizer state from a checkpoint.
        NrW   )	r   r(   Úrequires_gradr#   Ú
zeros_likeÚautogradÚVariableÚgradrT   )r2   r   Úts      r   Ú
init_statez_NamedOptimizer.init_state-  so   € ð Ô*×1Ò1Ñ3Ô3ð 	8ð 	8ˆEØÔ"ð 8ÝÔ$ UÑ+Ô+�Ý"œ^×4Ò4°QÑ7Ô7�”
øà�	Š	˜$ˆ	ÑÔÐÐÐr!   c                 ó~   — t          | j        t          ¦  «        r"t          j        | j        | j        |d¬¦  «        S |S )NT)Úis_named_optimizer)r9   r   ÚFSDPÚoptim_state_dict_to_loadr)   ©r2   rG   s     r   r\   z$_NamedOptimizer._pre_load_state_dict:  sG   € õ �d”k¥4Ñ(Ô(ð 	ÝÔ0Ø”˜Tœ_¨jÈTðñ ô ð ð Ðr!   c                 óz   — t          | j        t          ¦  «        r t          j        | j        | j        |¦  «         |S rR   )r9   r   r„   Úoptim_state_dictr)   r†   s     r   rI   z _NamedOptimizer._post_state_dictC  s8   € õ �d”k¥4Ñ(Ô(ð 	LÝÔ! $¤+¨t¬À
ÑKÔKÐKØÐr!   )NN)r   NrR   )rP   Nr   N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Ústrr#   r;   r
   r	   Ú	Optimizerr   r   ÚnnÚModuleÚtupler'   r6   r&   rG   r   rT   r   ÚfloatÚpropertyrD   re   ry   r�   r\   rI   r   r!   r   r   r      sƒ  € € € € € ð)ð )ð^ >BØ#'ð)9ð )9à! # u¤|°mÑ'CÐ"CÔDð)9ð œð)9ð ! ¨¨c¨Ô!2Ô3°dÑ:ð	)9ð
 ”	˜DÑ ð)9ð �S˜#�XŒð)9ð �s˜C�x”.ð)9ð 
ð)9ð )9ð )9ð )9ðV/ð /ð /ð /ð&W˜D  c œNð Wð Wð Wð Wð4 Ø5Ð5Ð5Ð5ñ „XØ5àØ>˜H R¨ YÔ/Ð>°EÐ>Ð>Ð>ñ „XØ>ð5ð 5˜H R¨ YÔ/°$Ñ6ð 5À%È$Á,ð 5ð 5ð 5ð 5ð ð%�w˜uœ|¨SÐ0Ô1ð %ð %ð %ñ „Xð%ðh8¨$¨s°C¨x¬.ð h8¸Tð h8ð h8ð h8ð h8ðT9¨7°3¸°8Ô+<ð 9Àð 9ð 9ð 9ð 9ð2 ð  ð  ð  ð¨t°C¸°H¬~ð À$ÀsÈCÀxÄ.ð ð ð ð ð¨4°°S°¬>ð ¸dÀ3ÈÀ8¼nð ð ð ð ð ð r!   r   rL   r   c                 óF   — d                      t          | ¦  «        ¦  «        S )zFConcatenate all param keys as a unique identifier for one param group.rZ   )ÚjoinrH   )rL   s    r   rd   rd   K  s   € à�8Š8•F˜:Ñ&Ô&Ñ'Ô'Ð'r!   )Úloggingr-   Úcollections.abcr   r   r   Úcopyr   Útypingr   r   r#   Útorch.nnr�   r	   Ú'torch.distributed._shard.sharded_tensorr
   Útorch.distributed.fsdpr   r„   r   r*   r�   Ú__annotations__Ú	getLoggerr‰   ÚloggerrŽ   r   rd   r   r!   r   ú<module>r       s?  ðØ €€€€Ø €€€Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø Ð Ð Ð Ð Ð Ø  Ð  Ð  Ð  Ð  Ð  Ð  Ð  à €€€Ø Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð Ø AÐ AÐ AÐ AÐ AÐ AØ CÐ CÐ CÐ CÐ CÐ Cð €ˆˆcŒÐ Ð Ñ à	ˆÔ	˜8Ñ	$Ô	$€ðuð uð uð uð u�e”oñ uô uð uðp	( T¨#¤Yð (°3ð (ð (ð (ð (ð (ð (r!   