Ë
    Ž\;j§W  ã                   ó  — d 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Z	g Z
 G d„ de«      Z G d„ de«      Zd	„ Zd
„ Zd„ Zd„ Zd„ Zd„ Zd„ Z ej(                  «       ai ad„ Zd„ Zej2                  ddfd„Zej6                  ddfd„Zy)z#
Utilities of Auto SParsity (ASP).
é    N)ÚEnum)Úpermutationsc                   ó   — e Zd ZdZdZdZdZy)ÚMaskAlgoz’
    A collection of all mask generating algorithms.
    There currently are three algorithms, `MASK_1D`, `MASK_2D_GREEDY` and `MASK_2D_BEST`
    Úget_mask_1dÚget_mask_2d_greedyÚget_mask_2d_bestN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚMASK_1DÚMASK_2D_GREEDYÚMASK_2D_BEST© ó    úbG:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/incubate/asp/utils.pyr   r      s   „ ñð €GØ)€NØ%�Lr   r   c                   ó(   — e Zd ZdZdZdZed„ «       Zy)ÚCheckMethodzz
    A collection of all sparsity checking approaches.
    There currently are two methods, `CHECK_1D` and `CHECK_2D`
    Úcheck_mask_1dÚcheck_mask_2dc                 ó–   — t        | t        «      sJ d«       ‚| t        j                  k(  rt        j                  S t        j
                  S )aý  
        Get sparsity checking method by mask generating algorithm.

        Args:
            mask_algo (MaskAlgo): The algorithm of mask generating.
        Returns:
            CheckMethod: The corresponded sparsity checking method.
        Examples:
            .. code-block:: python

                >>> import numpy as np
                >>> from paddle.incubate.asp import CheckMethod, MaskAlgo
                >>> print(CheckMethod.get_checking_method(MaskAlgo.MASK_1D))
                CheckMethod.CHECK_1D
                >>> print(CheckMethod.get_checking_method(MaskAlgo.MASK_2D_GREEDY))
                CheckMethod.CHECK_2D
                >>> print(CheckMethod.get_checking_method(MaskAlgo.MASK_2D_BEST))
                CheckMethod.CHECK_2D
        z!mask_algo should be MaskAlgo type)Ú
isinstancer   r   r   ÚCHECK_1DÚCHECK_2D)Ú	mask_algos    r   Úget_checking_methodzCheckMethod.get_checking_method0   sK   € ô* Ø”xô
ð 	/à.ó	/ð 
ð œ×(Ñ(Ò(Ü×'Ñ'Ð'ä×'Ñ'Ð'r   N)r
   r   r   r   r   r   Ústaticmethodr   r   r   r   r   r   (   s%   „ ñð €HØ€Hàñ(ó ñ(r   r   c                 ó’   — | j                  «       }t        t        j                  |«      d   j                  «      |j                  z  S )aÐ  

    Return the density of the input tensor.

    Args:
        x (nparray): The input tensor.

    Returns:
        float, The density of :attr:`x`.

    Examples:
        .. code-block:: python

            >>> import paddle
            >>> import numpy as np

            >>> x = np.array([[0, 1, 3, 0],
            ...             [1, 1, 0, 1]])
            >>> out = paddle.incubate.asp.calculate_density(x)
            >>> print(out)
            0.625

    r   )ÚflattenÚfloatÚnpÚnonzeroÚsize)ÚxÚx_flatteneds     r   Úcalculate_densityr'   N   s9   € ð0 —)‘)“+€KÜ”—‘˜KÓ(¨Ñ+×0Ñ0Ó1°K×4DÑ4DÑDÐDr   c                 ó¨  — t        | j                  «      dk(  sJ d«       ‚| j                  d   |z  }| j                  d   |z  dkD  rot        j                  | j                  d   | j                  d   ||z
  z   f«      }| |dd…d| j                  d   …f<   |j                  }|j	                  d|«      |fS | j	                  d|«      | j                  fS )aæ  
    Reshape the input 2D matrix to shape (-1, m).
    If the second dimension of :attr:`mat` is not a multiples of :attr:`m`,
    then this function would pad the remainder with 0 before reshaping.

    .. math::

        remainder = mat.shape[1] % m

    Args:
        mat (nparray): The input 2D matrix.
        m (int): The second dimension of reshaped matrix.
    Returns:
        tuple: A pair of the reshaped and padded matrix and the shape of padded matrix (non-reshaping).
    é   ú$The input mat should be a 2D matrix!é   r   Néÿÿÿÿ)ÚlenÚshaper"   ÚzerosÚreshape)ÚmatÚmÚ	remainderÚ
