Ë
    (täi†  ã                   óÚ  — d dl Z d dlZd dlZd dlZd dlZd dlZd dlZd dl mZ d dlm	Z	 d dl
mZmZ d dlZd dlZ eedd«      �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 d	dlmZm Z m!Z!m"Z"m#Z#m$Z$m%Z%m&Z&  e&«       r,d dl'Z'e'jP                  jR                  jU                  «       rd dl)Z) e"«       rd dl+m,Z,  e#«       rd dl-m.Z.  e%«       rd dl/m0Z0  e$«       rd dl1Z1de2fd„Z3d„ Z4de5fd„Z6	 dCdededejn                  dejn                  dejn                  dejn                  dejn                  de8de9ejn                  ejn                  f   fd„Z:dede;e5ejn                  f   fd „Z<ejz                  fd!ej|                  j~                  e@ej|                  j~                     z  fd"„ZAd#e;e5ejn                  f   d$e5d%ej|                  j~                  fd&„ZBd'e;e5ej|                  j~                  f   de;e5ef   fd(„ZC	 	 	 	 	 dDd)e5d*e2d+e8d,e8d-e8d.ejˆ                  e5z  d/ejŠ                  dz  fd0„ZFdEd)e5fd1„ZGd2„ ZHed3d4œd5ej|                  j~                  ez  d.e5ejˆ                  z  d6eIfd7„«       ZJd8„ ZKd9„ ZLde;fd:„ZMde;fd;„ZN	 	 	 	 	 dFd!ej|                  j~                  d.e5ejˆ                  z  d6eId<eId=eId>e;e5ef   dz  d?eOePej|                  j~                        dz  defd@„ZQ G dA„ dB«      ZRy)Gé    N)Úcontextmanager)Úpartial)ÚAnyÚIterableÚdistributed)Ú
CPUOffloadÚShardingStrategy)ÚFullyShardedDataParallel)Útransformer_auto_wrap_policyé   )ÚUNet2DConditionModel)ÚDiffusionPipeline)ÚSchedulerMixin)Úconvert_state_dict_to_diffusersÚconvert_state_dict_to_peftÚ	deprecateÚis_accelerate_availableÚis_peft_availableÚis_torch_npu_availableÚis_torchvision_availableÚis_transformers_available)Ú
get_logger)Úset_peft_model_state_dict)Ú
transformsÚseedc                 ó(  — t        j                  | «       t        j                   j                  | «       t        j                  | «       t        «       r t        j                  j                  | «       yt        j                  j                  | «       y)z±
    Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`.

    Args:
        seed (`int`): The seed to set.

    Returns:
        `None`
    N)	Úrandomr   ÚnpÚtorchÚmanual_seedr   ÚnpuÚmanual_seed_allÚcuda)r   s    úg/Volumes/fast/ai/experiments/MLX_z-image/.venv/lib/python3.12/site-packages/diffusers/training_utils.pyÚset_seedr%   7   sX   € ô ‡K�K�ÔÜ‡I�I‡N�N�4ÔÜ	×Ñ�dÔÜÔÜ�	‰	×!Ñ! $Õ'ä�
‰
×"Ñ" 4Õ(ó    c                 óØ  — | j                   }|dz  }d|z
  dz  }|j                  |j                  ¬«      |   j                  «       }t	        |j
                  «      t	        |j
                  «      k  r1|d   }t	        |j
                  «      t	        |j
                  «      k  rŒ1|j                  |j
                  «      }|j                  |j                  ¬«      |   j                  «       }t	        |j
                  «      t	        |j
                  «      k  r1|d   }t	        |j
                  «      t	        |j
                  «      k  rŒ1|j                  |j
                  «      }||z  dz  }|S )a�  
    Computes SNR as per
    https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L847-L849
    for the given timesteps using the provided noise scheduler.

    Args:
        noise_scheduler (`NoiseScheduler`):
            An object containing the noise schedule parameters, specifically `alphas_cumprod`, which is used to compute
            the SNR values.
        timesteps (`torch.Tensor`):
            A tensor of timesteps for which the SNR is computed.

    Returns:
        `torch.Tensor`: A tensor containing the computed SNR values for each timestep.
    ç      à?ç      ð?©Údevice).Né   )Úalphas_cumprodÚtor+   ÚfloatÚlenÚshapeÚexpand)Únoise_schedulerÚ	timestepsr-   Úsqrt_alphas_cumprodÚsqrt_one_minus_alphas_cumprodÚalphaÚsigmaÚsnrs           r$   Úcompute_snrr:   K   sH  € ð  %×3Ñ3€NØ(¨#Ñ-ÐØ%(¨>Ñ%9¸cÑ$AÐ!ð .×0Ñ0¸	×8HÑ8HÐ0ÓIÈ)ÑT×ZÑZÓ\ÐÜ
Ð!×'Ñ'Ó
(¬3¨y¯©Ó+?Ò
?Ø1°)Ñ<Ðô Ð!×'Ñ'Ó
(¬3¨y¯©Ó+?Ó
?à×&Ñ& y§¡Ó7€Eà$A×$DÑ$DÈI×L\ÑL\Ð$DÓ$]Ð^gÑ$h×$nÑ$nÓ$pÐ!Ü
Ð+×1Ñ1Ó
2´S¸¿¹Ó5IÒ
IØ(EÀiÑ(PÐ%ô Ð+×1Ñ1Ó
2´S¸¿¹Ó5IÓ
Ià)×0Ñ0°·±ÓA€Eð �5‰=˜QÑ
€CØ€Jr&   Úinterpolation_typec                 ó  — t        «       st        d«      ‚| dk(  rt        j                  j                  }|S | dk(  rt        j                  j
                  }|S | dk(  rt        j                  j                  }|S | dk(  rt        j                  j                  }|S | dk(  rt        j                  j                  }|S | dk(  rt        j                  j                  }|S | dk(  rt        j                  j                  }|S t        d	| › d
�«      ‚)aÕ  
    Maps a string describing an interpolation function to the corresponding torchvision `InterpolationMode` enum. The
    full list of supported enums is documented at
    https://pytorch.org/vision/0.9/transforms.html#torchvision.transforms.functional.InterpolationMode.

    Args:
        interpolation_type (`str`):
            A string describing an interpolation method. Currently, `bilinear`, `bicubic`, `box`, `nearest`,
            `nearest_exact`, `hamming`, and `lanczos` are supported, corresponding to the supported interpolation modes
            in torchvision.

    Returns:
        `torchvision.transforms.InterpolationMode`: an `InterpolationMode` enum used by torchvision's `resize`
        transform.
    zhPlease make sure to install `torchvision` to be able to use the `resolve_interpolation_mode()` function.ÚbilinearÚbicubicÚboxÚnearestÚnearest_exactÚhammingÚlanczoszThe given interpolation mode z’ is not supported. Currently supported interpolation modes are `bilinear`, `bicubic`, `box`, `nearest`, `nearest_exact`, `hamming`, and `lanczos`.)r   ÚImportErrorr   ÚInterpolationModeÚBILINEARÚBICUBICÚBOXÚNEARESTÚNEAREST_EXACTÚHAMMINGÚLANCZOSÚ
ValueError)r;   Úinterpolation_modes     r$   Úresolve_interpolation_moderO   p   s;  € ô  $Ô%ÜØvó
ð 	
ð ˜ZÒ'Ü'×9Ñ9×BÑBÐð& Ðð% 
˜yÒ	(Ü'×9Ñ9×AÑAÐð" Ðð! 
˜uÒ	$Ü'×9Ñ9×=Ñ=Ðð Ðð 
˜yÒ	(Ü'×9Ñ9×AÑAÐð Ðð 
˜Ò	.Ü'×9Ñ9×GÑGÐð Ðð 
˜yÒ	(Ü'×9Ñ9×AÑAÐð Ðð 
˜yÒ	(Ü'×9Ñ9×AÑAÐð Ðô Ø+Ð,>Ð+?ð @mð nó
ð 	
r&   Úunetr3   r4   ÚnoiseÚnoisy_latentsÚtargetÚencoder_hidden_statesÚdream_detail_preservationÚreturnc                 óX  — |j                   j                  |j                  «      |dddf   }d|z
  dz  }	|	|z  }
d}t        j                  «       5   | |||«      j
                  }ddd«       d\  }}|j                  j                  dk(  rO|}||z
  j                  «       }|j                  |
«       |j                  |	|z  «      }|j                  |«      }||fS |j                  j                  dk(  rt        d«      ‚t        d|j                  j                  › �«      ‚# 1 sw Y   Œ¼xY w)	aò  
    Implements "DREAM (Diffusion Rectification and Estimation-Adaptive Models)" from
    https://huggingface.co/papers/2312.00210. DREAM helps align training with sampling to help training be more
    efficient and accurate at the cost of an extra forward step without gradients.

    Args:
        `unet`: The state unet to use to make a prediction.
        `noise_scheduler`: The noise scheduler used to add noise for the given timestep.
        `timesteps`: The timesteps for the noise_scheduler to user.
        `noise`: A tensor of noise in the shape of noisy_latents.
        `noisy_latents`: Previously noise latents from the training loop.
        `target`: The ground-truth tensor to predict after eps is removed.
        `encoder_hidden_states`: Text embeddings from the text model.
        `dream_detail_preservation`: A float value that indicates detail preservation level.
          See reference.

    Returns:
        `tuple[torch.Tensor, torch.Tensor]`: Adjusted noisy_latents and target.
    Nr)   r(   )NNÚepsilonÚv_predictionz/DREAM has not been implemented for v-predictionzUnknown prediction type )r-   r.   r+   r   Úno_gradÚsampleÚconfigÚprediction_typeÚdetachÚmul_ÚaddÚNotImplementedErrorrM   )rP   r3   r4   rQ   rR   rS   rT   rU   r-   r6   Údream_lambdaÚpredÚ_noisy_latentsÚ_targetÚpredicted_noiseÚdelta_noises                   r$   Ú compute_dream_and_update_latentsrh   œ   s7  € ð: %×3Ñ3×6Ñ6°y×7GÑ7GÓHÈÐTXÐZ^Ð`dÐIdÑe€NØ%(¨>Ñ%9¸cÑ$AÐ!ð 1Ð2KÑK€Là€DÜ	�‰�Ù�M 9Ð.CÓD×KÑKˆ÷ 
ð +Ñ€N�GØ×Ñ×-Ñ-°Ò:ØˆØ˜Ñ.×6Ñ6Ó8ˆØ×Ñ˜Ô&Ø&×*Ñ*Ð+HÈ;Ñ+VÓWˆØ—*‘*˜[Ó)ˆð ˜7Ð"Ð"ð 
×	Ñ	×	/Ñ	/°>Ò	AÜ!Ð"SÓTÐTäÐ3°O×4JÑ4J×4ZÑ4ZÐ3[Ð\Ó]Ð]÷ 
ˆús   ÁD Ä D)c                 óÖ   — i }| j                  «       D ]S  \  }}t        |d«      sŒt        |d«      }|€Œ"|j                  «       }|j	                  «       D ]  \  }}|||› d|› �<   Œ ŒU |S )zL
    Returns:
        A state dict containing just the LoRA parameters.
    Úset_lora_layerÚ
