Ë
    g^(h$?  ã                   óþ   — d dl Z d dlmZm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 d dlmZmZmZ d dlmZ  e j*                  e«      Zdeeef   fd„Zd	ej6                  defd
„Z G d„ de«      Z G d„ de«      Zy)é    N)ÚabcÚdefaultdict)ÚIterable)ÚAnyÚOptionalÚoverloadÚUnion)Ú_MultiDeviceReplicatorÚ
GradScalerÚOptState)ÚProcessGroupÚreturnc                  ó(   — t         j                  i dœS )N)ÚstageÚfound_inf_per_device)r   ÚREADY© ó    úh/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/distributed/fsdp/sharded_grad_scaler.pyÚ_refresh_per_optimizer_stater      s   € Ü—^‘^¸RÑ@Ð@r   Útensorc                 ó’   — | j                   xs: | j                  j                  dddddt        j                  j                  «       fv S )NÚxlaÚcpuÚhpuÚmtiaÚxpu)Úis_cudaÚdeviceÚtypeÚtorchÚ_CÚ_get_privateuse1_backend_name)r   s    r   Ú_is_supported_devicer$      sH   € Ø�>‰>ò ˜VŸ]™]×/Ñ/ØØØØØÜ�‰×.Ñ.Ó0ð4ð ð r   c                   ó4   — e Zd ZdZdej
                  ddfd„Zy)Ú_GeneralMultiDeviceReplicatorz‡
    Lazily serves tensor to request device. This class extends
    _MultiDeviceReplicator to allow support for "cpu" as a device.
    Úmaster_tensorr   Nc                 ó:   — t        |«      sJ ‚|| _        i | _        y ©N)r$   ÚmasterÚ_per_device_tensors)Úselfr'   s     r   Ú__init__z&_GeneralMultiDeviceReplicator.__init__%   s   € Ü# MÔ2Ð2Ð2Ø#ˆŒØEGˆÕ r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r!   ÚTensorr-   r   r   r   r&   r&      s!   „ ñð
H e§l¡lð H°tô Hr   r&   c                   ój  ‡ — e Zd ZdZddddddej
                  j                  fded	ed
edede	de
dee   ddfˆ fd„Zedej                   dej                   fd„«       Zedeej                      deej                      fd„«       Zedeej                   df   deej                   df   fd„«       Zedeej                      deej                      fd„«       Zdeej                   eej                      f   deej                   eej                      f   fd„Z	 d"dej,                  j.                  dej                   dej                   de
deej2                  ej                   f   f
d„Zdej,                  j.                  ddfd„Zdej                   ddfd„Zd#d eeeej                   f      ddfd!„Zˆ xZS )$ÚShardedGradScaleraA	  
    ShardedGradScaler helps perform gradient scaling in a shard aware manner. It extends
    functionality from GradScaler:
    * Supports Pytorch DDP and FSDP implementations
    * Support CPU offloaded tensors (as used in fully sharded data parallel[FSDP])
    * Supports the custom Mixed Precision loss dtype (fp16, bf16) that FSDP returns
    * Sync inf/nan for scaled gradient tensors on any torch.device (where tensors are placed) across
    nodes

    Example::

        # Creates a ShardedGradScaler once at the beginning of training.
        scaler = ShardedGradScaler()

        for epoch in epochs:
            for input, target in data:
                optimizer.zero_grad()
                output = model(input)
                loss = loss_fn(output, target)

                # Scales loss.  Calls backward() on scaled loss to create scaled gradients.
                scaler.scale(loss).backward()

                # scaler.step() first unscales gradients of the optimizer's params.
                # If gradients don't contain infs/NaNs, optimizer.step() is then called,
                # otherwise, optimizer.step() is skipped.
                scaler.step(optimizer)

                # Updates the scale for next iteration.
                scaler.update()

    See :class:`GradScaler` for explanation of scaling/unscaling and more use cases.

    Args:
        init_scale (float, optional, default=2.**16):  Initial scale factor.
        growth_factor (float, optional, default=2.0):  Factor by which the scale is multiplied during
            :meth:`update` if no inf/NaN gradients occur for ``growth_interval`` consecutive iterations.
        backoff_factor (float, optional, default=0.5):  Factor by which the scale is multiplied during
            :meth:`update` if inf/NaN gradients occur in an iteration.
        growth_interval (int, optional, default=2000):  Number of consecutive iterations without inf/NaN gradients
            that must occur for the scale to be multiplied by ``growth_factor``.
        enabled (bool, optional):  If ``False``, disables gradient scaling. :meth:`step` simply
            invokes the underlying ``optimizer.step()``, and other methods become no-ops.
            Default: ``True``
        process_group (ProcessGroup, optional, default=torch.distributed.group.WORLD):
            process group for sharding
    Úcudag      ð@g      à?g       @iÐ  Tr   Ú