mat_paddedr.   s        r   Ú_reshape_1dr5   j   sË   € ô  ˆs�y‰y‹>˜QÒÐFÐ FÓFÐà—	‘	˜!‘˜qÑ €IØ
‡y�y��|�aÑ˜!ÒÜ—X‘X˜sŸy™y¨™|¨S¯Y©Y°q©\¸QÀ¹]Ñ-KÐLÓMˆ
Ø(+ˆ
’1�n˜Ÿ	™	 !™�nÐ$Ñ%Ø× Ñ ˆØ×!Ñ! " aÓ(¨%Ð/Ð/à�{‰{˜2˜qÓ! 3§9¡9Ð,Ð,r   c                 ó  — t        | j                  «      dk  r-t        | j                  d| j                  d   «      |«      \  }}nt        | |«      \  }}|D ],  }t	        j
                  |«      d   j                  ||z
  kD  sŒ, y y)aæ  
    Check if every row of the input matrix :attr:`mat` is in 1D `n:m` sparse pattern.
    This function would pad the second dimension of :attr:`mat` by zero
    to be a multiples of :attr:`m` if necessary.

    1D `n:m` sparse pattern: At least :attr:`n` zeros in every :math:`1 \times m` block.

    Args:
        mat (nparray): The input matrix.
        n (int): n of `n:m` sparse pattern.
        m (int): m of `n:m` sparse pattern.
    Returns:
        bool: True if every row of :attr:`mat` is in 1D n:m sparse pattern, else False.
    Examples:
        .. code-block:: python

          >>> import numpy as np
          >>> import paddle.incubate.asp as sparsity

          >>> x = np.array([[0, 1, 3, 0],
          ...               [1, 0, 0, 1]])
          >>> y = sparsity.check_mask_1d(x, 2, 4)
          >>> print(y)
          True

          >>> x = np.array([[0, 1, 5, 4],
          ...               [1, 0, 0, 1]])
          >>> y = sparsity.check_mask_1d(x, 2, 4)
          >>> print(y)
          False

          >>> # x would be padded to shape (2, 8)
          >>> x = np.array([[0, 1, 0, 4, 6],
          ...               [1, 0, 0, 1, 7]])
          >>> y = sparsity.check_mask_1d(x, 2, 4)
          >>> print(y)
          True
    r+   r   FT)r-   r.   r5   r0   r"   r#   r$   )r1   Únr2   Úmat_flatternr.   Úsub_mats         r   r   r   †   s|   € ôN ˆ3�9‰9ƒ~˜ÒÜ)¨#¯+©+°a¸¿¹À1¹Ó*FÈÓJÑˆ‘eä)¨#¨qÓ1Ñˆ�eãˆÜ�:‰:�gÓ˜qÑ!×&Ñ&¨!¨a©%Ó0Ùð  ð r   c                 ó   — t        | |«      \  }}t        j                  |«      }t        j                  | «      }t        |j                  d   «      D ]G  }||   }t        j
                  t        j                  |«      «      }	d|||	d| j                  «       f<   ŒI |j                  |«      }|dd…d| j                  d   …f   |dd…dd…f<   |S )aÂ  
    Generate 1D `n:m` sparse pattern mask of the input matrix :attr:`mat`
    in row-directory. This function would pad the second dimension of :attr:`mat`
    by zero to be a multiples of :attr:`m` before mask generation.

    1D `n:m` sparse pattern: At least :attr:`n` zeros in every :math:`1 \times m` block.

    Args:
        mat (nparray): The input matrix.
        n (int): n of `n:m` sparse pattern.
        m (int): m of `n:m` sparse pattern.
    Returns:
        nparray: The 1D `n:m` sparse mask of :attr:`mat`.
    Examples:
        .. code-block:: python

          >>> import numpy as np
          >>> import paddle.incubate.asp as sparsity
          >>> mat = np.array([[0, 1, 5, 4],
          ...                 [2, 7, 3, 6]])
          >>> mask = sparsity.get_mask_1d(mat, 2, 4)
          >>> print(mask)
          [[0 0 1 1]
          [0 1 0 1]]
          >>> y = sparsity.check_mask_1d(mask, 2, 4)
          >>> print(y)
          True
    r   Nr+   )	r5   r"   Ú	ones_likeÚranger.   ÚargsortÚabsoluteÚtolistr0   )