lora_layerz.lora.)Únamed_modulesÚhasattrÚgetattrÚ
state_dictÚitems)rP   Úlora_state_dictÚnameÚmodulerk   Úcurrent_lora_layer_sdÚlora_layer_matrix_nameÚ
lora_params           r$   Úunet_lora_state_dictrw   Ò   s…   € ð
 €Oà×*Ñ*Ö,‰ˆˆfÜ�6Ð+Õ,Ü  ¨Ó6ˆJØÑ%Ø(2×(=Ñ(=Ó(?Ð%Ø:O×:UÑ:UÖ:WÑ6Ð*¨JàOY�O t f¨FÐ3IÐ2JÐ$KÒLñ ;Xð -ð Ðr&   Úmodelc                 ó¨   — t        | t        «      s| g} | D ]:  }|j                  «       D ]%  }|j                  sŒ|j	                  |«      |_        Œ' Œ< y)zä
    Casts the training parameters of the model to the specified data type.

    Args:
        model: The PyTorch model whose parameters will be cast.
        dtype: The data type to which the model parameters will be cast.
    N)Ú
isinstanceÚlistÚ
parametersÚrequires_gradr.   Údata)rx   ÚdtypeÚmÚparams       r$   Úcast_training_paramsr‚   å   sG   € ô �eœTÔ"Ø�ˆÛˆØ—\‘\–^ˆEà×"Ó"Ø"ŸX™X e›_�•
ñ $ñ r&   rq   ÚprefixÚtext_encoderc                 óà   — | j                  «       D ��ci c]+  \  }}|j                  |«      sŒ|j                  |d«      › |“Œ- }}}t        t	        |«      «      }t        ||d¬«       yc c}}w )aD  
    Sets the `lora_state_dict` into `text_encoder` coming from `transformers`.

    Args:
        lora_state_dict: The state dictionary to be set.
        prefix: String identifier to retrieve the portion of the state dict that belongs to `text_encoder`.
        text_encoder: Where the `lora_state_dict` is to be set.
    Ú Údefault)Úadapter_nameN)rp   Ú