init_scaleÚbackoff_factorÚgrowth_factorÚgrowth_intervalÚenabledÚprocess_groupr   Nc                 ó€   •— t         ‰| �  ||||||¬«       | j                  r|| _        t	        t
        «      | _        y y )N)r6   r7   r8   r9   r:   )Úsuperr-   Ú_enabledr;   r   r   Ú_per_optimizer_states)	r,   r   r6   r7   r8   r9   r:   r;   Ú	__class__s	           €r   r-   zShardedGradScaler.__init__\   sM   ø€ ô 	‰ÑØØ!Ø)Ø'Ø+Øð 	ô 	
ð �=Š=Ø!.ˆDÔÜ)4Ô5QÓ)RˆDÕ&ð r   Úoutputsc                  ó   — y r)   r   ©r,   rA   s     r   ÚscalezShardedGradScaler.scaler   s   € Ø<?r   c                  ó   — y r)   r   rC   s     r   rD   zShardedGradScaler.scaleu   s   € ØHKr   .c                  ó   — y r)   r   rC   s     r   rD   zShardedGradScaler.scalex   s   € ØTWr   c                  ó   — y r)   r   rC   s     r   rD   zShardedGradScaler.scale{   s   € ØPSr   c                 óæ  ‡ ‡‡— ‰ j                   s|S t        |t        j                  «      r‡t	        |«      sJ ‚‰ j
                  €‰ j                  |j                  «       ‰ j
                  €J ‚|‰ j
                  j                  |j                  d¬«      z  }|j                  |j                  «      S g Šdt        t        j                  t        t        j                     f   fˆˆ ˆfd„Š ‰|«      S )NT©r   Únon_blockingÚvalc                 óL  •— t        | t        j                  «      r°t        | «      sJ ‚t	        ‰«      dk(  rY‰j
                  €‰j                  | j                  «       ‰j
                  €J ‚‰j                  t        ‰j
                  «      «       | ‰d   j                  | j                  «      z  }|j                  | j                  «      S t        | t        j                  «      r5t        ‰| «      }t        | t         t"        f«      r t        | «      |«      S |S t%        d«      ‚)Nr   z2outputs must be a Tensor or an iterable of Tensors)Ú
isinstancer!   r2   r$   ÚlenÚ_scaleÚ_lazy_init_scale_growth_trackerr   Úappendr&   Úgetr    Údtyper   r   ÚmapÚlistÚtupleÚ
ValueError)rK   Ú
scaled_valÚiteratorÚapply_scaler,   Ústashs      €€€r   rZ   z,ShardedGradScaler.scale.<locals>.apply_scale“   sæ   ø€ Ü˜#œuŸ|™|Ô,Ü+¨CÔ0Ð0Ð0Ü�u“: ’?Ø—{‘{Ð*Ø×<Ñ<¸S¿Z¹ZÔHØŸ;™;Ð2Ð2Ð2Ø—L‘LÔ!>¸t¿{¹{Ó!KÔLØ  5¨¡8§<¡<°·
±
Ó#;Ñ;�
ð "—‘ s§y¡yÓ1Ð1Ü˜#œsŸ|™|Ô,Ü˜{¨CÓ0�Ü˜c¤D¬% =Ô1Ø$œ4 ›9 XÓ.Ð.Ø�ÜÐQÓRÐRr   )r>   rM   r!   r2   r$   rO   rP   r   Útor    rS   r	   r   )r,   rA   Úscaled_outputrZ   r[   s   `  @@r   rD   zShardedGradScaler.scale~   sÎ   ú€ ð �}Š}ØˆNä�gœuŸ|™|Ô,Ü'¨Ô0Ð0Ð0Ø�{‰{Ð"Ø×4Ñ4°W·^±^ÔDØ—;‘;Ð*Ð*Ð*Ø# d§k¡k§n¡nØ—~‘~°Dð '5ó 'ñ ˆMð !×%Ñ% g§m¡mÓ4Ð4à57ˆð	SœU¤5§<¡<´¼%¿,¹,Ñ1GÐ#GÑH÷ 	Sñ( ˜7Ó#Ð#r   Ú	optimizerÚ	inv_scaleÚ	found_infÚ
allow_fp16c           
      ó†  — t        |«      }t        |«      }t        d„ «      }t        j                  «       5  |j                  D �]9  }|d   D �]-  }	|	j
                  €Œ|s2|	j
                  j                  t        j                  k(  rt        d«      ‚|	j
                  j                  rœ|	j
                  j                  t        j                  u r[|	j
                  j                  t        j                  «      j                  «       }
|
j                  t        j                  «      |	_        |	j
                  j                  «       }n|	j
                  }||j                     |j                     j                  |«       �Œ0 �Œ< |j!                  «       D ]O  \  }}|j#                  «       D ]7  }t        j$                  ||j'                  |«      |j'                  |«      «       Œ9 ŒQ 	 d d d «       |j(                  s3| j*                  €J ‚|j'                  | j*                  j                  «       |j(                  S # 1 sw Y   ŒTxY w)Nc                  ó    — t        t        «      S r)   )r   rU   r   r   r   ú<lambda>z3ShardedGradScaler._unscale_grads_.<locals>.<lambda>¹   s   € ¼ÄTÓ9J€ r   Úparamsz%Attempting to unscale FP16 gradients.)r&   r   r!   Úno_gradÚparam_groupsÚgradrS   Úfloat16rW   Ú	is_sparser    Úfloat32ÚcoalesceÚ_valuesr   rQ   ÚitemsÚvaluesÚ*_amp_foreach_non_finite_check_and_unscale_rR   r+   rO   )r,   r^   r_   r`   ra   Úper_device_inv_scaleÚper_device_found_infÚper_device_and_dtype_gradsÚgroupÚparamÚparam_grad_fp32Ú
