Ë
    g^(h~5  ã                   óÊ   — d dl Z d dl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mZmZmZmZ d	gZej(                  hZg d
¢Z G d„ d	e j.                  «      Zy)é    N)Údefaultdict)ÚAnyÚOptional)Únn)Úparametrize)Útype_before_parametrizationsé   )ÚFakeSparsityÚget_arg_info_from_tensor_fqnÚmodule_contains_paramÚmodule_to_fqnÚswap_moduleÚBaseSparsifier)ÚmoduleÚ
module_fqnÚtensor_namec            
       ó*  ‡ — e Zd ZdZd!deeeef      fˆ fd„Zdeeef   fd„Z	deeeeef   f   ddfd„Z
d	„ Zdeeef   fd
„Zd"deeef   defd„Zefdej"                  deeej(                        ddfd„Zd„ Zd„ Z	 	 d#deeedf      deeeeedf   f      fd„Zddefdej"                  deeeej"                     eej"                     f      dedeej"                     fd„Zd"deddfd„Zej<                  dej"                  defd „«       Zˆ xZ S )$r   a'  Base class for all sparsifiers.

    Abstract methods that need to be implemented:

    - update_mask: Function to compute a new mask for all keys in the
        `groups`.

    Args:
        - model [nn.Module]: model to configure. The model itself is not saved
            but used for the state_dict saving / loading.
        - config [list]: configuration elements should be a dict map that includes
            `tensor_fqn` of tensors to sparsify
        - defaults [dict]: default configurations will be attached to the
            configuration. Only the keys that don't exist in the `config` will
            be updated.

    Example::

        >>> # xdoctest: +SKIP("Can't instantiate abstract class BaseSparsifier with abstract method update_mask")
        >>> config = [{'tensor_fqn': 'layer1.weight', 'tensor_fqn': 'linear2.weight2', 'sparsity_level': 0.5}]
        >>> defaults = {'sparsity_level': 0.7}
        >>> # model.layer1.weight will have `sparsity_level` = 0.7 (getting default)
        >>> sparsifier = BaseSparsifier(config, defaults)
    NÚdefaultsc                 ó|   •— t         ‰| �  «        |xs i | _        t        t        «      | _        g | _        d| _        y )NT)ÚsuperÚ__init__r   r   ÚdictÚstateÚgroupsÚenable_mask_update)Úselfr   Ú	__class__s     €úi/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/ao/pruning/sparsifier/base_sparsifier.pyr   zBaseSparsifier.__init__7   s4   ø€ Ü‰ÑÔØ(0ª°BˆŒä&1´$Ó&7ˆŒ
Ø,.ˆŒØ"&ˆÕó    Úreturnc                 óJ   — | j                   | j                  | j                  dœS )N©r   r   r   r"   )r   s    r   Ú__getstate__zBaseSparsifier.__getstate__?   s!   € àŸ™Ø—Z‘ZØ—k‘kñ
ð 	
r   r   c                 ó:   — | j                   j                  |«       y ©N)Ú__dict__Úupdate)r   r   s     r   Ú__setstate__zBaseSparsifier.__setstate__F   s   € Ø�‰×Ñ˜UÕ#r   c                 ó  — | j                   j                  dz   }t        | j                  «      D ]T  \  }}|d   }|dz  }|d|› d�z  }|d|› d�z  }t	        |j                  «       «      D ]  }|dk(  rŒ	|d|› d||   › d�z  }Œ ŒV |dz  }|S )	Nz (r   ú
z	Group z	    module: z	    z: ú))r   Ú__name__Ú	enumerater   ÚsortedÚkeys)r   Úformat_stringÚiÚsparse_argsr   Úkeys         r   Ú__repr__zBaseSparsifier.__repr__I   sÈ   € ØŸ™×/Ñ/°$Ñ6ˆÜ'¨¯©Ó4ò 	F‰NˆAˆ{Ø  Ñ*ˆFØ˜TÑ!ˆMØ˜x¨ s¨"Ð-Ñ-ˆMØ˜~¨f¨X°RÐ8Ñ8ˆMÜ˜k×.Ñ.Ó0Ó1ò F�Ø˜(’?ØØ 6¨#¨¨b°¸SÑ1AÐ0BÀ"Ð!EÑE‘ñFð	Fð 	˜ÑˆØÐr   c           
      ó    — | j                   D �cg c]&  }t        t        d„ |j                  «       «      «      ‘Œ( }}| j                  |dœS c c}w )aƒ  Returns the state of the optimizer as a :class:`dict`.

        It contains:
        * state - current state of the sparsification.
        * groups - a list containing all sparsity configuration groups
            with the key 'tensor_fqn' specifying the path to the sparsified tensor within a model

        TODO: Need a clean way of loading the state of the "prepared" module
        c                 ó   — | d   t         vS )Nr   )ÚKEYS_NOT_IN_STATE_DICT)Ú	key_values    r   ú<lambda>z+BaseSparsifier.state_dict.<locals>.<lambda>e   s   €  i°¡lÔ:PÐ&P€ r   ©r   r   )r   r   ÚfilterÚitemsr   )r   Úmgr   s      r   Ú