startswithÚreplacer   r   r   )rq   rƒ   r„   ÚkÚvÚtext_encoder_state_dicts         r$   Ú!_set_state_dict_into_text_encoderrŽ   ö   su   € ð 3B×2GÑ2GÔ2IôÙ2I©$¨!¨QÈQÏ\É\ÐZ`ÕMaˆ1�9‰9�V˜RÓ Ð
! AÑ%Ð2Ið ñ ô 9Ô9XÐYpÓ9qÓrÐÜ˜lÐ,CÐR[Ö\ùó	s
   ”A*®A*Úmodules_to_savec                 ó†   — i }| j                  «       D ]+  \  }}|€Œ	|j                  d   j                  «       ||› d�<   Œ- |S )Nr‡   Ú_lora_adapter_metadata)rp   Úpeft_configÚto_dict)r�   Ú	metadatasÚmodule_namers   s       r$   Ú_collate_lora_metadatar–   	  sT   € Ø€IØ.×4Ñ4Ö6Ñˆ�VØÑØ@F×@RÑ@RÐS\Ñ@]×@eÑ@eÓ@gˆI˜˜Ð%;Ð<Ò=ð  7ð Ðr&   Úweighting_schemeÚ
batch_sizeÚ
logit_meanÚ	logit_stdÚ
mode_scaler+   Ú	generatorc                 ó„  — | dk(  rFt        j                  |||f||¬«      }t         j                  j                  j	                  |«      }|S | dk(  rVt        j
                  |f||¬«      }d|z
  |t        j                  t        j                  |z  dz  «      dz  dz
  |z   z  z
  }|S t        j
                  |f||¬«      }|S )a  
    Compute the density for sampling the timesteps when doing SD3 training.

    Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.

    SD3 paper reference: https://huggingface.co/papers/2403.03206v1.
    Úlogit_normal)ÚmeanÚstdÚsizer+   rœ   Úmode)r¡   r+   rœ   r   r,   )	r   ÚnormalÚnnÚ
functionalÚsigmoidÚrandÚcosÚmathÚpi)r—   r˜   r™   rš   r›   r+   rœ   Úus           r$   Ú%compute_density_for_timestep_samplingr¬     s¾   € ð  ˜>Ò)Ü�L‰L˜j¨i¸z¸mÐTZÐfoÔpˆÜ�H‰H×Ñ×'Ñ'¨Ó*ˆð €Hð 
˜VÒ	#Ü�J‰J˜Z˜M°&ÀIÔNˆØ�‰E�J¤%§)¡)¬D¯G©G°a©K¸!©OÓ"<ÀÑ"AÀAÑ"EÈÑ"IÑJÑJˆð €Hô �J‰J˜Z˜M°&ÀIÔNˆØ€Hr&   c                 óÀ   — | dk(  r|dz  j                  «       }|S | dk(  r)dd|z  z
  d|dz  z  z   }dt        j                  |z  z  }|S t        j                  |«      }|S )zë
    Computes loss weighting scheme for SD3 training.

    Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.

    SD3 paper reference: https://huggingface.co/papers/2403.03206v1.
    Ú
sigma_sqrtg       ÀÚcosmapr   r,   )r/   r©   rª   r   Ú	ones_like)r—   ÚsigmasÚ	weightingÚbots       r$   Úcompute_loss_weighting_for_sd3r´   ,  sz   € ð ˜<Ò'Ø˜T‘\×(Ñ(Ó*ˆ	ð Ðð 
˜XÒ	%Ø�!�f‘*‰n˜q 6¨1¡9™}Ñ,ˆØœŸ™ 3™Ñ'ˆ	ð Ðô —O‘O FÓ+ˆ	ØÐr&   c                  ó"  — t        j                  «        t        j                  j	                  «       rt        j                  j                  «        yt        j                  j                  j	                  «       rt        j                  j                  «        yt        «       rt        j                  j                  «        yt        t        d«      r>t        j                  j	                  «       rt        j                  j                  «        yyy)zV
    Runs garbage collection. Then clears the cache of the available accelerator.
    ÚxpuN)ÚgcÚcollectr   r#   Úis_availableÚempty_cacheÚbackendsÚmpsr   Ú	torch_npur!   rm   r¶   © r&   r$   Úfree_memoryr¿   >  sš   € ô ‡J�J„Lä‡z�z×ÑÔ Ü�
‰
×ÑÕ Ü	�‰×	Ñ	×	(Ñ	(Ô	*Ü�	‰	×ÑÕÜ	Ô	!Ü�‰×!Ñ!Õ#Ü	”˜Ô	¤5§9¡9×#9Ñ#9Ô#;Ü�	‰	×ÑÕð $<Ð	r&   T)ÚoffloadÚmodulesrÀ   c              '   óÊ  K  — |r~t        d„ |D «       «       }|r1|D �cg c]%  }t        |j                  «       «      j                  ‘Œ' }}n t	        |«      dk(  sJ ‚|d   j                  g}|D ]  }|j                  | «       Œ 	 d–— |r&t        |«      D ]  \  }}|j                  |«       Œ yyc c}w # |r&t        |«      D ]  \  }}|j                  |«       Œ w w xY w­w)a  
    Context manager that, if offload=True, moves each module to `device` on enter, then moves it back to its original
    device on exit.

    Args:
        device (`str` or `torch.Device`): Device to move the `modules` to.
        offload (`bool`): Flag to enable offloading.
    c              3   ó<   K  — | ]  }t        |t        «      –— Œ y ­w©N)rz   r   )Ú.0r€   s     r$   Ú	<genexpr>z!offload_models.<locals>.<genexpr>Y  s   è ø€ ÐMÁWÀœ: aÔ):×;ÁWùs   ‚r   r   N)ÚanyÚnextr|   r+   r0   r.   Úzip)r+   rÀ   rÁ   Úis_modelr€   Úoriginal_devicesÚorig_devs          r$   Úoffload_modelsrÍ   N  sä   è ø€ ñ ÜÑMÁWÓMÓMÐMˆáÙELÓMÁWÀ¤ Q§\¡\£^Ó 4× ;Ó ;ÀWÐÑMä�w“< 1Ò$Ð$Ð$à '¨¡
× 1Ñ 1Ð2ÐãˆAØ�D‰D��Lð ðÛáä" 7Ð,<Ö=‘��8Ø—‘�X•ñ  >ð ùò  Nøñ ä" 7Ð,<Ö=‘��8Ø—‘�X•ñ  >ð üs(   ‚C#ž*B1Á;C#ÂB6 Â.C#Â6*C Ã C#c                 ó0  — | st        d«      ‚| j                  «       j                  d«      }g }|D ]²  }t        j                  d|«      }|st        d|› d�«      ‚	 t        |j                  d«      «      }t        |j                  d«      «      }|dk  s|dk  rt        d	«      ‚|d
z  dk7  s|d
z  dk7  rt        j                  d|› d|› d�«       |j                  ||f«       Œ´ |st        d«      ‚|S # t         $ r}t        d|› d|› �«      |‚d}~ww xY w)zGParses a string defining buckets into a list of (height, width) tuples.zBucket string cannot be empty.Ú;z^\s*(\d+)\s*,\s*(\d+)\s*$zInvalid bucket format: 'z'. Expected 'height,width'.r   r,   r   z,Bucket dimensions must be positive integers.é   zBucket dimension (Ú,z.) not divisible by 8. This might cause issues.z Invalid integer in bucket pair 'z': Nz.No valid buckets found in the provided string.)
rM   ÚstripÚsplitÚreÚmatchÚintÚgroupÚwarningsÚwarnÚappend)Úbuckets_strÚbucket_pairsÚparsed_bucketsÚpair_strrÕ   ÚheightÚwidthÚes           r$   Úparse_buckets_stringrâ   n  s1  € áÜÐ9Ó:Ð:à×$Ñ$Ó&×,Ñ,¨SÓ1€LØ€NÛ ˆÜ—‘Ð5°xÓ@ˆÙÜÐ7¸°zÐA\Ð]Ó^Ð^ð		YÜ˜Ÿ™ Q›Ó(ˆFÜ˜Ÿ™ A›Ó'ˆEØ˜Š{˜e qšjÜ Ð!OÓPÐPØ˜‰z˜QŠ %¨!¡)¨q¢.Ü—‘Ð 2°6°(¸!¸E¸7ÐBpÐqÔrØ×!Ñ! 6¨5 /Õ2ð !ñ ÜÐIÓJÐJàÐøô ò 	YÜÐ?À¸zÈÈQÈCÐPÓQÐWXÐXûð	Yús   ÁBC5Ã5	DÃ>DÄDc                 óˆ   — t        d«      }d}t        |«      D ]&  \  }\  }}t        | |z  ||z  z
  «      }||k  sŒ#|}|}Œ( |S )z6Finds the closes bucket to the given height and width.ÚinfN)r/   Ú	enumerateÚabs)	ÚhÚwÚbucket_optionsÚ
min_metricÚbest_bucket_idxÚ
bucket_idxÚbucket_hÚbucket_wÚmetrics	            r$   Úfind_nearest_bucketrð   Š  s[   € ä�u“€JØ€OÜ,5°nÖ,EÑ(ˆ
Ñ(�X˜xÜ�Q˜‘\ A¨¡LÑ0Ó1ˆØ�ZÓØˆJØ(‰Oð	 -Fð
 Ðr&   c           	      óâ   — | j                  «       D ��ci c]N  \  }}|t        |t        j                  «      r,|j	                  «       j                  «       j                  «       n|“ŒP c}}S c c}}w rÄ   )rp   rz   r   ÚTensorr^   ÚcpuÚ
contiguous)Ústate_dictsr‹   rŒ   s      r$   Ú_to_cpu_contiguousrö   –  sW   € Ø_j×_pÑ_pÔ_rÔsÑ_rÑW[ÐWXÐZ[ˆA´
¸1¼e¿l¹lÔ0Kˆq�x‰x‹z�~‰~Ó×*Ñ*Ô,ÐQRÑRÐ_rÒsÐsùÓss   ”AA+c                 óî   — i }t        | j                  dd«      }|€t        d«      ‚| j                  j                  }|€t        j
                  |d<   |S |j                  xs t        j
                  |d<   |S )zT
    Extract and convert FSDP config from Accelerator into PyTorch FSDP kwargs.
    Úfsdp_pluginNzLAccelerate isn't configured to handle FSDP. Please update your installation.Úsharding_strategy)rn   ÚstaterM   rø   r	   Ú
FULL_SHARDrù   )ÚacceleratorÚkwargsÚ
fsdp_staterø   s       r$   Ú get_fsdp_kwargs_from_acceleratorrÿ   š  s‚   € ð
 €FÜ˜×*Ñ*¨M¸4Ó@€JàÐÜÐgÓhÐhà×#Ñ#×/Ñ/€KàÐä&6×&AÑ&AˆÐ"Ñ#ð
 €Mð '2×&CÑ&CÒ&bÔGW×GbÑGbˆÐ"Ñ#à€Mr&   Úuse_orig_paramsÚlimit_all_gathersÚfsdp_kwargsÚtransformer_layer_clsc                 óJ  — t        t        «      }|€Jt        | j                  j                  j
                  d   «      }|j                  d|j                  › �«       t        t        |h¬«      }||rt        |¬«      nd|||dœ}	|r|	j                  |«       t        | fi |	¤Ž}
|
S )uY  
    Wrap a model with FSDP using common defaults and optional transformer auto-wrapping.

    Args:
        model: Model to wrap
        device: Target device (e.g., accelerator.device)
        offload: Whether to enable CPU parameter offloading
        use_orig_params: Whether to use original parameters
        limit_all_gathers: Whether to limit all gathers
        fsdp_kwargs: FSDP arguments (sharding_strategy, etc.) â€” usually from Accelerate config
        transformer_layer_cls: Classes for auto-wrapping (if not using policy from fsdp_kwargs)

    Returns:
        FSDP-wrapped model
    Nr   z8transformer_layer_cls is not provided, auto-inferred as )r  )Úoffload_params)Ú	device_idÚcpu_offloadr   r  Úauto_wrap_policy)r   Ú__name__Útyperx   Úlanguage_modelÚlayersÚinfor   r   r   ÚupdateÚFSDP)rx   r+   rÀ   r   r  r  r  Úloggerr  r\   Ú
fsdp_models              r$   Úwrap_with_fsdpr  ±  s«   € ô2 œÓ!€FàÐ$ä $ U§[¡[×%?Ñ%?×%FÑ%FÀqÑ%IÓ JÐØ�‰ÐNÐOd×OmÑOmÐNnÐoÔpô Ô;ÐTiÐSjÔkÐð Ù=D”z°Õ9È$Ø*Ø.Ø,ñ€Fñ Ø�‰�kÔ"ä�eÑ&˜vÑ&€JØÐr&   c                   ó  — e Zd ZdZ	 	 	 	 	 	 	 	 	 ddeej                  j                     dedede	de
dee	z  d	ee	z  d
e
dedz  deeef   dz  fd„Zeddd„«       Zd„ Zde	defd„Z ej&                  «       deej                  j                     fd„«       Zdeej                  j                     ddfd„Zdd„Zd dd„Zdefd„Zdeej                  j                     ddfd„Zdeej                  j                     ddfd„Zdeddfd„Zy)!ÚEMAModelz6
    Exponential Moving Average of models weights
    Nr|   ÚdecayÚ	min_decayÚupdate_after_stepÚuse_ema_warmupÚ	inv_gammaÚpowerÚforeachÚ	model_clsÚmodel_configc                 óÌ  — t        |t        j                  j                  «      r#d}t	        dd|d¬«       |j                  «       }d}|j                  dd«      �d	}t	        dd|d¬«       |d   }|j                  d
d«      �d}t	        d
d|d¬«       |d
   }t        |«      }|D �cg c]   }|j                  «       j                  «       ‘Œ" c}| _
        |j                  dd«      �&d}t	        dd|d¬«       | j                  |d   ¬«       d| _        || _        || _        || _        || _        || _        || _        d| _        d| _        || _        |	| _        |
| _        yc c}w )ar  
        Args:
            parameters (Iterable[torch.nn.Parameter]): The parameters to track.
            decay (float): The decay factor for the exponential moving average.
            min_decay (float): The minimum decay factor for the exponential moving average.
            update_after_step (int): The number of steps to wait before starting to update the EMA weights.
            use_ema_warmup (bool): Whether to use EMA warmup.
            inv_gamma (float):
                Inverse multiplicative factor of EMA warmup. Default: 1. Only used if `use_ema_warmup` is True.
            power (float): Exponential factor of EMA warmup. Default: 2/3. Only used if `use_ema_warmup` is True.
            foreach (bool): Use torch._foreach functions for updating shadow parameters. Should be faster.
            device (str | torch.device | None): The device to store the EMA weights on. If None, the EMA
                        weights will be stored on CPU.

        @crowsonkb's notes on EMA Warmup:
            If gamma=1 and power=1, implements a simple average. gamma=1, power=2/3 are good values for models you plan
            to train for a million or more steps (reaches decay factor 0.999 at 31.6K steps, 0.9999 at 1M steps),
            gamma=1, power=3/4 for models you plan to train for less (reaches decay factor 0.999 at 10K steps, 0.9999
            at 215.4k steps).
        zzPassing a `torch.nn.Module` to `ExponentialMovingAverage` is deprecated. Please pass the parameters of the module instead.z9passing a `torch.nn.Module` to `ExponentialMovingAverage`ú1.0.0F©Ústandard_warnTÚ	max_valueNzCThe `max_value` argument is deprecated. Please use `decay` instead.Ú	min_valuezGThe `min_value` argument is deprecated. Please use `min_decay` instead.r+   z=The `device` argument is deprecated. Please use `to` instead.r*   r   )rz   r   r¤   ÚModuler   r|   Úgetr{   Úcloner^   Úshadow_paramsr.   Útemp_stored_paramsr  r  r  r  r  r  Úoptimization_stepÚcur_decay_valuer  r  r  )Úselfr|   r  r  r  r  r  r  r  r  r  rý   Údeprecation_messageÚps                 r$   Ú__init__zEMAModel.__init__é  sx  € ôF �j¤%§(¡(§/¡/Ô2ðDð  ô ØKØØ#Ø#õ	ð $×.Ñ.Ó0ˆJð "ˆNà�:‰:�k 4Ó(Ð4Ø"gÐÜ�k 7Ð,?ÈuÕUØ˜;Ñ'ˆEà�:‰:�k 4Ó(Ð4Ø"kÐÜ�k 7Ð,?ÈuÕUØ˜{Ñ+ˆIä˜*Ó%ˆ
Ù:DÓE¹*°Q˜aŸg™g›i×.Ñ.Õ0¸*ÑEˆÔà�:‰:�h Ó%Ð1Ø"aÐÜ�h Ð)<ÈEÕRØ�G‰G˜6 (Ñ+ˆGÔ,à"&ˆÔàˆŒ
Ø"ˆŒØ!2ˆÔØ,ˆÔØ"ˆŒØˆŒ
Ø!"ˆÔØ#ˆÔØˆŒà"ˆŒØ(ˆÕùò) Fs   Â'%E!rV   c                 ó¾   — |j                  |d¬«      \  }}|j                  |«      } | |j                  «       ||j                  |¬«      }|j	                  |«       |S )NT)Úreturn_unused_kwargs)r  r  r  )Úfrom_configÚfrom_pretrainedr|   r\   Úload_state_dict)ÚclsÚpathr  r  Ú_Ú
ema_kwargsrx   Ú	ema_models           r$   r2  zEMAModel.from_pretrained=  s]   € à!×-Ñ-¨dÈÐ-ÓN‰ˆˆ:Ø×)Ñ)¨$Ó/ˆá˜×(Ñ(Ó*°iÈeÏlÉlÐdkÔlˆ	à×!Ñ! *Ô-ØÐr&   c                 ór  — | j                   €t        d«      ‚| j                  €t        d«      ‚| j                   j                  | j                  «      }| j	                  «       }|j                  dd «        |j                  di |¤Ž | j                  |j                  «       «       |j                  |«       y )NzJ`save_pretrained` can only be used if `model_cls` was defined at __init__.zM`save_pretrained` can only be used if `model_config` was defined at __init__.r'  r¾   )
r  rM   r  r1  ro   ÚpopÚregister_to_configÚcopy_tor|   Úsave_pretrained)r+  r5  rx   ro   s       r$   r=  zEMAModel.save_pretrainedG  sš   € Ø�>‰>Ð!ÜÐiÓjÐjà×ÑÐ$ÜÐlÓmÐmà—‘×*Ñ*¨4×+<Ñ+<Ó=ˆØ—_‘_Ó&ˆ
Ø�‰�¨Ô-à ˆ× Ñ Ñ. :Ò.Ø�‰�U×%Ñ%Ó'Ô(Ø×Ñ˜dÕ#r&   r)  c                 ó  — t        d|| j                  z
  dz
  «      }|dk  ry| j                  r$dd|| j                  z  z   | j                   z  z
  }nd|z   d|z   z  }t        || j                  «      }t        || j                  «      }|S )zN
        Compute the decay factor for the exponential moving average.
        r   r   ç        é
   )Úmaxr  r  r  r  Úminr  r  )r+  r)  Ústepr*  s       r$   Ú	get_decayzEMAModel.get_decayV  s�   € ô �1Ð'¨$×*@Ñ*@Ñ@À1ÑDÓEˆà�1Š9Øà×ÒØ 1 t¨d¯n©nÑ'<Ñ#<À$Ç*Á*ÀÑ"LÑL‰Oà  4™x¨B°©IÑ6ˆOä˜o¨t¯z©zÓ:ˆä˜o¨t¯~©~Ó>ˆØÐr&   c           	      óÐ  — t        |t        j                  j                  «      r!d}t	        dd|d¬«       |j                  «       }t        |«      }| xj                  dz  c_        | j                  | j                  «      }|| _	        d|z
  }t        j                  «       }| j                  �rZt        «       rIt        j                  j                   j#                  «       r!t         j$                  j'                  |d ¬«      }|5  |D �cg c]  }|j(                  sŒ|‘Œ }}t+        | j,                  |«      D ��cg c]  \  }}|j(                  sŒ|‘Œ }	}}t/        |«      t/        |«      k  rgt        j0                  t+        | j,                  |«      D ��cg c]  \  }}|j(                  rŒ|‘Œ c}}|D �cg c]  }|j(                  rŒ|‘Œ c}d¬	«       t        j2                  |	t        j4                  |	|«      |¬
«       d d d «       y t+        | j,                  |«      D ]˜  \  }}t        «       rIt        j                  j                   j#                  «       r!t         j$                  j'                  |d ¬«      }|5  |j(                  r|j7                  |||z
  z  «       n|j9                  |«       d d d «       Œš y c c}w c c}}w c c}}w c c}w # 1 sw Y   y xY w# 1 sw Y   ŒÇxY w)NzPassing a `torch.nn.Module` to `ExponentialMovingAverage.step` is deprecated. Please pass the parameters of the module instead.z>passing a `torch.nn.Module` to `ExponentialMovingAverage.step`r  Fr   r   )Úmodifier_rankT)Únon_blocking)r7   )rz   r   r¤   r$  r   r|   r{   r)  rD  r*  Ú
contextlibÚnullcontextr  r   ÚtransformersÚintegrationsÚ	deepspeedÚis_deepspeed_zero3_enabledÚzeroÚGatheredParametersr}   rÉ   r'  r0   Ú_foreach_copy_Ú_foreach_sub_Ú_foreach_subÚsub_Úcopy_)
r+  r|   r,  r  Úone_minus_decayÚcontext_managerr�   Úparams_gradÚs_paramÚs_params_grads
             r$   rC  zEMAModel.stepi  sn  € ä�j¤%§(¡(§/¡/Ô2ðDð  ô ØPØØ#Ø#õ	ð $×.Ñ.Ó0ˆJä˜*Ó%ˆ
à×Ò !Ñ#Õð —‘˜t×5Ñ5Ó6ˆØ$ˆÔØ˜e™)ˆä$×0Ñ0Ó2ˆà�<‹<Ü(Ô*¬|×/HÑ/H×/RÑ/R×/mÑ/mÔ/oÜ"+§.¡.×"CÑ"CÀJÐ^bÐ"CÓ"c�â Ù2<ÓT±*¨À×@SÓ@Sšu°*�ÐTä25°d×6HÑ6HÈ*Ô2Uô!Ù2U¡ ¨ÐY^×YlÓYl’GÐ2Uð ñ !ô �{Ó#¤c¨*£oÒ5Ü×(Ñ(Ü7:¸4×;MÑ;MÈzÔ7ZÔvÑ7Z¡^ W¨eÐbg×buÓbušÐ7ZÒvÙ,6ÓR©J 5¸e×>QÓ>Qš¨JÑRØ%)õô ×#Ñ#Ø!¤5×#5Ñ#5°mÀ[Ó#QÐYhõ÷ !�ô$ #& d×&8Ñ&8¸*Ö"E‘�˜Ü,Ô.´<×3LÑ3L×3VÑ3V×3qÑ3qÔ3sÜ&/§n¡n×&GÑ&GÈÐ]aÐ&GÓ&b�Oâ$Ø×*Ò*ØŸ™ _¸À%¹Ñ%HÕIàŸ™ eÔ,÷	 %�_ñ	 #Fùò# Uùó!ùó wùÚR÷ !�ú÷, %�_úsg   ÄKÄJ:ÄJ:ÄKÄ9J?ÅJ?ÅAKÆKÆ)KÆ-KÆ4KÇKÇ
5KÉ96KÊ:KËKËK%	c           
      óò  — t        |«      }| j                  ryt        j                  |D �cg c]  }|j                  ‘Œ c}t        | j                  |«      D ��cg c]*  \  }}|j                  |j                  «      j                  ‘Œ, c}}«       yt        | j                  |«      D ]C  \  }}|j                  j                  |j                  |j                  «      j                  «       ŒE yc c}w c c}}w )aa  
        Copy current averaged parameters into given collection of parameters.

        Args:
            parameters: Iterable of `torch.nn.Parameter`; the parameters to be
                updated with the stored moving averages. If `None`, the parameters with which this
                `ExponentialMovingAverage` was initialized will be used.
        N)
r{   r  r   rP  r~   rÉ   r'  r.   r+   rT  )r+  r|   r�   rX  s       r$   r<  zEMAModel.copy_to£  s½   € ô ˜*Ó%ˆ
Ø�<Š<Ü× Ñ Ù)3Ó4© �—“¨Ñ4ÜEHÈ×I[ÑI[Ð]gÔEhÔiÑEh±>°7¸E�—‘˜EŸL™LÓ)×.Ó.ÐEhÒiõô
 #& d×&8Ñ&8¸*Ö"E‘�˜Ø—
‘
× Ñ  §¡¨E¯L©LÓ!9×!>Ñ!>Õ?ñ #Fùò	 5ùÛis   «C.Á/C3c                 óh   — | j                   D �cg c]  }|j                  «       ‘Œ c}| _         yc c}w )zª
        Move internal buffers of the ExponentialMovingAverage to pinned memory. Useful for non-blocking transfers for
        offloading EMA params to the host.
        N)r'  Ú
pin_memory)r+  r-  s     r$   r\  zEMAModel.pin_memory¶  s,   € ð 7;×6HÒ6HÓIÑ6H°˜aŸl™l�nÐ6HÑIˆÕùÒIs   �/c                 ó¶   — | j                   D �cg c]9  }|j                  «       r|j                  |||¬«      n|j                  ||¬«      ‘Œ; c}| _         yc c}w )z£
        Move internal buffers of the ExponentialMovingAverage to `device`.

        Args:
            device: like `device` argument to `torch.Tensor.to`
        )r+   r   rG  )r+   rG  N)r'  Úis_floating_pointr.   )r+  r+   r   rG  r-  s        r$   r.   zEMAModel.to¾  sf   € ð ×'Ò'ó	
ñ (�ð ×"Ñ"Ô$ð �D‰D˜ e¸,ˆDÔGà—‘˜V°,�Ó?ñ@ð (ñ	
ˆÕùò 
s   �>Ac           	      ó¸   — | j                   | j                  | j                  | j                  | j                  | j
                  | j                  | j                  dœS )z©
        Returns the state of the ExponentialMovingAverage as a dict. This method is used by accelerate during
        checkpointing to save the ema state dict.
        ©r  r  r)  r  r  r  r  r'  r`  )r+  s    r$   ro   zEMAModel.state_dictÍ  sN   € ð —Z‘ZØŸ™Ø!%×!7Ñ!7Ø!%×!7Ñ!7Ø"×1Ñ1ØŸ™Ø—Z‘ZØ!×/Ñ/ñ	