r1   r7   r2   r8   r.   Úmask_flatternÚmaskÚir9   Úmin_order_indicess
             r   r   r   ¸   sÆ   € ô: & c¨1Ó-Ñ€L�%ä—L‘L Ó.€MÜ�<‰<˜Ó€DÜ�<×%Ñ% aÑ(Ö)ˆØ˜q‘/ˆÜŸJ™J¤r§{¡{°7Ó';Ó<ÐØ;<ˆ�aÐ*¨2¨AÐ.×5Ñ5Ó7Ð7Ò8ð *ð "×)Ñ)¨%Ó0€MØšq . C§I¡I¨a¡L .Ð0Ñ1€DŠŠAˆ�JØ€Kr   c                 ó  — t        | j                  «      dk(  sJ d«       ‚| j                  d   |z  }| j                  d   |z  }|dk(  r| j                  d   n| j                  d   ||z
  z   |dk(  r| j                  d   n| j                  d   ||z
  z   f}t        j                  |«      }| |d| j                  d   …d| j                  d   …f<   t        j                  |«      j                  d||z  «      }d}t        d|j                  d   |«      D ]b  }||z   }	t        d|j                  d   |«      D ]>  }
|
|z   }t        j                  |||	…|
|…f   j                  d«      «      }|||<   |dz  }Œ@ Œd ||j                  fS )a3  
    Reshape the input 2D matrix to shape (-1, :math:`m \times m`).
    In each dimension of :attr:`mat`, if it is not a multiples of :attr:`m`,
    then this function would pad the remainder with 0 before reshaping.

    .. math::

        remainder_0 = mat.shape[0] % m \\
        remainder_1 = mat.shape[1] % m

    Args:
        mat (nparray): The input 2D matrix.
        m (int): The square root of second dimension of reshaped matrix.
    Returns:
        tuple: A pair of the reshaped and padded matrix and the shape of padded matrix (non-reshaping).
    r)   r*   r   r+   Nr,   )r-   r.   r"   r/   Úemptyr0   r<   Úsqueeze)r1   r2   Úremainder_0Úremainder_1Ú	new_shaper4   r8   Úcurr_idxÚ	row_startÚrow_endÚ	col_startÚcol_endr9   s                r   Ú_reshape_2drO   â   s˜  € ô" ˆs�y‰y‹>˜QÒÐFÐ FÓFÐà—)‘)˜A‘, Ñ"€KØ—)‘)˜A‘, Ñ"€Kð $ qÒ(ˆ�	‰	�!Š¨c¯i©i¸©l¸aÀ+¹oÑ.NØ# qÒ(ˆ�	‰	�!Š¨c¯i©i¸©l¸aÀ+¹oÑ.Nð€Iô —‘˜)Ó$€JØ14€Jˆ~�—‘˜1‘ˆ~˜~ §¡¨1¡˜~Ð-Ñ.ä—8‘8˜IÓ&×.Ñ.¨r°1°q±5Ó9€LØ€HÜ˜1˜j×.Ñ.¨qÑ1°1Ö5ˆ	Ø˜a‘-ˆÜ˜q *×"2Ñ"2°1Ñ"5°qÖ9ˆIØ !‘mˆGÜ—j‘jØ˜9 WÐ,¨i¸Ð.?Ð?Ñ@×HÑHÈÓLóˆGð &-ˆL˜Ñ"Ø˜‰M‰Hñ :ð 6ð ˜×)Ñ)Ð)Ð)r   c           	      óx  — t        | |«      \  }}|D ]¦  }t        j                  t        j                  |j	                  ||«      «      «      dkD  }t        j
                  t        j
                  |d¬«      ||z
  kD  «      dk7  sŒrt        j
                  t        j
                  |d¬«      ||z
  kD  «      dk7  sŒ¦ y y)a•  
    Check if every :math:`m \times m` block of the input matrix :attr:`mat` is in 2D `n:m` sparse pattern.
    This function would pad each dimension of :attr:`mat` by zero to be a multiples of
    :attr:`m` if necessary.

    2D `n:m` sparse pattern: At least :math:`n \times n` zeros in every :math:`m \times m` block
    under the constraint of at least :attr:`n` zeros for each row and column.

    Args:
        mat (nparray): The input matrix.
        n (int): n of `n:m` sparse pattern.
        m (int): m of `n:m` sparse pattern.
    Returns:
        bool: True if  every :math:`m \times m` block of the input matrix :attr:`mat` is in 2D `n:m` sparse pattern, else False.
    Examples:
        .. code-block:: python

          >>> import numpy as np
          >>> import paddle.incubate.asp as sparsity

          >>> x = np.array([[0, 8, 9, 0],
          ...               [9, 0, 0, 10],
          ...               [5, 0, 0, 6],
          ...               [0, 4, 6, 0]])
          >>> y = sparsity.check_mask_2d(x, 2, 4)
          >>> print(y)
          True

          >>> x = np.array([[0, 8, 0, 9],
          ...               [9, 0, 0, 10],
          ...               [0, 5, 0, 6],
          ...               [0, 4, 6, 0]])
          >>> y = sparsity.check_mask_2d(x, 2, 4)
          >>> print(y)
          True

          >>> # x would be padded to shape (8, 8)
          >>> x = np.array([[0, 8, 0, 9],
          ...               [9, 0, 7, 0],
          ...               [0, 5, 0, 6],
          ...               [3, 0, 6, 0],
          ...               [1, 1, 0, 1]])
          >>> y = sparsity.check_mask_2d(x, 2, 4)
          >>> print(y)
          True
    r   r+   ©ÚaxisFT)rO   r"   r>   rF   r0   Úsum)r1   r7   r2   r4   r.   r9   Úsub_masks          r   r   r     s™   € ô^ $ C¨Ó+Ñ€J�ÛˆÜ—;‘;œrŸz™z¨'¯/©/¸!¸QÓ*?Ó@ÓAÀAÑEˆÜ�F‰F”2—6‘6˜(¨Ô+¨q°1©uÑ5Ó6¸!Ó;Ü�F‰F”2—6‘6˜(¨Ô+¨q°1©uÑ5Ó6¸!Ó;áð ð r   c                 óÀ  — t        | |«      \  }}t        j                  |«      j                  d||«      }t	        t        |«      «      D �]
  }t        j                  t        j                  ||   «      «      }t        j                  ||   «      }t        j                  |«      }	|	D �
cg c]  }
t        |
|z  «      |
|z  f‘Œ }}
t        j                  «       }t        j                  «       }t	        t        |	«      dz
  dd«      D ]K  }||   }||d      |k(  s||d      |k(  rŒd||d   |d   f<   ||d   xx   dz  cc<   ||d   xx   dz  cc<   ŒM �Œ t        j                  |«      }d}t	        d|d   |«      D ]4  }||z   }t	        d|d   |«      D ]  }||z   }||   |||…||…f<   |dz  }Œ Œ6 |d| j                  d   …d| j                  d   …f   S c c}
w )a  
    Greedily generate 2D `n:m` sparse pattern mask of the input matrix :attr:`mat`.
    This function would pad each dimension of :attr:`mat` by zero to be a multiples of :attr:`m` before mask generation.

    2D `n:m` sparse pattern: At least :math:`n \times n` zeros in every :math:`m \times m` block
    under the constraint of at least :attr:`n` zeros for each row and column.
    Greedily generating: For each :math:`m \times m` block, selecting values to keep in descent order.

    Args:
        mat (nparray): The input matrix.
        n (int): n of `n:m` sparse pattern.
        m (int): m of `n:m` sparse pattern.
    Returns:
        nparray: The 2D `n:m` sparse mask of :attr:`mat`.
    Examples:
        .. code-block:: python

          >>> import numpy as np
          >>> import paddle.incubate.asp as sparsity

          >>> mat = np.array([[9, 8, 3, 7],
          ...                 [9, 2, 1, 10],
          ...                 [5, 1, 3, 6],
          ...                 [2, 4, 6, 1]])
          >>> mask = sparsity.get_mask_2d_greedy(mat, 2, 4)
          >>> print(mask)
          [[1. 1. 0. 0.]
          [1. 0. 0. 1.]
          [0. 0. 1. 1.]
          [0. 1. 1. 0.]]
          >>> y = sparsity.check_mask_2d(mask, 2, 4)
          >>> print(y)
          True
    r,   r+   r   g      ð?N)rO   r"   Ú