to_unscaler   Úper_dtype_gradsÚgradss                  r   Ú_unscale_grads_z!ShardedGradScaler._unscale_grads_©   sà  € ô  =¸YÓGÐÜ<¸YÓGÐô &1Ñ1JÓ%KÐ"Ü�]‰]‹_ñ 	Ø"×/Ñ/ó )�Ø" 8™_ó )�EØ—z‘zÐ)Ø Ù&¨E¯J©J×,<Ñ,<ÄÇÁÒ,MÜ(Ð)PÓQÐQØ—z‘z×+Ò+ð
 !Ÿ:™:×+Ñ+¬u¯}©}Ñ<à.3¯j©j¯o©o¼e¿m¹mÓ.L×.UÑ.UÓ.W˜OØ)8×)=Ñ)=¼e¿m¹mÓ)L˜EœJØ%*§Z¡Z×%7Ñ%7Ó%9™
à%*§Z¡Z˜
à.¨z×/@Ñ/@ÑAØ"×(Ñ(ñç‘f˜ZÖ(ò))ð)ð. ,F×+KÑ+KÓ+Mò Ñ'�˜Ø,×3Ñ3Ó5ò �EÜ×DÑDØØ,×0Ñ0°Ó8Ø,×0Ñ0°Ó8õññ÷1	ðD $×7Ò7Ø—;‘;Ð*Ð*Ð*Ø ×$Ñ$ T§[¡[×%7Ñ%7Ô8Ø#×7Ñ7Ð7÷K	ð 	ús   ·F,H7È7I c                 óž  — | j                   sy | j                  d«       | j                  t        |«         }|d   t        j
                  u rt        d«      ‚|d   t        j                  u rt        d«      ‚| j                  €J ‚| j                  j                  «       j                  «       j                  «       }t        j                  ddt        j                  | j                  j                  ¬«      }| j!                  |||d«      |d	<   t        j
                  |d<   | j                  t        |«         }g }g }g }|d	   j#                  «       D ]Ê  }| j$                  d