ð 		
r&   c                 óŒ   — |D �cg c].  }|j                  «       j                  «       j                  «       ‘Œ0 c}| _        yc c}w )zµ
        Saves the current parameters for restoring later.

        Args:
            parameters: Iterable of `torch.nn.Parameter`. The parameters to be temporarily stored.
        N)r^   ró   r&  r(  )r+  r|   r�   s      r$   ÚstorezEMAModel.storeà  s8   € ñ NXÓ"XÉZÀE 5§<¡<£>×#5Ñ#5Ó#7×#=Ñ#=Õ#?ÈZÑ"XˆÕùÒ"Xs   …3Ac                 ó¢  — | j                   €t        d«      ‚| j                  rXt        j                  |D �cg c]  }|j
                  ‘Œ c}| j                   D �cg c]  }|j
                  ‘Œ c}«       d| _         yt        | j                   |«      D ]*  \  }}|j
                  j                  |j
                  «       Œ, d| _         yc c}w c c}w )aG  
        Restore the parameters stored with the `store` method. Useful to validate the model with EMA parameters
        without: affecting the original optimization process. Store the parameters before the `copy_to()` method. After
        validation (or model saving), use this to restore the former parameters.

        Args:
            parameters: Iterable of `torch.nn.Parameter`; the parameters to be
                updated with the stored parameters. If `None`, the parameters with which this
                `ExponentialMovingAverage` was initialized will be used.
        NzGThis ExponentialMovingAverage has no `store()`ed weights to `restore()`)r(  ÚRuntimeErrorr  r   rP  r~   rÉ   rT  )r+  r|   r�   Úc_params       r$   ÚrestorezEMAModel.restoreé  sµ   € ð ×"Ñ"Ð*ÜÐhÓiÐiØ�<Š<Ü× Ñ Ù)3Ó4© �—“¨Ñ4ÐSW×SjÒSjÓ6kÑSjÈ°w·|³|ÐSjÑ6kôð #'ˆÕô	 #& d×&=Ñ&=¸zÖ"J‘�˜Ø—
