Ë
    ÿÍ:j¹  ã                  ó˜   — d Z ddlmZ ddlZddlmZ ddlZddlmZ ddl	m
Z
  ej                  d«      Zddd„Z	 d	 	 	 	 	 	 	 dd	„Zdd
„Zy)z*Utilities that can be used with Deepspeed.é    )ÚannotationsN)ÚAny)Ú_PATH)Ú_DEEPSPEED_AVAILABLEÚcpuc                óÆ  — |€€t         j                  j                  | d«      }t         j                  j                  |«      r3t	        |«      5 }|j                  «       j                  «       }d d d «       nt        d|› �«      ‚t         j                  j                  | |«      }t         j                  j                  |«      st        dt        › d�«      ‚|S # 1 sw Y   Œ]xY w)NÚlatestz Unable to find 'latest' file at zDirectory 'z' doesn't exist)ÚosÚpathÚjoinÚisfileÚopenÚreadÚstripÚ
ValueErrorÚisdirÚFileNotFoundErrorÚds_checkpoint_dir)Úcheckpoint_dirÚtagÚlatest_pathÚfdÚ	directorys        úz/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/lightning/pytorch/utilities/deepspeed.pyr   r      s¶   € Ø
€{Ü—g‘g—l‘l >°8Ó<ˆÜ�7‰7�>‰>˜+Ô&Ü�kÓ"ð ( bØ—g‘g“i—o‘oÓ'�÷(ð (ô Ð?À¸}ÐMÓNÐNä—‘—‘˜^¨SÓ1€Iä�7‰7�=‰=˜Ô#Ü +Ô.?Ð-@ÀÐ PÓQÐQØÐ÷(ð (ús   ÁCÃC c                ó   — t         st        t        t         «      «      ‚ddlm}m}m}  || |«      }g d¢}t        | «      }  || «      }t        j                  |d   t        d¬«      }	|	d   d   }
 || |
«      }t        j                  |t        d¬«      }|j                  «       D ��ci c]  \  }}||vsŒ||“Œ }}}|D �ci c]  }t        |d«      ||   “Œ }}||d	<   t        d
|› �«       t        j                  ||«       |S c c}}w c c}w )a  Convert ZeRO 2 or 3 checkpoint into a single fp32 consolidated ``state_dict`` file that can be loaded with
    ``torch.load(file)`` + ``load_state_dict()`` and used for training without DeepSpeed. It gets copied into the top
    level checkpoint dir, so the user can easily do the conversion at any point in the future. Once extracted, the
    weights don't require DeepSpeed and can be used in any application. Additionally the script has been modified to
    ensure we keep the lightning state inside the state dict for being able to run
    ``LightningModule.load_from_checkpoint('...')```.

    Args:
        checkpoint_dir: path to the desired checkpoint folder.
            (one that contains the tag-folder, like ``global_step14``)
        output_file: path to the pytorch fp32 state_dict output file (e.g. path/pytorch_model.bin)
        tag: checkpoint tag used as a unique identifier for checkpoint. If not provided will attempt
            to load tag in the file named ``latest`` in the checkpoint folder, e.g., ``global_step14``

    Examples::

        # Lightning deepspeed has saved a directory instead of a file
        convert_zero_checkpoint_to_fp32_state_dict(
            "lightning_logs/version_0/checkpoints/epoch=0-step=0.ckpt/",
            "lightning_model.pt"
        )

    r   )Ú(get_fp32_state_dict_from_zero_checkpointÚget_model_state_fileÚget_optim_files)ÚmoduleÚ	optimizerÚlr_schedulerÚcsr_tensor_module_namesÚskipped_stepsÚglobal_stepsÚdp_world_sizeÚmp_world_sizeF)Úmap_locationÚweights_onlyÚoptimizer_state_dictÚ
zero_stagez_forward_module.Ú
state_dictzSaving fp32 state dict to )r   ÚModuleNotFoundErrorÚstrÚdeepspeed.utils.zero_to_fp32r   r   r   r   ÚtorchÚloadÚ
CPU_DEVICEÚitemsÚ_remove_prefixÚprintÚsave)r   Úoutput_filer   r   r   r   r+   Údeepspeed_statesÚoptim_filesÚoptim_stater*   Ú
model_fileÚclient_stateÚkeyÚvalueÚks                   r   Ú*convert_zero_checkpoint_to_fp32_state_dictr?   .   s"  € õ4  Ü!¤#Ô&:Ó";Ó<Ð<÷ñ ñ :¸.È#ÓN€Jò	Ðô ' ~Ó6€NÙ! .Ó1€KÜ—*‘*˜[¨™^¼*ÐSXÔY€KØÐ3Ñ4°\ÑB€JÙ% n°jÓA€JÜ—:‘:˜j´zÐPUÔV€LØ1=×1CÑ1CÓ1E×e¡: 3¨ÈÐTdÒId�C˜‘JÐe€LÑeð Q[Ö[È1”. Ð$6Ó7¸ÀA¹ÑFÐ[€JÐ[Ø!+€L�Ñä	Ð& { mÐ
4Ô5Ü	‡J�Jˆ|˜[Ô)àÐùó fùò \s   Â(DÂ5DÃDc                óD   — | j                  |«      r| t        |«      d  S | S ©N)Ú
startswithÚlen)r<   Úprefixs     r   r3   r3   p   s#   € Ø!$§¡°Ô!7ˆ3Œs�6‹{ˆ}ÐÐ@¸SÐ@ó    rA   )r   r   r   ú
str | NoneÚreturnr-   )r   r   r6   r   r   rF   rG   zdict[str, Any])r<   r-   rD   r-   rG   r-   )Ú__doc__Ú
__future__r   r
   Útypingr   r/   Ú lightning.fabric.utilities.typesr   Ú&lightning.pytorch.strategies.deepspeedr   Údevicer1   r   r?   r3   © rE   r   ú<module>rO      sd   ðñ 1å "ã 	Ý ã å 2Ý GàˆU�\‰\˜%Ó €
ôð$ BFð?Øð?Ø(-ð?Ø4>ð?àó?ôDArE   