k7  rˆ|j                  j&                  d
k(  ro|j)                  |«       |j+                  | j$                  «      }|j)                  |«       |j)                  t-        j.                  |d| j0                  ¬«      «       Œš|j)                  t-        j.                  |d| j0                  ¬«      «       ŒÌ |D ]  }	|	j3                  «        Œ |rt        j4                  ||«       y y )NÚunscale_r   zMunscale_() has already been called on this optimizer since the last update().z(unscale_() is being called after step().)é   g        )rS   r   Tr   r   )Úasync_oprt   )r>   Ú_check_scale_growth_trackerr?   Úidr   ÚUNSCALEDÚRuntimeErrorÚSTEPPEDrO   ÚdoubleÚ
reciprocalÚfloatr!   Úfullrk   r   rz   ro   Ú_devicer    rQ   r\   ÚdistÚ
all_reducer;   ÚwaitÚ_foreach_copy_)
r,   r^   Úoptimizer_stater_   r`   ÚworksÚfound_inf_on_cpusÚfound_inf_on_devicesÚfound_inf_on_deviceÚworks
             r   r|   zShardedGradScaler.unscale_á   s  € Ø�}Š}Øà×(Ñ(¨Ô4à×4Ñ4´R¸	³]ÑCˆà˜7Ñ#¤x×'8Ñ'8Ñ8ÜØ_óð ð ˜WÑ%¬×)9Ñ)9Ñ9ÜÐIÓJÐJð �{‰{Ð&Ð&Ð&Ø—K‘K×&Ñ&Ó(×3Ñ3Ó5×;Ñ;Ó=ˆ	Ü—J‘JØ�#œUŸ]™]°4·;±;×3EÑ3Eô
ˆ	ð 37×2FÑ2FØ�y )¨Tó3
ˆÐ.Ñ/ô $,×#4Ñ#4ˆ˜Ñ ð ×4Ñ4´R¸	³]ÑCˆØˆØÐØ!Ðà(Ð)?Ñ@×GÑGÓIò 	ˆIØ�|‰|˜uÒ$¨×)9Ñ)9×)>Ñ)>À%Ò)GØ!×(Ñ(¨Ô3Ø&/§l¡l°4·<±<Ó&@Ð#Ø$×+Ñ+Ð,?Ô@Ø—‘Ü—O‘OØ+°dÀ$×BTÑBTôõð —‘Ü—O‘O I¸ÀD×DVÑDVÔWõð	ð ò 	ˆDØ�I‰I�Kð	áÜ× Ñ Ð!2Ð4HÕIð r   c                 ó”  — | j                   �| j                  €J ‚|j                  «       dk\  r;| xj                   | j                  z  c_         | j                  j	                  d«       y| j                  dz   }|| j
                  k(  r;| xj                   | j                  z  c_         | j                  j	                  d«       y|| _        y)zÜ
        If found_inf is 1.0 (True), then scale is multiplied by backoff_factor and growth_tracker is set to zero.
        Otherwise, scale is multiplied by the growth factor when the growth interval is reached.
        Ng      ð?r   r}   )rO   Ú_growth_trackerÚitemÚ_backoff_factorÚfill_Ú_growth_intervalÚ_growth_factor)r,   r`   Ú
successfuls      r   Ú_amp_update_scale_cpu_z(ShardedGradScaler._amp_update_scale_cpu_  s¤   € ð
 �{‰{Ð&¨4×+?Ñ+?Ð+KÐKÐKà�>‰>Ó˜sÒ"Ø�KŠK˜4×/Ñ/Ñ/�KØ× Ñ ×&Ñ& qÕ)à×-Ñ-°Ñ1ˆJØ˜T×2Ñ2Ò2Ø—’˜t×2Ñ2Ñ2•Ø×$Ñ$×*Ñ*¨1Õ-à'1�Õ$r   Ú	new_scalec           	      ó  — | j                   sy| j                  d«      \  }}|�¥t        |t        «      r| j                  j                  |«       �n•d}|j                  j                  | j                  k(  sJ |«       ‚|j                  «       dk(  sJ |«       ‚|j                  du sJ |«       ‚| j                  j                  |«       �n| j                  j                  «       D ��cg c]7  }|d   j                  «       D ]  }|j                  |j                  d¬«      ‘Œ! Œ9 }}}t        |«      d	kD  sJ d
«       ‚|d	   }t        |«      dkD  r"t!        dt        |«      «      D ]
  }	|||	   z  }Œ |j                  j                  dk(  r| j#                  |«       nLt%        j&                  | j                  | j(                  || j*                  | j,                  | j.                  «       t1        t2        «      | _        yc c}}w )a™  
        Updates the scale factor.
        If any optimizer steps were skipped the scale is multiplied by ``backoff_factor``
        to reduce it. If ``growth_interval`` unskipped iterations occurred consecutively,
        the scale is multiplied by ``growth_factor`` to increase it.
        Passing ``new_scale`` sets the new scale value manually. (``new_scale`` is not
        used directly, it's used to fill GradScaler's internal scale tensor. So if
        ``new_scale`` was a tensor, later in-place changes to that tensor will not further
        affect the scale GradScaler uses internally.)
        Args:
            new_scale (float or :class:`torch.Tensor`, optional, default=None):  New scale factor.
        .. warning::
            :meth:`update` should only be called at the end of the iteration, after ``scaler.step(optimizer)`` has
            been invoked for all optimizers used this iteration.
        NÚupdatez„new_scale should be a float or a 1-element torch.cuda.FloatTensor or                     torch.FloatTensor with requires_grad=False.r}   Fr   TrI   r   z,No inf checks were recorded prior to update.r   )r>   r   rM   r†   rO   r—   r   r    rˆ   ÚnumelÚrequires_gradÚcopy_r?   ro   r\   rN   Úranger›   r!   Ú_amp_update_scale_r”   r™   r–   r˜   r   r   )