‘
× Ñ  §¡Õ.ð #Kð #'ˆÕùò 5ùÒ6ks   ·CÁC
ro   c                 óò  — t        j                  |«      }|j                  d| j                  «      | _        | j                  dk  s| j                  dkD  rt	        d«      ‚|j                  d| j
                  «      | _        t        | j
                  t        «      st	        d«      ‚|j                  d| j                  «      | _        t        | j                  t        «      st	        d«      ‚|j                  d	| j                  «      | _
        t        | j                  t        «      st	        d
«      ‚|j                  d| j                  «      | _        t        | j                  t        «      st	        d«      ‚|j                  d| j                  «      | _        t        | j                  t        t        f«      st	        d«      ‚|j                  d| j                  «      | _        t        | j                  t        t        f«      st	        d«      ‚|j                  dd«      }|�T|| _        t        | j                  t         «      st	        d«      ‚t#        d„ | j                  D «       «      st	        d«      ‚yy)a  
        Loads the ExponentialMovingAverage state. This method is used by accelerate during checkpointing to save the
        ema state dict.

        Args:
            state_dict (dict): EMA state. Should be an object returned
                from a call to :meth:`state_dict`.
        r  r?  r)   zDecay must be between 0 and 1r  zInvalid min_decayr)  zInvalid optimization_stepr  zInvalid update_after_stepr  zInvalid use_ema_warmupr  zInvalid inv_gammar  zInvalid powerr'  Nzshadow_params must be a listc              3   óP   K  — | ]  }t        |t        j                  «      –— Œ  y ­wrÄ   )rz   r   rò   )rÅ   r-  s     r$   rÆ   z+EMAModel.load_state_dict.<locals>.<genexpr>/  s   è ø€ ÐOÑ<N°q”z !¤U§\¡\×2Ñ<Nùs   ‚$&z!shadow_params must all be Tensors)ÚcopyÚdeepcopyr%  r  rM   r  rz   r/   r)  rÖ   r  r  Úboolr  r  r'  r{   Úall)r+  ro   r'  s      r$   r3  zEMAModel.load_state_dict  sé  € ô —]‘] :Ó.ˆ