state_dictzBaseSparsifier.state_dictW   s[   € ð$ —k‘kö(
ð ô ÜÙPØ—H‘H“Jóõð(
ˆð (
ð —Z‘ZØñ
ð 	
ùò(
s   �+Ar>   Ústrictc           	      ó|  — t        j                  |d   «      }|d   }|j                  «       D ]ø  \  }}t        | j                  |«      }|d   }|d   }	|r|€t        d|› d�«      ‚d}
|j                  |	   D ]  }t        |t        «      sŒd}
 n |
sIt        t        j                  t        ||	«      j                  «      «      }t        j                  ||	|«       |j                  d	d «      �|j!                  d	«      }|_        |D ]  }|d
   |k(  sŒ|j%                  |«       Œ Œú | j'                  ||dœ«       y )Nr   r   r   r   zError loading z into the modelFTÚmaskÚ
tensor_fqnr:   )ÚcopyÚdeepcopyr<   r   ÚmodelÚRuntimeErrorÚparametrizationsÚ
isinstancer
   ÚtorchÚonesÚgetattrÚshaper   Úregister_parametrizationÚgetÚpoprA   r'   r(   )r   r>   r?   r   ÚstatesrB   ÚsÚarg_infor   r   ÚfoundÚprA   r=   s                 r   Úload_state_dictzBaseSparsifier.load_state_dictq   sA  € Ü—‘˜z¨(Ñ3Ó4ˆØ˜GÑ$ˆØ#Ÿ\™\›^ò 	(‰MˆJ˜Ü3°D·J±JÀ
