Ë
    •\;j®,  ã                   ó‚   — 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 d	d
lmZ g Z G d„ de«      Zy)é    )Ú_C_opsÚin_dynamic_mode)Úunique_nameé   )Úbase)Ú	framework)ÚVarDesc)Úcheck_variable_and_dtype)Ú_current_expected_placeé   )ÚInitializerc                   ó,   ‡ — e Zd ZdZdˆ fd„	Zdd„Zˆ xZS )ÚDiraca	  Initialize the 3D/4D/5D Tensor with Dirac delta function.

    It can reserve the feature of convolution layer input, which means that
    as many channels are reserved as possible.

    In this initialize method, elements in the middle of convolution kernels will
    be set to 1 . The formula can be described as follow.

    .. math::

        X[d, d, shape[2]//2, shape[3]//2, ...]=1,  \   d=0,1...N

    where, ``N`` is the minimum value of ``in_channels`` and ``out_channels``

    Args:
        groups(int, optional): 0-dimension of the Tensor will be divided by groups,
            each group has the same value. Default: 1.
        name(str, optional): The default value is None. Normally there is no need for user to set this
            property. For more information, please refer to :ref:`api_guide_Name`.

    Returns:
        Dirac initializer instance objects.

    Examples:
        .. code-block:: python

            >>> import paddle

            >>> # 1. For kernel_size is uneven number:
            >>> attr = paddle.ParamAttr(initializer=paddle.nn.initializer.Dirac())
            >>> conv = paddle.nn.Conv1D(3, 2, 3, weight_attr=attr)
            >>> print(conv.weight)
            Parameter containing:
            Tensor(shape=[2, 3, 3], dtype=float32, place=CPUPlace, stop_gradient=False,
            [[[0., 1., 0.],
              [0., 0., 0.],
              [0., 0., 0.]],
             [[0., 0., 0.],
              [0., 1., 0.],
              [0., 0., 0.]]])
            >>> input = paddle.rand([8, 3, 10])
            >>> output = conv(input)
            >>> output == input[:, 0:2, 1:9]
            >>> print(output.shape)
            [8, 2, 8]
            >>> # It means output is almost the same with input, 2 channels are reserved

            >>> # 2. For kernel_size is even number:
            >>> attr = paddle.ParamAttr(initializer=paddle.nn.initializer.Dirac())
            >>> conv = paddle.nn.Conv1D(3, 2, 4, weight_attr=attr)
            >>> print(conv.weight)
            Parameter containing:
            Tensor(shape=[2, 3, 4], dtype=float32, place=CPUPlace, stop_gradient=False,
            [[[0., 0., 1., 0.],
              [0., 0., 0., 0.],
              [0., 0., 0., 0.]],
             [[0., 0., 0., 0.],
              [0., 0., 1., 0.],
              [0., 0., 0., 0.]]])
    c                 óh   •— |dkD  rt        |t        «      sJ d«       ‚t        ‰| �  «        || _        y )Nr   z& 'groups' must be a positive integer. )Ú
