Ë
    Ž\;j�+  ã                   óÚ   — d dl Z d dlZd dlZd dlmZmZmZ d dlmZ d dl	m
Z
mZ ej                  d	d„«       Zej                  d	d„«       Zej                   e«        e«       ddfd„«       Zy)
é    N)ÚbackwardÚcoreÚ	framework)Úprim_config)ÚprimxÚutilsc                 ó,  ‡	— t        j                  «       st        d«      ‚t        | t        j
                  t        j                  f«      st        dt        | «      › d�«      ‚t        |t        j
                  t        j                  f«      st        dt        |«      › d�«      ‚t        j                  | «      t        j                  |«      t        j                  |«      }}}t	        j                  «       j                  «       Š	t        ˆ	fd„||z   D «       «      rt        d«      ‚t        j                  ‰	«       t        j                   |d   j"                  «      }|j%                  |||«      \  }}t        | t        j
                  «      r|d   S |S )aS  Forward mode of automatic differentiation.

    Note:
        **ONLY available in the static graph mode and primitive operators.**

    Args:
        outputs(Tensor|Sequence[Tensor]): The output tensor or tensors.
        inputs(Tensor|Sequence[Tensor]): The input tensor or tensors.
        grad_inputs(Tensor|Sequence[Tensor]): Optional, the gradient Tensor or
            Tensors of inputs which has the same shape with inputs, Defaults to
            None, in this case is equivalent to all ones.

    Returns:
        grad_outputs(Tensor|Sequence[Tensor]): The gradients for outputs.

    Examples:

        .. code-block:: python

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

            >>> paddle.enable_static()
            >>> paddle.incubate.autograd.enable_prim()

            >>> startup_program = paddle.static.Program()
            >>> main_program = paddle.static.Program()

            >>> with paddle.static.program_guard(main_program, startup_program):
            ...     x = paddle.static.data('x', shape=[1], dtype='float32')
            ...     y = x * x
            ...     y_grad = paddle.incubate.autograd.forward_grad(y, x)
            ...     paddle.incubate.autograd.prim2orig()
            ...
            >>> exe = paddle.static.Executor()
            >>> exe.run(startup_program)
            >>> y_grad = exe.run(main_program, feed={'x': np.array([2.]).astype('float32')}, fetch_list=[y_grad])
            >>> print(y_grad)
            [array([4.], dtype=float32)]

            >>> paddle.incubate.autograd.disable_prim()
            >>> paddle.disable_static()
    zRforward_grad must be running on primitiveoperators, use enable_prim to turn it on.ú5Expected outputs is Tensor|Sequence[Tesnor], but got Ú.ú4Expected inputs is Tensor|Sequence[Tesnor], but got c              3   ó<   •K  — | ]  }|j                   ‰k7  –— Œ y ­w©N©Úblock©Ú.0Úxr   s     €úiG:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/incubate/autograd/primapi.pyÚ	<genexpr>zforward_grad.<locals>.<genexpr>^   s   øè ø€ Ð
-¡W ˆ1�7‰7�eÕ¡Wùs   ƒzMVariable in inputs and targets should exist in current block of main program.r   )r   Úprim_enabledÚRuntimeErrorÚ
isinstancer   ÚVariableÚtypingÚSequenceÚ	TypeErrorÚtypeÚ
as_tensorsÚdefault_main_programÚcurrent_blockÚanyr   Ú	orig2primÚ	Transformr   Ú	linearize)
ÚoutputsÚinputsÚgrad_inputsÚysÚxsÚxs_dotÚadÚ_Úys_dotr   s
            @r   Úforward_gradr.      sh  ø€ ôZ ×ÑÔÜð8ó