ÓKˆHØ˜hÑ'ˆFØ" =Ñ1ˆKÙ˜&˜.Ü" ^°J°<¸Ð#OÓPÐPàˆEØ×,Ñ,¨[Ñ9ò �Ü˜a¤Õ.Ø �EÙðñ Ü ¤§¡¬G°F¸KÓ,H×,NÑ,NÓ!OÓP�Ü×4Ñ4°V¸[È!ÔLØ�u‰u�V˜TÓ"Ð.Ø—u‘u˜V“}�Ø�”àò (�Ø�lÑ# zÓ1Ø—I‘I˜hÕ'ñ(ð'	(ð, 	×Ñ F°fÑ=Õ>r   rE   ÚSUPPORTED_MODULESc                 ó.  — g | _         |g}|r‰|j                  «       }|j                  «       D ]b  \  }}t        |«      |v r?t	        ||«      }t        |t        «      sJ ‚| j                   j                  d|dz   i«       ŒR|j                  |«       Œd |rŒˆy y )NrB   z.weight)ÚconfigrO   Únamed_childrenÚtyper   rH   ÚstrÚappend)r   rE   rV   Ústackr   Ú_nameÚchildr   s           r   Úmake_config_from_modelz%BaseSparsifier.make_config_from_modelŒ   s’   € ð
 ˆŒØ�ˆÙØ—Y‘Y“[ˆFØ &× 5Ñ 5Ó 7ò (‘��uÜ˜“;Ð"3Ñ3Ü!.¨u°eÓ!<�JÜ% j´#Ô6Ð6Ð6Ø—K‘K×&Ñ&¨°jÀ9Ñ6LÐ'MÕNà—L‘L Õ'ð(ô r   c                 ó�  — || _         || _        | j                  €| j                  |«       | j                  D ]ü  }t        |t        «      sJ d«       ‚t        | j
                  t        «      sJ ‚t        j                  | j
                  «      }|j                  |«       |j                  dd«      }|€J d«       ‚t        ||«      }|j                  «       D ]1  }||v sŒ||   ||   k(  rŒ|dk(  rd||   z   ||   k(  rŒ(J d|› d�«       ‚ |j                  |«       | j                  j                  |«       Œþ | j                  «        y)zÁPrepares a model, by adding the parametrizations.

        Note::

            The model is modified inplace. If you need to preserve the original
            model, use copy.deepcopy.
        Nznconfig elements should be dicts not modules i.e.:[{`tensor_fqn`: `foo.bar.weight`}, {`tensor_fqn`: ... }, ...]rB   zttensor_fqn is a required argument in the sparsity config whichreplaces previous `module` and [module]`fqn` argumentsú.zGiven both `z?` and `tensor_fqn` in the config, it is expected them to agree!)rE   rX   r`   rH   r   r   rC   rD   r'   rN   r   r/   r   r\   Ú_prepare)r   rE   rX   Úmodule_configÚ
local_argsrB   Úinfo_from_tensor_fqnr3   s           r   ÚpreparezBaseSparsifier.prepare�   sq  € ð ˆŒ
ØˆŒð �;‰;ÐØ×'Ñ'¨Ô.ð "Ÿ[™[ò  	+ˆMÜ˜m¬TÔ2ð ðPóÐ2ô
 ˜dŸm™m¬TÔ2Ð2Ð2ÜŸ™ t§}¡}Ó5ˆJØ×Ñ˜mÔ,à#Ÿ™¨°dÓ;ˆJØÐ)ð ðIóÐ)ô $@ÀÀzÓ#RÐ ð ,×0Ñ0Ó2ò 	k�Ø˜*Ò$à,¨SÑ1°ZÀ±_ÓDà <Ò/Ø #Ð&:¸3Ñ&?Ñ ?À:ÈcÁ?Ó Rð	kð & c UÐ*iÐjókðð	kð ×ÑÐ2Ô3Ø�K‰K×Ñ˜zÕ*ðA 	+ðB 	�‰�r   c           
      ó(  — | j                   D ]ƒ  }|d   }|d   }|j                  dt        «      }|j                  dt        j                  t        ||«      «      «      }|| j                  |d      d<   t        j                  || ||«      «       Œ… y)z-Adds mask parametrization to the layer weightr   r   ÚparametrizationrA   rB   N)	r   rN   r
   rI   Ú	ones_likerK   r   r   rM   )r   ÚargsÚkwargsrX   r   r   ri   rA   s           r   rc   zBaseSparsifier._prepareÐ   sŒ   € à—k‘kò 	ˆFØ˜HÑ%ˆFØ  Ñ/ˆKØ$Ÿj™jÐ):¼LÓIˆOØ—:‘:˜f¤e§o¡o´g¸fÀkÓ6RÓ&SÓTˆDØ7;ˆD�J‰J�v˜lÑ+Ñ,¨VÑ4Ü×0Ñ0Ø˜¡_°TÓ%:õñ	r   Úparams_to_keep.Úparams_to_keep_per_layerc                 ó\  — | j                   D ]“  }|d   }|d   }t        j                  ||d¬«       i }|�$|D �	ci c]  }	|	||	   “Œ
 }
}	|j                  |
«       |�;|j	                  |d   d«      }|�$|D �	ci c]  }	|	||	   “Œ
 }}	|j                  |«       |sŒ�||_        Œ• yc c}	w c c}	w )a=	  Squashes the sparse masks into the appropriate tensors.

        If either the `params_to_keep` or `params_to_keep_per_layer` is set,
        the module will have a `sparse_params` dict attached to it.

        Args:
            params_to_keep: List of keys to save in the module or a dict
                            representing the modules and keys that will have
                            sparsity parameters saved
            params_to_keep_per_layer: Dict to specify the params that should be
                            saved for specific layers. The keys in the dict
                            should be the module fqn, while the values should
                            be a list of strings with the names of the variables
                            to save in the `sparse_params`

        Examples:
            >>> # xdoctest: +SKIP("locals are undefined")
            >>> # Don't save any sparse params
            >>> sparsifier.squash_mask()
            >>> hasattr(model.submodule1, 'sparse_params')
            False

            >>> # Keep sparse params per layer
            >>> sparsifier.squash_mask(
            ...     params_to_keep_per_layer={
            ...         'submodule1.linear1': ('foo', 'bar'),
            ...         'submodule2.linear42': ('baz',)
            ...     })
            >>> print(model.submodule1.linear1.sparse_params)
            {'foo': 42, 'bar': 24}
            >>> print(model.submodule2.linear42.sparse_params)
            {'baz': 0.1}

            >>> # Keep sparse params for all layers
            >>> sparsifier.squash_mask(params_to_keep=('foo', 'bar'))
            >>> print(model.submodule1.linear1.sparse_params)
            {'foo': 42, 'bar': 24}
            >>> print(model.submodule2.linear42.sparse_params)
            {'foo': 42, 'bar': 24}

            >>> # Keep some sparse params for all layers, and specific ones for
            >>> # some other layers
            >>> sparsifier.squash_mask(
            ...     params_to_keep=('foo', 'bar'),
            ...     params_to_keep_per_layer={
            ...         'submodule2.linear42': ('baz',)
            ...     })
            >>> print(model.submodule1.linear1.sparse_params)
            {'foo': 42, 'bar': 24}
            >>> print(model.submodule2.linear42.sparse_params)
            {'foo': 42, 'bar': 24, 'baz': 0.1}
        r   r   T)Úleave_parametrizedNr   )r   r   Úremove_parametrizationsr'   rN   Úsparse_params)r   rm   rn   rk   rl   rX   r   r   rr   ÚkÚglobal_paramsÚparamsÚper_layer_paramss                r   Úsquash_maskzBaseSparsifier.squash_maskÜ   sÜ   € ðv —k‘kò 	5ˆFØ˜HÑ%ˆFØ  Ñ/ˆKÜ×/Ñ/Ø˜¸õð ˆMØÐ)Ø7EÖ F°!  F¨1¡I¡Ð F�Ð FØ×$Ñ$ ]Ô3Ø'Ð3Ø1×5Ñ5°f¸\Ñ6JÈDÓQ�ØÐ%Ø>DÖ'E¸¨¨6°!©9©Ð'EÐ$Ð'EØ!×(Ñ(Ð)9Ô:Úà'4�Õ$ñ#	5ùò !Gùò
 (Fs   ºB$Á7B)Fr   ÚmappingÚinplaceÚparameterizationc                 óR  — |€t        d«      ‚|st        j                  |«      }i }|j                  «       D ]F  \  }}t	        ||«      rt        |«      |v rt        ||«      ||<   Œ/| j                  ||d|¬«      ||<   ŒH |j                  «       D ]  \  }}	|	|j                  |<   Œ |S )aí  Converts submodules in input module to a different module according to `mapping`
        by calling `from_dense` method on the target module class
        Args:
            module: input module
            mapping: a dictionary that maps from source module type to target
                module type, can be overwritten to allow swapping user defined
                Modules
            inplace: carry out model transformations in-place, the original module
                is mutated
        zNeed to auto generate mapping T)rx   ry   rz   )
ÚNotImplementedErrorrC   rD   rY   r   r   r   Úconvertr<   Ú_modules)
r   r   rx   ry   rz   ÚreassignÚnameÚmodr3   Úvalues
             r   r}   zBaseSparsifier.convert*  sÇ   € ð" ˆ?Ü%Ð&FÓGÐGÙÜ—]‘] 6Ó*ˆFàˆØ×.Ñ.Ó0ò 	‰IˆD�#ô & cÐ+;Ô<Ü0°Ó5¸Ñ@ä!,¨S°'Ó!:�˜’ð "&§¡ØØ#Ø Ø%5ð	 ".ó "�˜’ð	ð  #Ÿ.™.Ó*ò 	)‰JˆC�Ø#(ˆF�O‰O˜CÒ ð	)ð ˆr   Úuse_pathc                 ó¸   — | j                   sy t        j                  «       5  | j                  D ]  } | j                  di |¤Ž Œ 	 d d d «       y # 1 sw Y   y xY w)N© )r   rI   Úno_gradr   Úupdate_mask)r   rƒ   rX   s      r   ÚstepzBaseSparsifier.stepV  sR   € Ø×&Ò&ØÜ�]‰]‹_ñ 	+ØŸ+™+ò +�Ø �× Ñ Ñ* 6Ó*ñ+÷	+÷ 	+ñ 	+ús   ¢$AÁAr   c                  ó   — y r%   r…   )r   r   r   rl   s       r   r‡   zBaseSparsifier.update_mask]  s   € àr   r%   )T)NN)!r,   Ú
__module__Ú__qualname__Ú__doc__r   r   r[   r   r   r#   r(   r4   r>   ÚboolrU   rV   r   ÚModuleÚsetrZ   ÚLinearr`   rg   rc   Útuplerw   r
   r}   rˆ   ÚabcÚabstractmethodr‡   Ú__classcell__)r   s   @r   r   r      s»  ø„ ññ2' ¨$¨s°C¨x©.Ñ!9õ 'ð
