Ë
    g^(h~  ã                   óF   — d dl Z d dlZd dlmZ d dlmZ dgZ G d„ d«      Zy)é    N)Úwraps)ÚBaseSparsifierÚBaseSchedulerc                   óH   — e Zd Zdd„Zd„ Zd„ Zd„ Zd„ Zdd„Zd„ Z	dd	„Z
d
„ Zy)r   c                 ó˜  — t        |t        «      s!t        t        |«      j                  › d�«      ‚|| _        |j                  D �cg c]  }|d   ‘Œ	 c}| _        || _        d„ } || j
                  j                  «      | j
                  _	        d| j
                  _
        d| _
        || _        d| _        | j                  «        y c c}w )Nz6 is not an instance of torch.ao.pruning.BaseSparsifierÚsparsity_levelc                 óÜ   ‡‡‡— t        | dd«      r| S t        j                  | j                  «      Š| j                  Š ‰«       j
                  Š~ t        ‰«      ˆˆˆfd„«       }d|_        |S )NÚ_with_counterFc                  óp   •—  ‰«       }|xj                   dz  c_         ‰j                  |‰«      } || i |¤ŽS )Né   )Ú_step_countÚ__get__)ÚargsÚkwargsÚinstanceÚwrappedÚclsÚfuncÚinstance_refs       €€€úg/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/ao/pruning/scheduler/base_scheduler.pyÚwrapperz=BaseScheduler.__init__.<locals>.with_counter.<locals>.wrapper+   s;   ø€ á'›>�Ø×$Ò$¨Ñ)Õ$ØŸ,™, x°Ó5�Ù Ð/¨Ñ/Ð/ó    T)ÚgetattrÚweakrefÚrefÚ__self__Ú__func__Ú	__class__r   r
   )Úmethodr   r   r   r   s     @@@r   Úwith_counterz,BaseScheduler.__init__.<locals>.with_counter   sf   ú€ Ü�v˜°Ô6à�ô #Ÿ;™; v§¡Ó7ˆLà—?‘?ˆDÙ“.×*Ñ*ˆCØä�4‹[õ0ó ð0ð %)ˆGÔ!ØˆNr   r   F)Ú
isinstancer   Ú	TypeErrorÚtypeÚ__name__Ú
sparsifierÚgroupsÚbase_slÚ
last_epochÚstepr   ÚverboseÚ_get_sl_called_within_step)Úselfr%   r(   r*   Úgroupr    s         r   Ú__init__zBaseScheduler.__init__   s¸   € ä˜*¤nÔ5ÜÜ˜
Ó#×,Ñ,Ð-Ð-cÐdóð ð %ˆŒð >H×=NÑ=NÖO°E˜Ð.Ó/ÒOˆŒØ$ˆŒò
	ñ2  ,¨D¯O©O×,@Ñ,@ÓAˆ�‰ÔØ&'ˆ�‰Ô#Ø !ˆÔØˆŒð 16ˆÔ'à�	‰	�ùòO Ps   ÁCc                 óv   — | j                   j                  «       D ��ci c]  \  }}|dk7  sŒ||“Œ c}}S c c}}w )z¦Returns the state of the scheduler as a :class:`dict`.

        It contains an entry for every variable in self.__dict__ which
        is not the sparsifier.
        r%   )Ú__dict__Úitems)r,   ÚkeyÚvalues      r   Ú
state_dictzBaseScheduler.state_dictA   s=   € ð *.¯©×)<Ñ)<Ó)>÷
Ù%˜3 À#ÈÓBUˆC�‰Jó
ð 	
ùó 
s   ž5¬5c                 ó:   — | j                   j                  |«       y)z³Loads the schedulers state.

        Args:
            state_dict (dict): scheduler state. Should be an object returned
                from a call to :meth:`state_dict`.
        N)r0   Úupdate)r,   r4   s     r   Úload_state_dictzBaseScheduler.load_state_dictK   s   € ð 	�‰×Ñ˜ZÕ(r   c                 ó   — | j                   S )z9Return last computed sparsity level by current scheduler.)Ú_last_sl©r,   s    r   Úget_last_slzBaseScheduler.get_last_slT   s   € à�}‰}Ðr   c                 óP   — | j                   st        j                  d«       t        ‚)NzUTo get the last sparsity level computed by the scheduler, please use `get_last_sl()`.)r+   ÚwarningsÚwarnÚNotImplementedErrorr:   s    r   Úget_slzBaseScheduler.get_slX   s&   € ð ×.Ò.Ü�M‰Mð.ôô "Ð!r   Nc           	      ód   — |r.|€t        d|› d|d›d�«       yt        d|d›d|› d|d›d�«       yy)	z#Display the current sparsity level.Nz"Adjusting sparsity level of group z to z.4eú.zEpoch Ú5dz$: adjusting sparsity level of group )Úprint)r,   Ú