zeros_liker0   r<   r-   r>   rF   r=   ÚintÚcollectionsÚCounterrE   r.   )r1   r7   r2   r4   r.   Úmask_paddedÚidxr9   rT   Úmin_order_1d_indicesr%   Úmin_order_2d_indicesÚrow_counterÚcol_counterrB   Úmatrix_entryrA   rJ   rK   rL   rM   rN   s                         r   r   r   F  s  € ôF $ C¨Ó+Ñ€J�Ü—-‘- 
Ó+×3Ñ3°B¸¸1Ó=€Kä”S˜“_×%ˆÜ—+‘+œbŸj™j¨°C©Ó9Ó:ˆÜ—:‘:˜k¨#Ñ.Ó/ˆä!Ÿz™z¨'Ó2Ðá)=ó 
Ù)= AŒS��Q‘‹Z˜˜Q™ÒÐ)=ð 	ð  
ô "×)Ñ)Ó+ˆÜ!×)Ñ)Ó+ˆä”sÐ/Ó0°1Ñ4°b¸"Ö=ˆAØ/°Ñ2ˆLØ˜L¨™OÑ,°Ò1Ø˜L¨™OÑ,°Ò1àà9<ˆH�\ !‘_ l°1¡oÐ5Ñ6Ø˜ Q™Ó(¨AÑ-Ó(Ø˜ Q™Ó(¨AÑ-Ô(ò >ð &ô, �8‰8�E‹?€DØ€HÜ˜1˜e A™h¨Ö*ˆ	Ø˜a‘-ˆÜ˜q %¨¡(¨AÖ.ˆIØ !‘mˆGØ9DÀXÑ9NˆD�˜7Ð" I¨gÐ$5Ð5Ñ6Ø˜‰M‰Hñ /ð +ð ��#—)‘)˜A‘,�  #§)¡)¨A¡, Ð.Ñ/Ð/ùò3 
s   Â*Gc           
      ó~  — |› d| › �}|t         v r	t         |   S t        j                  |«      }d|d|  t        t	        t        |j                  «       «      «      «      }||z   }t        j                  t        t	        t        ||«      «      «      «      }|j                  d¬«      | k  j                  d¬«      |k(  j                  «       d   j                  d«      }t        j                  |j                  d   ||f«      }||dd    |dd t        j                  «        |t         |<   t        j                  «        |S )a¾  
    Compute all vaild 2D `n:m` sparse patterns.

    2D `n:m` sparse pattern: At least :math:`n \times n` zeros in every :math:`m \times m` block
    under the constraint of at least :attr:`n` zeros for each row and column.

    Args:
        n (int): n of `n:m` sparse pattern.
        m (int): m of `n:m` sparse pattern.
    Returns:
        dictionary: A dictionary with key: *m_n* (string) and value: all vaild 2D `n:m` sparse patterns.
    Ú_r+   NrQ   r   r,   )Ú_valid_2d_patternsr"   r/   ÚlistÚsetr   r?   ÚasarrayrS   r#   r0   rE   r.   Ú_valid_2d_patterns_lockÚacquireÚrelease)r7   r2   Ú	valid_keyÚpatternsÚvalidÚvalid_patternss         r   Ú_compute_valid_2d_patternsrn   ‘  s"  € ð  �#�Q�q�c�
€IØÔ&Ñ&Ü! )Ñ,Ð,ä—8‘8˜A“;ˆØˆ��!ˆÜœœL¨¯©Ó):Ó;Ó<Ó=ˆØ˜hÑ&ˆÜ—:‘:œd¤3¤|°H¸aÓ'@Ó#AÓBÓCˆð �l‰l ˆlÓ" aÑ'×,Ñ,°!Ð,Ó4¸Ñ9ß‰W‹Y�qñç‰W�R‹[ð 	ô
 Ÿ™ 5§;¡;¨q¡>°1°aÐ"8Ó9ˆØ$ U©1 XÑ.ˆ‘qÐä×'Ñ'Ô)Ø(6Ô˜9Ñ%Ü×'Ñ'Ô)àÐr   c           
      óJ  — t        ||«      }t        | |«      \  }}t        j                  |«      j	                  d||«      }t        j
                  t        j                  ||j	                  |j                  d   ||z  «      j                  «      d¬«      }||dd    |dd t        j                  |«      }d}	t        d|d   |«      D ]4  }
|
|z   }t        d|d   |«      D ]  }||z   }||	   ||
|…||…f<   |	dz  }	Œ Œ6 |d| j                  d   …d| j                  d   …f   S )a´  
    Generate 2D `n:m` sparse pattern mask of the input matrix :attr:`mat`
    to form sparse matrix with maximun L1 norm .This function would pad each
    dimension of :attr:`mat` by zero to be a multiples of :attr:`m` before mask generation.

    2D `n:m` sparse pattern: At least :math:`n \times n` zeros in every :math:`m \times m` block
    under the constraint of at least :attr:`n` zeros for each row and column.

    *Note*: L1 norm of sparse matrix from `Best` API is greater than or equal to the one from `Greedy`.

    Args:
        mat (nparray): The input matrix.
        n (int): n of `n:m` sparse pattern.
        m (int): m of `n:m` sparse pattern.
    Returns:
        nparray: The 1D `n:m` sparse mask of :attr:`mat`.
    Examples:
        .. code-block:: python

          >>> import numpy as np
          >>> import paddle.incubate.asp as sparsity

          >>> mat = np.array([[2, 8, 9, 9],
          ...                 [9, 1, 3, 9],
          ...                 [5, 6, 3, 9],
          ...                 [2, 4, 6, 9]])
          >>> mask_greedy = sparsity.get_mask_2d_greedy(mat, 2, 4)
          >>> mask_best = sparsity.get_mask_2d_best(mat, 2, 4)
          >>> print("L1 norm of `greedy` sparse matrix", np.multiply(mat, mask_greedy).sum())
          L1 norm of `greedy` sparse matrix 56.0
          >>> print("L1 norm of `best` sparse matrix", np.multiply(mat, mask_best).sum())
          L1 norm of `best` sparse matrix 61.0
    r,   r   r+   rQ   N)rn   rO   r"   r;   r0   ÚargmaxÚmatmulr.   ÚTrE   r<   )r1   r7   r2   rk   r8   r.   r@   ÚpmaxrA   rJ   rK   rL   rM   rN   s                 r   r	   r	   º  s6  € ôD *¨!¨QÓ/€Hä% c¨1Ó-Ñ€L�%Ü—L‘L Ó.×6Ñ6°r¸1¸aÓ@€MÜ�9‰9Ü