à—^‘^ G¨T¯Z©ZÓ8ˆŒ
Ø�:‰:˜Ò˜tŸz™z¨CÒ/ÜÐ<Ó=Ð=à#Ÿ™¨°T·^±^ÓDˆŒÜ˜$Ÿ.™.¬%Ô0ÜÐ0Ó1Ð1à!+§¡Ð0CÀT×E[ÑE[Ó!\ˆÔÜ˜$×0Ñ0´#Ô6ÜÐ8Ó9Ð9à!+§¡Ð0CÀT×E[ÑE[Ó!\ˆÔÜ˜$×0Ñ0´#Ô6ÜÐ8Ó9Ð9à(Ÿn™nÐ-=¸t×?RÑ?RÓSˆÔÜ˜$×-Ñ-¬tÔ4ÜÐ5Ó6Ð6à#Ÿ™¨°T·^±^ÓDˆŒÜ˜$Ÿ.™.¬5´#¨,Ô7ÜÐ0Ó1Ð1à—^‘^ G¨T¯Z©ZÓ8ˆŒ
Ü˜$Ÿ*™*¤u¬c lÔ3Ü˜_Ó-Ð-à"Ÿ™ ¸Ó=ˆØÐ$Ø!.ˆDÔÜ˜d×0Ñ0´$Ô7Ü Ð!?Ó@Ð@ÜÑO¸D×<NÒ<NÓOÔOÜ Ð!DÓEÐEð Pð	 %r&   )	g§èH.ÿï?r?  r   Fr)   gUUUUUUå?FNN)F)rV   r  )rV   N)NNF)r	  Ú
__module__Ú__qualname__Ú__doc__r   r   r¤   Ú	Parameterr/   rÖ   rk  r   ÚdictÚstrr.  Úclassmethodr2  r=  rD  rZ   rC  r<  r\  r.   ro   rb  rf  r3  r¾   r&   r$   r  r  ä  s·  „ ñð ØØ!"Ø$Ø!$Ø"ØØ $Ø.2ñR)à˜UŸX™X×/Ñ/Ñ0ðR)ð ðR)ð ð	R)ð
 ðR)ð ðR)ð ˜3‘;ðR)ð �s‰{ðR)ð ðR)ð ˜‘:ðR)ð ˜3 ˜8‘n tÑ+óR)ðh óó ðò$ð¨3ð °5ó ð& €U‡]�]ƒ_ð7-˜x¨¯©×(:Ñ(:Ñ;ò 7-ó ð7-ðr@ (¨5¯8©8×+=Ñ+=Ñ">ð @À4ó @ó&Jô
ð
˜Dó 
ð&Y ¨¯©×);Ñ);Ñ <ð YÀó Yð' (¨5¯8©8×+=Ñ+=Ñ">ð 'À4ó 'ð2.F¨$ð .F°4ô .Fr&   r  )r)   )NNNró   NrÄ   )TTTNN)SrH  ri  r·   r©   r   rÔ   rØ   r   Ú	functoolsr   Útypingr   r   Únumpyr   r   rn   Útorch.distributed.fsdpr   r	   r
   r  Útorch.distributed.fsdp.wrapr   Úmodelsr   Ú	pipelinesr   Ú
schedulersr   Úutilsr   r   r   r   r   r   r   r   rJ  rK  rL  rM  Úaccelerate.loggingr   Úpeftr   Útorchvisionr   r½   rÖ   r%   r:   rr  rO   rò   r/   Útuplerh   rq  rw   Úfloat32r¤   r$  r{   r‚   rŽ   r–   r+   Ú	Generatorr¬   r´   r¿   rk  rÍ   râ   rð   rö   rÿ   Úsetr
  r  r  r¾   r&   r$   Ú<module>r„     s‘  ðÛ Û Û 	Û Û Û 	Û Ý %Ý ß  ã Û ñ ˆ5�- Ó&Ð2ßCÝGÝHå (Ý (Ý &÷	÷ 	ó 	ñ ÔÛà× Ñ ×*Ñ*×EÑEÔGÛáÔÝ-áÔÝ.áÔÝ&áÔÛð)�3ó )ò("ðJ)°3ó )ðh (+ñ3#Ø
ð3#à#ð3#ð �|‰|ð3#ð �<‰<ð	3#ð
 —<‘<ð3#ð �L‰Lð3#ð !Ÿ<™<ð3#ð  %ð3#ð ˆ5�<‰<˜Ÿ™Ð%Ñ&ó3#ðlÐ3ð ¸¸SÀ%Ç,Á,Ð=NÑ8Oó ð& PUÏ}É}ñ - §¡§¡°$°u·x±x·±Ñ2GÑ Gó -ð"]Ø˜#˜uŸ|™|Ð+Ñ,ð]Ø69ð]ØINÏÉÏÉó]ð&¨D°°e·h±h·o±oÐ1EÑ,Fð È4ÐPSÐUXÐPXÉ>ó ð ØØØ!&Ø(,ñØðàðð ðð ð	ð
 ðð �L‰L˜3Ñðð �‰ Ñ%óñ6°Só ò$ ð  Ønrò ˜UŸX™XŸ_™_Ð/@Ñ@ð È#ÐPU×P\ÑP\ÑJ\ð Ðgkò ó ðò>ò8	ðt tó tð°Tó ð4 Ø Ø"Ø)-Ø?Cñ/Ø�8‰8�?‰?ð/à�%—,‘,Ñð/ð ð/ð ð	/ð
 ð/ð �c˜3�h‘ $Ñ&ð/ð ˜t E§H¡H§O¡OÑ4Ñ5¸Ñ<ð/ð 
ó/÷fLFò LFr&   