r,   rœ   rO   r”   ÚreasonÚstater`   Ú
found_infsÚfound_inf_combinedÚis
             r   rž   zShardedGradScaler.update'  sê  € ð" �}Š}Øà"&×"BÑ"BÀ8Ó"LÑˆ�àÐ ä˜)¤UÔ+Ø—‘×!Ñ! )Ö,ðAð ð !×'Ñ'×,Ñ,°·±Ò<ÐD¸fÓDÐ<Ø —‘Ó(¨AÒ-Ð5¨vÓ5Ð-Ø ×.Ñ.°%Ñ7Ð?¸Ó?Ð7Ø—‘×!Ñ! )Ö,ð "×7Ñ7×>Ñ>Ó@÷àØ!&Ð'=Ñ!>×!EÑ!EÓ!Gòð ð —‘ F§M¡MÀ�ÕEðØEðˆJñ ô �z“? QÒ&ÐVÐ(VÓVÐ&à!+¨A¡ÐÜ�:‹ Ò"Ü˜q¤# j£/Ó2ò 8�AØ&¨*°Q©-Ñ7Ñ&ð8ð �}‰}×!Ñ! UÒ*Ø×+Ñ+Ð,>Õ?ä×(Ñ(Ø—K‘KØ×(Ñ(Ø&Ø×'Ñ'Ø×(Ñ(Ø×)Ñ)ôô &1Ô1MÓ%NˆÕ"ùó5s   Ã&<G;)Tr)   )r.   r/   r0   r1   r‰   rt   ÚWORLDÚstrr†   ÚintÚboolr   r   r-   r   r!   r2   rD   rU   rV   r   r	   ÚoptimÚ	OptimizerÚdictr   rz   r|   r›   rž   Ú__classcell__)r@   s   @r   r4   r4   +   s=  ø„ ñ.ðd Ø#Ø #Ø"Ø#ØØ04·
±
×0@Ñ0@ñSàðSð ðSð ð	Sð
 ðSð ðSð ðSð   Ñ-ðSð 
õSð, Ø?˜UŸ\™\Ð?¨e¯l©lÒ?ó Ø?àØK˜T %§,¡,Ñ/ÐK°D¸¿¹Ñ4FÒKó ØKàØW˜U 5§<¡<°Ð#4Ñ5ÐW¸%ÀÇÁÈcÐ@QÑ:RÒWó ØWàØS˜X e§l¡lÑ3ÐS¸ÀÇÁÑ8NÒSó ØSð)$Ø˜UŸ\™\¨8°E·L±LÑ+AÐAÑBð)$à	ˆu�|‰|˜X e§l¡lÑ3Ð3Ñ	4ó)$ð`  ñ68à—;‘;×(Ñ(ð68ð —<‘<ð68ð —<‘<ð	68ð
 ð68ð 
ˆe�l‰l˜EŸL™LÐ(Ñ	)ó68ðp2J %§+¡+×"7Ñ"7ð 2J¸Dó 2Jðh2°·±ð 2Àó 2ñ$@O ¨¨u°e·l±lÐ/BÑ)CÑ Dð @OÐPT÷ @Or   r4   ) ÚloggingÚcollectionsr   r   Úcollections.abcr   Útypingr   r   r   r	   r!   Útorch.distributedÚdistributedr‰   Útorch.amp.grad_scalerr
   r   r   Ú"torch.distributed.distributed_c10dr   Ú	getLoggerr.   Úloggerr¯   rª   r   r2   r¬   r$   r&   r4   r   r   r   ú<module>r»      sƒ   ðã ß (Ý $ß 1Ó 1ã Ý  ß NÑ NÝ ;ð 
ˆ×	Ñ	˜8Ó	$€ðA d¨3°¨8¡nó Að §¡ð °$ó ô	HÐ$:ô 	Hô|O˜
õ |Or   