�	‰	�, × 0Ñ 0°·±ÀÑ1BÀAÈÁEÓ J× LÑ LÓMØô€Dð
   ¡Q Ñ(€M‘!ÐÜ�8‰8�E‹?€Dà€HÜ˜1˜e A™h¨Ö*ˆ	Ø˜a‘-ˆÜ˜q %¨¡(¨AÖ.ˆIØ !‘mˆGØ9FÀxÑ9PˆD�˜7Ð" I¨gÐ$5Ð5Ñ6Ø˜‰M‰Hñ /ð +ð ��#—)‘)˜A‘,�  #§)¡)¨A¡, Ð.Ñ/Ð/r   r)   é   c                 óŒ  — | j                   }| j                  }| j                  t        «      }t	        |t
        «      sJ dt        |«      › �«       ‚t        t        j                  t           |j                  d«      }t        |«      dk(  r|j                  d|d   «      }nút        |«      dk(  r|j                  |d   |d   «      }nÓt        |«      dk(  r|j                  |d   |d   z  |d   «      }n¦t        |«      dk(  r�|j                  g d¢«      j                  |d   |d   z  |d   z  |d   «      } ||||¬	«      }|j                  |d   |d   |d   |d   g«      j                  g d¢«      j                  |«      S t        d
t        |«      › �«      ‚ ||||¬	«      }|j                  |«      j                  |«      S )aH  
    Create `n:m` sparse pattern mask of the input tensor via function given by :attr:`func_name`.
    Currently only support tensor with dimension less than or equal to 4.

    Args:
        tensor (nparray): The input tensor.
        func_name (MaskAlgo, optional): The function name to generate spase mask. Default is `MaskAlgo.MASK_1D`. All options please refer to `MaskAlgo`.
        n (int, optional): n of `n:m` sparse pattern. Default is 2.
        m (int, optional): m of `n:m` sparse pattern. Default is 4.
    Returns:
        nparray: The `n:m` sparse mask of :attr:`tensor` generated by :attr:`func_name`.
    Examples:
        .. code-block:: python

          >>> import numpy as np
          >>> import paddle.incubate.asp as sparsity

          >>> tensor = np.array([[2, 8, 9, 9],
          ...                    [9, 1, 3, 9],
          ...                    [5, 6, 3, 9],
          ...                    [2, 4, 6, 9]])
          >>> mask_1d = sparsity.create_mask(tensor, func_name=sparsity.MaskAlgo.MASK_1D)
          >>> print(mask_1d)
          [[0 0 1 1]
          [1 0 0 1]
          [0 1 0 1]
          [0 0 1 1]]
          >>> mask_2d = sparsity.create_mask(tensor, func_name=sparsity.MaskAlgo.MASK_2D_BEST)
          >>> print(mask_2d)
          [[0 1 1 0]
          [1 0 0 1]
          [1 1 0 0]
          [0 0 1 1]]
    zLfunc_name argumet of create_mask is only accepted as type MaskAlgo. But got Nr+   r   r)   é   rt   ©r   r+   rv   r)   ©r7   r2   úgThe dimension of input tensor is not supported in create_mask, Only dimension < 4 is supported but got )r.   ÚdtypeÚastyper!   r   r   ÚtypeÚgetattrÚsysÚmodulesr
   Úvaluer-   r0   Ú	transposeÚ
ValueError)	ÚtensorÚ	func_namer7   r2   r.   rz   ÚtÚfuncrA   s	            r   Úcreate_maskr‡   ò  s¸  € ðF �L‰L€EØ�L‰L€EØ�‰”eÓ€Aä�i¤Ô*ð ð	Ü˜	“?Ð#ð	%óÐ*ô ”3—;‘;œxÑ(¨)¯/©/¸4Ó@€DÜ
ˆ5ƒz�Q‚Ø�I‰I�a˜˜q™Ó"‰Ü	ˆU‹�qŠØ�I‰I�e˜A‘h  a¡Ó)‰Ü	ˆU‹�qŠØ�I‰I�e˜A‘h  q¡Ñ)¨5°©8Ó4‰ä	ˆU‹�qŠØ�K‰KšÓ%×-Ñ-Ø�!‰H�u˜Q‘xÑ %¨¡(Ñ*¨E°!©Hó
ˆñ �A˜˜aÔ ˆà�L‰L˜% ™( E¨!¡H¨e°A©h¸¸a¹ÐAÓBß‰Y’|Ó$ß‰V�E‹]ð	
ô ð7Ü7:¸5³z°lðDó
ð 	
ñ
 ��Q˜!Ô€DØ�<‰<˜Ó×%Ñ% eÓ,Ð,r   c                 ó¤  — | j                   }| j                  t        «      }t        |«      t        k(  sJ dt        |«      › �«       ‚t        t        j                  t           |j                  d«      }t        |«      dk(  r|j                  d|d   «      }n°t        |«      dk(  r|j                  |d   |d   «      }n‰t        |«      dk(  r|j                  |d   |d   z  |d   «      }n\t        |«      dk(  r7|j                  g d¢«      j                  |d   |d   z  |d   z  |d   g«      }nt        d	t        |«      › �«      ‚ ||||¬