is_verboser-   ÚslÚepochs        r   Úprint_slzBaseScheduler.print_slc   sS   € áØˆ}ÜÐ:¸5¸'ÀÀbÈÀXÈQÐOÕPäØ˜U 2˜JÐ&JÈ5È'ÐQUÐVXÐY\ÐU]Ð]^Ð_õð	 r   c                 ó˜   — | j                   j                  dz   }|dz  }|d| j                  › d�z  }|d| j                  › d�z  }|dz  }|S )Nz (ú
zSparsifier z    base_sl: ú))r   r$   r%   r'   )r,   Úformat_strings     r   Ú__repr__zBaseScheduler.__repr__m   s_   € ØŸ™×/Ñ/°$Ñ6ˆØ˜ÑˆØ˜; t§¡Ð&7°rÐ:Ñ:ˆØ˜=¨¯©¨°bÐ9Ñ9ˆØ˜ÑˆØÐr   c                 óö  — | j                   dk(  rnt        | j                  j                  d«      st	        j
                  dt        «       n3| j                  j                   dk  rt	        j
                  dt        «       | xj                   dz  c_          G d„ d«      } || «      5  | xj                  dz  c_        | j                  «       }d d d «       t        t        | j                  j                  «      «      D ]-  \  }}|\  }}||d<   | j                  | j                  |||«       Œ/ | j                  j                  D �cg c]  }|d   ‘Œ	 c}| _        d| j                  _        y # 1 sw Y   Œ xY wc c}w )	Nr   r
   z¤Seems like `sparsifier.step()` has been overridden after sparsity scheduler initialization. Please, make sure to call `sparsifier.step()` before `scheduler.step()`.z�Detected call of `scheduler.step()` before `sparsifier.step()`. You have to make sure you run the sparsifier.step() BEFORE any calls to the scheduler.step().c                   ó   — e Zd Zd„ Zd„ Zd„ Zy)ú/BaseScheduler.step.<locals>._enable_get_sl_callc                 ó   — || _         y ©N)Úo)r,   rS   s     r   r.   z8BaseScheduler.step.<locals>._enable_get_sl_call.__init__Œ   s	   € Ø�•r   c                 ó(   — d| j                   _        | S )NT©rS   r+   r:   s    r   Ú	__enter__z9BaseScheduler.step.<locals>._enable_get_sl_call.__enter__�   s   € Ø48�—‘Ô1Ø�r   c                 ó&   — d| j                   _        y )NFrU   )r,   r#   r3   Ú	tracebacks       r   Ú__exit__z8BaseScheduler.step.<locals>._enable_get_sl_call.__exit__“   s   € Ø49�—‘Õ1r   N)r$   Ú
__module__Ú__qualname__r.   rV   rY   © r   r   Ú_enable_get_sl_callrP   ‹   s   „ òòó:r   r]   r   T)r   Úhasattrr%   r)   r=   r>   ÚUserWarningr(   r@   Ú	enumerateÚzipr&   rH   r*   r9   Úenable_mask_update)	r,   rG   r]   ÚvaluesÚiÚdataÚparam_grouprF   r-   s	            r   r)   zBaseScheduler.stepu   sF  € ð ×Ñ˜qÒ Ü˜4Ÿ?™?×/Ñ/°ÔAÜ—‘ð*ô  õ	ð —‘×,Ñ,¨qÒ0Ü—‘ð5ô  ô	ð 	×Ò˜AÑÕ÷		:ñ 		:ñ ! Ó&ñ 	#Ø�OŠO˜qÑ �OØ—[‘[“]ˆF÷	#ô !¤ T§_¡_×%;Ñ%;¸VÓ!DÓEò 	6‰GˆAˆtØ"‰OˆK˜Ø,.ˆKÐ(Ñ)Ø�M‰M˜$Ÿ,™,¨¨2¨uÕ5ð	6ð
 ?C¿o¹o×>TÑ>TÖU°U˜Ð/Ó0ÒUˆŒØ-1ˆ�‰Õ*÷	#ð 	#üò Vs   Â%&E*ÅE6Å*E3c                 óÞ   — t        | j                  j                  «      }t        |t        t
        f«      s|g|z  S t        |«      |k7  rt        d|› dt        |«      › �«      ‚t	        |«      S )zPUtility that extends it to the same length as the .groups, ensuring it is a listzExpected variable of length z
, but got )Úlenr%   r&   r!   ÚlistÚtupleÚ
ValueError)r,   ÚvarÚns      r   Ú_make_sure_a_listzBaseScheduler._make_sure_a_list¢   sc   € ä�—‘×&Ñ&Ó'ˆÜ˜#¤¤e˜}Ô-Ø�5˜1‘9Ðä�3‹x˜1Š}Ü Ð#?À¸sÀ*ÌSÐQTËXÈJÐ!WÓXÐXÜ˜“9Ðr   )éÿÿÿÿFrR   )r$   rZ   r[   r.   r4   r7   r;   r@   rH   rM   r)   rn   r\   r   r   r   r      s1   „ ó1òf
ò)òò	"óòó+2óZr   )r=   r   Ú	functoolsr   Ú+torch.ao.pruning.sparsifier.base_sparsifierr   Ú__all__r   r\   r   r   ú<module>rs      s)   ðó Û Ý å Fð Ð
€÷]ò ]r   