isinstanceÚintÚsuperÚ__init__Ú_groups)ÚselfÚgroupsÚnameÚ	__class__s      €údG:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/nn/initializer/dirac.pyr   zDirac.__init__Z   s<   ø€ Ø˜ŠzœjØ”Cô
ð 	4à3ó	4ð 
ô 	‰ÑÔØˆ�ó    c           
      óf  — | j                  |«      }t        |t        j                  «      sJ ‚t        |t        j                  «      sJ ‚t        |dg d¢d«       t        |j                  «      dv sJ d«       ‚|j                  d   | j                  z  dk(  sJ d«       ‚|j                  t        j                  j                  k7  r€|j                  t        j                  dj!                  d	|j"                  d
g«      «      |j                  t        j                  j                  t        j                  j$                  d¬«      }n|}d}t        j&                  «       rqt(        j*                  j-                  «       5  t/        «       }t1        j2                  ||j                  t5        t7        d«      «      |j                  |«       ddd«       n9|j9                  di d|it7        d«      |j                  |j                  dœd¬«       |j                  }|d   | j                  z  }t;        ||d   «      }g }	g }
g }d}t=        |«      D ]  }|j?                  d|«       ||z  }Œ tA        | j                  «      D ]y  }tA        |«      D ]i  }|
jC                  d«       d}tE        |«      D ]5  \  }}|dk(  r||||z  z   |z  z  }Œ|dk(  r	|||z  z  }Œ(|||   dz  |z  z  }Œ7 |	jC                  |«       Œk Œ{ t        j&                  «       rPt(        j*                  j-                  «       5  t1        jF                  |dg«      }|jI                  |«       ddd«       n�|j                  t        j                  dj!                  |j"                  dg«      «      |j                  |j                  t        j                  j$                  dd¬«      }|j9                  dd|iddgi||dœd¬«       |j                  t        j                  d«      dd¬«      }t        j&                  «       r�t(        j*                  j-                  «       5  t        jJ                  «       }t1        jL                  |t        |	«      gt        j                  jN                  |	t/        «       «       |jI                  |«       ddd«       n=|j9                  dd|it        j                  jN                  t        |	«      g|	d œd¬!«       |j                  t        j                  d"«      dd¬«      }t        j&                  «       r�t(        j*                  j-                  «       5  t        jJ                  «       }t1        jL                  |t        |
«      gt        j                  j                  |
t/        «       «       |jI                  |«       ddd«       n=|j9                  dd|it        j                  j                  t        |
«      g|
d#œd¬!«       t        j&                  «       rÑt(        j*                  j-                  «       5  t1        jP                  |||d«      }|jI                  |«       t1        jF                  ||«      }|jI                  |«       |j                  t        j                  j                  k7  r1t1        jR                  ||j                  «      }|jI                  |«       ddd«       �n|j9                  d$|||d%œd&did|id¬«      }|j                  t        j                  dj!                  |j"                  dg«      «      |j                  |j                  t        j                  j$                  dd¬«      }|j9                  dd|id|i||dœd¬«       |j                  t        j                  j                  k7  r1|j9                  d'd|id|i|j                  |j                  d(œd¬«       tU        «       s||_+        |S # 1 sw Y   �ŒöxY w# 1 sw Y   �Œ#xY w# 1 sw Y   �Œ(xY w# 1 sw Y   �Œ-xY w# 1 sw Y   ŒPxY w))a’  Initialize the input tensor with dirac initializer.

        Args:
            var(Tensor): Tensor that needs to be initialized.
            block(Block, optional): The block in which initialization ops
                   should be added. Used in static graph only, default None.

        Returns:
            The most critical OP(scatter) in this initializer, which contains 7~8 ops in total.
        ÚOut)Úfloat16Úbfloat16Úfloat32Úfloat64r   )r   é   é   z=Only Tensor with 3/4/5 dimensions can be initialized by Diracr   z.Tensor 0-dimension must be divisible by groupsÚ.ÚdiracÚtmpF)r   ÚshapeÚdtypeÚtypeÚpersistableNÚfill_constant)Úvaluer(   r'   T)r)   ÚinputsÚoutputsÚattrsÚstop_gradientr   g      ð?é   éÿÿÿÿÚXShape)r   r(   r'   r)   r*   r0   Úreshape2ÚXr'   )r   r3   )r)   r-   r/   r.   r0   Úscatter_index)r   r*   r0   Úassign_value)r(   r'   Úint64_values)r)   r.   r/   r0   Úscatter_value)r(   r'   Úfp32_valuesÚscatter)r5   ÚIdsÚUpdatesÚ	overwriteÚcast)Úin_dtypeÚ	out_dtype),Ú_check_blockr   r   Ú	ParameterÚBlockr
   Úlenr'   r   r(   r	   ÚVarTypeÚFP32Ú
create_varr   ÚgenerateÚjoinr   Ú
LOD_TENSORÚin_dygraph_moder   ÚdygraphÚno_gradr   r   Úfull_ÚstrÚfloatÚ	append_opÚminÚreversedÚinsertÚrangeÚappendÚ	enumerateÚreshapeÚ_share_underline_tensor_toÚ_create_tensorÚassign_value_ÚINT64r;   r?   r   Úop)r   ÚvarÚblockÚout_varr^   ÚplaceÚorigin_shapeÚnum_per_groupÚ	min_shapeÚidx_listÚ
value_listÚstridesÚprodÚdimÚiÚjÚoffsetÚkÚstrideÚtmp_outÚx_shapeÚindex_tensorÚ
tmp_tensorÚvalue_tensorÚtmp_reshape_outÚtmp_cast_outs                             r   Ú__call__zDirac.__call__a   sU  € ð ×!Ñ! %Ó(ˆÜ˜#œy×2Ñ2Ô3Ð3Ð3Ü˜%¤§¡Ô1Ð1Ð1Ü Ø�ÒEÀwô	
ô �3—9‘9‹~ð "
ñ 
ð 	Kð Kó		Kð 
ð �I‰I�a‰L˜4Ÿ<™<Ñ'Øòð 	Aà@ó	Að ð �9‰9œŸ™×,Ñ,Ò,Ø×&Ñ&Ü ×)Ñ)¨#¯(©(°G¸S¿X¹XÀuÐ3MÓ*NÓOØ—i‘iÜ—o‘o×*Ñ*Ü—_‘_×/Ñ/Ø!ð 'ó ‰Gð ˆGØˆÜ×$Ñ$Ô&Ü—‘×%Ñ%Õ'Ü/Ó1�Ü—‘Ø˜WŸ]™]¬C´°a³«M¸7¿=¹=È%ô÷ (Ð'ð �O‰OØ$ØØ Ð(ä" 1›XØ$Ÿ]™]Ø$Ÿ]™]ñð
 #ð ô 
ð —y‘yˆØ$ Q™¨4¯<©<Ñ7ˆÜ˜ |°A¡Ó7ˆ	àˆØˆ
ØˆØˆÜ˜LÖ)ˆCØ�N‰N˜1˜dÔ#Ø�C‰K‰Dð *ô �t—|‘|Ö$ˆAÜ˜9Ö%�Ø×!Ñ! #Ô&Ø�Ü!*¨7Ö!3‘I�A�vØ˜A’vØ 1 q¨=Ñ'8Ñ#8¸FÑ"BÑB™Ø˜ašØ ! f¡*Ñ,™à ,¨q¡/°QÑ"6¸Ñ"?Ñ?™ð "4ð —‘ Õ'ñ &ð %ô ×$Ñ$Ô&Ü—‘×%Ñ%Õ'Ü Ÿ.™.¨°2°$Ó7�Ø×2Ñ2°7Ô;÷ (Ð'ð ×&Ñ&Ü ×)Ñ)¨#¯(©(°G·L±LÀ(Ð3KÓ*LÓMØ—m‘mØ—m‘mÜ—_‘_×/Ñ/Ø!Ø"ð 'ó ˆGð �O‰OØØ˜W�~Ø  �oØ '°7Ñ;Ø"ð ô ð ×'Ñ'Ü×%Ñ% oÓ6ØØð (ó 
ˆô ×$Ñ$Ô&Ü—‘×%Ñ%Õ'Ü&×5Ñ5Ó7�
Ü×$Ñ$ØÜ˜“]�OÜ—O‘O×)Ñ)ØÜ+Ó-ôð ×5Ñ5°lÔC÷ (Ð'ð �O‰OØ#Ø Ð-ä$Ÿ_™_×2Ñ2Ü! (›m˜_Ø$,ñð
 #ð ô 	ð ×'Ñ'Ü×%Ñ% oÓ6ØØð (ó 
ˆô ×$Ñ$Ô&Ü—‘×%Ñ%Õ'Ü&×5Ñ5Ó7�
Ü×$Ñ$ØÜ˜“_Ð%Ü—O‘O×(Ñ(ØÜ+Ó-ôð ×5Ñ5°lÔC÷ (Ð'ð �O‰OØ#Ø Ð-ä$Ÿ_™_×1Ñ1Ü! *›oÐ.Ø#-ñð
 #ð ô 	ô ×$Ñ$Ô&Ü—‘×%Ñ%Õ'Ü Ÿ.™.Ø˜\¨<¸ó�ð ×2Ñ2°7Ô;Ü"(§.¡.°¸,Ó"G�Ø×:Ñ:¸7ÔCØ—9‘9¤§¡× 4Ñ 4Ò4Ü#)§;¡;¨w¸¿	¹	Ó#B�LØ ×;Ñ;¸CÔ@÷ (Ñ'ð —‘Øà Ø'Ø+ñð
 # DÐ)Ø Ð(Ø"ð !ó 
ˆBð ×&Ñ&Ü ×)Ñ)¨#¯(©(°G·L±LÀ(Ð3KÓ*LÓMØ—m‘mØ—m‘mÜ—_‘_×/Ñ/Ø!Ø"ð 'ó ˆGð �O‰OØØ˜W�~Ø Ð-Ø '°7Ñ;Ø"ð ô ð �y‰yœGŸO™O×0Ñ0Ò0Ø—‘ØØ ˜>Ø" C˜LØ'.§}¡}À3Ç9Á9ÑMØ"&ð  ô ô Ô ØˆCŒFØˆ	÷U (Ñ'ú÷V (Ñ'ú÷6 (Ñ'ú÷: (Ñ'ú÷0 (Ð'ús@   Å;A
_3Ì&)` Ñ A)`ÕA)`Ø)B)`'ß3_=à `
à`à`$à'`0)r   N)N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   rw   Ú__classcell__)r   s   @r   r   r      s   ø„ ñ;õz÷Qr   r   N)Úpaddler   r   Úpaddle.utilsr   Ú r   r   Ú	base.corer	   Úbase.data_feederr
   Úbase.frameworkr   Úinitializerr   Ú__all__r   © r   r   Ú<module>r†      s2   ð÷ +Ý $å Ý Ý  Ý 8Ý 5Ý $à
€ôVˆKõ Vr   