«      S )aŠ  
    Check if input tensor is in `n:m` sparse pattern via function given by :attr:`func_name`.
    Currently only support tensor with dimension less than or equal to 4.

    Args:
        tensor (nparray): The input tensor.
        func_name (CheckMethod, optional): The function name to generate spase mask. Default is `CheckMethod.CHECK_1D`. All options please refer to `CheckMethod`.
        n (int, optional): n of `n:m` sparse pattern. Default is 2.
        m (int, optional): m of `n:m` sparse pattern. Default is 4.
    Returns:
        bool: True if tensor pass checking of function given by :attr:`func_name`, else False.
    Examples:
        .. code-block:: python

          >>> import numpy as np
          >>> import paddle.incubate.asp as sparsity

          >>> tensor = np.array([[2, 8, 9, 9],
          ...                    [9, 1, 3, 9],
          ...                    [5, 6, 3, 9],
          ...                    [2, 4, 6, 9]])
          >>> mask_1d = sparsity.create_mask(tensor, func_name=sparsity.MaskAlgo.MASK_1D)
          >>> print(mask_1d)
          [[0 0 1 1]
          [1 0 0 1]
          [0 1 0 1]
          [0 0 1 1]]
          >>> y = sparsity.check_sparsity(mask_1d, func_name=sparsity.CheckMethod.CHECK_1D)
          >>> print(y)
          True
          >>> y = sparsity.check_sparsity(mask_1d, func_name=sparsity.CheckMethod.CHECK_2D)
          >>> print(y)
          True
    zRfunc_name argumet of check_sparsity is only accepted as type CheckMethod. But got Nr+   r   r)   rv   rt   rw   ry   rx   )r.   r{   r!   r|   r   r}   r~   r   r
   r€   r-   r0   r�   r‚   )rƒ   r„   r7   r2   r.   r…   r†   s          r   Úcheck_sparsityr‰   9  sN  € ðF �L‰L€EØ�‰”eÓ€Aä�	‹?œkÒ)ð ð	Ü˜	“?Ð#ð	%óÐ)ô ”3—;‘;œxÑ(¨)¯/©/¸4Ó@€DÜ
ˆ5ƒz�Q‚Ø�I‰I�a˜˜q™Ó"‰Ü	ˆU‹�qŠØ�I‰I�e˜A‘h  a¡Ó)‰Ü	ˆU‹�qŠØ�I‰I�e˜A‘h  q¡Ñ)¨5°©8Ó4‰ä	ˆU‹�qŠØ�K‰KšÓ%×-Ñ-Ø�1‰X˜˜a™Ñ  5¨¡8Ñ+¨U°1©XÐ6ó
‰ô ð7Ü7:¸5³z°lðDó
ð 	
ñ
 ��Q˜!ÔÐr   )r   rX   r~   Ú	threadingÚenumr   Ú	itertoolsr   Únumpyr"   Ú__all__r   r   r'   r5   r   r   rO   r   r   ÚLockrg   rc   rn   r	   r   r‡   r   r‰   r   r   r   Ú<module>r�      s³   ðñó Û 
Û Ý Ý "ã à
€ô&ˆtô &ô#(�$ô #(òLEò8-ò8/òd'òT(*òV6òrD0ðN )˜)Ÿ.™.Ó*Ð ØÐ ò&òR50ðp #+×"2Ñ"2°a¸1ó D-ðN &1×%9Ñ%9¸QÀ!ô <r   