ð 	
ô
 �g¤	× 2Ñ 2´F·O±OÐDÔEÜðÜ˜G“}�o Qð(ó
ð 	
ô
 �fœy×1Ñ1´6·?±?ÐCÔDÜðÜ˜F“|�n Að'ó
ð 	
ô 	×Ñ˜Ó!Ü×Ñ˜Ó Ü×Ñ˜Ó%ð ˆ€Bô ×*Ñ*Ó,×:Ñ:Ó<€EÜ
Ó
- R¨"¢WÓ
-Ô-Üðó
ð 	
ô
 
‡O�O�EÔÜ	�‰˜˜A™Ÿ™Ó	%€BØ—‘˜R  VÓ,�I€A€vä" 7¬I×,>Ñ,>Ô?ˆ6�!‰9ÐKÀVÐKó    c                 ó"  ‡— t        j                  «       s`t        j                  | ||«      }t	        |t
        j                  «      r-t	        |t        j                  «      rt        |«      dkD  r|d   S |S t	        | t
        j                  t        j                  f«      st        dt        | «      › d�«      ‚t	        |t
        j                  t        j                  f«      st        dt        |«      › d�«      ‚t        j                  | «      t        j                  |«      t        j                  |«      }}}t        j                  «       j                  «       Št        ˆfd„||z   D «       «      rt!        d«      ‚t#        j$                  ‰«       t#        j&                  ‰«      }|j)                  ||«      \  }}	t        d„ |	D «       «      rt!        d«      ‚|j+                  |	||«      \  }}
g }|D ]O  }|€Œ‰j,                  j/                  |j0                  «      }|dk  rt3        d	|› d�«      ‚|j5                  |«       ŒQ |j7                  t9        |«      «       |j;                  |«       t	        |t
        j                  «      r|
d   S |
S )
an  Reverse mode of automatic differentiation.

    Note:
        **ONLY available in the static graph mode and primitive operators**

    Args:
        outputs(Tensor|Sequence[Tensor]): The output Tensor or Tensors.
        inputs(Tensor|Sequence[Tensor]): The input Tensor or Tensors.
        grad_outputs(Tensor|Sequence[Tensor]): Optional, the gradient Tensor or
            Tensors of outputs which has the same shape with outputs, Defaults
            to None, in this case is equivalent to all ones.

    Returns:
        grad_inputs(Tensor|Tensors): The gradients for inputs.

    Examples:

        .. code-block:: python

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

            >>> paddle.enable_static()
            >>> paddle.incubate.autograd.enable_prim()

            >>> startup_program = paddle.static.Program()
            >>> main_program = paddle.static.Program()
            >>> with paddle.static.program_guard(main_program, startup_program):
            ...     x = paddle.static.data('x', shape=[1], dtype='float32')
            ...     x.stop_gradients = False
            ...     y = x * x
            ...     x_grad = paddle.incubate.autograd.grad(y, x)
            ...     paddle.incubate.autograd.prim2orig()
            ...
            >>> exe = paddle.static.Executor()
            >>> exe.run(startup_program)
            >>> x_grad = exe.run(main_program, feed={'x': np.array([2.]).astype('float32')}, fetch_list=[x_grad])
            >>> print(x_grad)
            [array([4.], dtype=float32)]

            >>> paddle.incubate.autograd.disable_prim()
            >>> paddle.disable_static()
    r   r
   r   r   c              3   óH   •K  — | ]  }|d uxr |j                   ‰k7  –— Œ y ­wr   r   r   s     €r   r   zgrad.<locals>.<genexpr>¸   s'   øè ø€ Ð
A¹°AˆA�TˆMÒ.˜aŸg™g¨Ñ.Ó.¹ùs   ƒ"zQVariable in inputs and outputs should be None or in current block of main programc              3   ó$   K  — | ]  }|d u –— Œ
 y ­wr   © )r   Úvars     r   r   zgrad.<locals>.<genexpr>Ä   s   è ø€ Ð
)¡&˜3ˆ3�$Œ;¡&ùs   ‚zEGrads cannot be computed. The given outputs does not depend on inputsz<op_index should be greater than or equal to 0, but op_index=)r   r   r   Ú	gradientsr   r   r   r   r   Úlenr   r   r   r   r    r!   r   r   r"   r#   r$   Ú	transposeÚopsÚindexÚopÚ
ValueErrorÚappendÚ	erase_opsÚsortedÚ
erase_dots)r%   r&   Úgrad_outputsr'   r(   r)   Úys_barr+   r*   r-   Úxs_barÚ
op_indexesr4   Úop_indexr   s                 @r   ÚgradrE   k   sM  ø€ ôZ ×ÑÔÜ×(Ñ(¨°&¸,ÓGˆô
 �vœy×1Ñ1Ô2Ü˜;¬¯©Ô8Ü�KÓ  1Ò$à˜q‘>Ð!àÐä�g¤	× 2Ñ 2´F·O±OÐDÔEÜðÜ˜G“}�o Qð(ó
ð 	
ô
 �fœy×1Ñ1´6·?±?ÐCÔDÜðÜ˜F“|�n Að'ó
ð 	
ô 	×Ñ˜Ó!Ü×Ñ˜Ó Ü×Ñ˜Ó&ð ˆ€Bô
 ×*Ñ*Ó,×:Ñ:Ó<€EÜ
Ó
A¸¸bºÓ
AÔAÜØ_ó
ð 	
ô 
‡O�O�EÔÜ	�‰˜Ó	€BØ—\‘\ " bÓ)�N€FˆFÜ
Ñ
)¡&Ó
)Ô)ÜØSó
ð 	
ð —\‘\ &¨&°&Ó9�N€FˆFð €JÛˆØ‰?Ø—y‘y—‘ s§v¡vÓ.ˆHØ˜!Š|Ü ØRÐS[ÐR\Ð\]Ð^óð ð ×Ñ˜hÕ'ð ð ‡L�L”˜
Ó#Ô$Ø‡M�M�&Ôä" 6¬9×+=Ñ+=Ô>ˆ6�!‰9ÐJÀFÐJr/   éÿÿÿÿc                 ó†  ‡‡— t        j                  «       syt        | t        j                  j
                  j                  «      r"t        j                  d«       | j                  }n�t        | t        j                  «      r]| D ]H  }t        |t        j                  j
                  j                  «      rŒ2t        dt        |«      › d�«      ‚ | d   j                  }nt        dt        | «      › d�«      ‚t        ‰t        t        f«      st        dt        ‰«      › d�«      ‚t        ‰t        t        f«      st        dt        ‰«      › d�«      ‚t         d	   ‰z  Št        j"                  |«      5  t        j                  d
«       t%        ‰«      dkD  rt%        ‰«      dkD  rˆˆfd„}nGt%        ‰«      dkD  rt%        ‰«      dk(  rˆfd„}n%t%        ‰«      dk(  rt%        ‰«      dkD  rˆfd„}nd„ }t'        j(                  | |||¬«       t         d   }t        j                  d|› �«       ddd«       y# 1 sw Y   yxY w)a—  Search nonbasic ops which have be registered composite rules and replace them with primitive ops.
    The operators in blacklist will be excluded from program when lowering into primitives, and only the
    operators in whitelist will be lowering. The priority of blacklist is higher than whitelist, it means
    an operator both in blacklist and whitelist will not be lowering.

    The finally set that will be lowering is:
        (blocks.ops & ops have decomposite rule & whitelist) - blacklist

    Args:
        blacklist(frozenset): The Operators that will be exclude when lowering into primitives.
        whitelist(frozenset): Only the operators in whitelist will be lowering into primitives.
        start_idx(int): If start_idx exceeds -1, ops[start_idx:] will be processed. Default: -1.
        backward_length(int): If backward_length exceeds -1, ops[:-backward_length] will be processed. Default: -1.
    Nz,Atomize composite op to primitive ops begin.z:Expect block or sequence of blocks, but sequence contains r   r   z,Expect block or sequence of blocks, but got z6Expected type of blacklisst is set|frozenset, but got z6Expected type of whiltelist is set|frozenset, but got Úforward_blacklistz'Lowering composite forward ops begin...c                 ó@   •— | j                   ‰v xr | j                   ‰vS r   ©r   )r   Ú	blacklistÚ	whitelists    €€r   Ú<lambda>zto_prim.<locals>.<lambda>  s   ø€  §¡¨)Ð 3Ò O¸¿¹ÀiÐ8OÐ Or/   c                 ó    •— | j                   ‰vS r   rJ   )r   rK   s    €r   rM   zto_prim.<locals>.<lambda>  s   ø€  §¡¨iÑ 7r/   c                 ó    •— | j                   ‰v S r   rJ   )r   rL   s    €r   rM   zto_prim.<locals>.<lambda>  s   ø€  §¡¨)Ñ 3r/   c                  ó   — y)NTr3   )r   s    r   rM   zto_prim.<locals>.<lambda>  s   €  r/   )Ú	start_idxÚbackward_lengthÚcomposite_ops_recordz'Lowering composite forward ops finish: )r   Ú_is_fwd_prim_enabledr   ÚpaddleÚbaser   ÚBlockÚloggingÚinfoÚprogramr   r   r   r   ÚsetÚ	frozensetr   Úprogram_guardr6   r   Ú_lower_composite)	ÚblocksrK   rL   rQ   rR   Úmain_programÚitemÚfilter_Úreplace_opss	    ``      r   Úto_primrd   Û   só  ù€ ô, ×$Ñ$Ô&ØÜ�&œ&Ÿ+™+×/Ñ/×5Ñ5Ô6Ü�‰ÐCÔDØ—~‘~‰Ü	�FœFŸO™OÔ	,ÛˆDÜ˜d¤F§K¡K×$9Ñ$9×$?Ñ$?Õ@ÜØPÔQUÐVZÓQ[ÐP\Ð\]Ð^óð ð ð
 ˜a‘y×(Ñ(‰äØ:¼4À»<¸.ÈÐJó
ð 	
ô �i¤#¤yÐ!1Ô2ÜØDÄTÈ)Ã_ÐDUÐUVÐWó
ð 	
ô �i¤#¤yÐ!1Ô2ÜØDÄTÈ)Ã_ÐDUÐUVÐWó
ð 	
ô Ð/Ñ0°9Ñ<€Iä	×	 Ñ	  Õ	.Ü�‰Ð>Ô?äˆy‹>˜AÒ¤# i£.°1Ò"4ÜO‰GÜ�‹^˜aÒ¤C¨	£N°aÒ$7Û7‰GÜ�‹^˜qÒ ¤S¨£^°aÒ%7Û3‰Gá$ˆGÜ×ÑØØØØ+õ		
ô "Ð"8Ñ9ˆÜ�‰Ð>¸{¸mÐLÔM÷% 
/×	.Ñ	.ús   Å4B:H7È7I r   )rX   r   rU   Úpaddle.baser   r   r   Úpaddle.base.corer   Úpaddle.incubate.autogradr   r   Ústatic_onlyr.   rE   r\   rd   r3   r/   r   Ú<module>ri      sŠ   ðó Û ã ß 1Ñ 1Ý (ß 1ð ×ÑòOLó ðOLðd ×ÑòlKó ðlKð^ ×Ññ ‹kÙ‹kØØòBNó ñBNr/   