˜d 3¨ 8™nó 
ð$ $ s¨D°°c°©NÐ':Ñ";ð $Àó $òð
˜D  c ™Nó 
ñ4?¨$¨s°C¨x©.ð ?À$ó ?ð< 3Dñ(à�y‰yð(ð ˜t B§I¡I™Ñ/ð(ð 
ó	(ò"1òf
ð 59ØIMñL5à   s¨C x¡Ñ1ðL5ð #+¨4°°U¸3À¸8±_Ð0DÑ+EÑ"FóL5ðb EIØØ,8ñ*à—	‘	ð*ð ˜$˜t B§I¡I™°°R·Y±Y±Ð?Ñ@ÑAð*ð ð	*ð
 ˜rŸy™y™/ó*ñX+˜Tð +¨Tó +ð 	×Ñð "§)¡)ð ¸#ò ó ôr   )r’   rC   Úcollectionsr   Útypingr   r   rI   r   Útorch.nn.utilsr   Útorch.nn.utils.parametrizer   Úutilsr
   r   r   r   r   Ú__all__r�   rV   r7   ÚABCr   r…   r   r   ú<module>rœ      sU   ðã 
Û Ý #ß  ã Ý Ý &Ý C÷õ ð Ð
€à—Y‘Y�KÐ â@Ð ôB�S—W‘Wõ Br   