Ë
    Ž\;j@.  ã                   óf   — d dl Z 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 g Z G d„ de«      Zy)	é    N)Ú	frameworkÚunique_name)Úbase)ÚVariable)ÚLayerHelper)Ú	Optimizerc                   ó®   ‡ — e Zd ZdZdZd
ˆ fd„	Zˆ fd„Zej                  e	j                  d„ «       «       Zd„ Zd„ Zd„ Ze	j                  	 dd	„«       Zˆ xZS )Ú	LookAheadaò  
    This implements the Lookahead optimizer of the
    paper : https://arxiv.org/abs/1907.08610.

    Lookahead keeps two sets of params: the fast_params and
    the slow_params. inner_optimizer update fast_params every
    training step. Lookahead updates the slow_params and fast_params
    every k training steps as follows:

    .. math::

        slow\_param_t &= slow\_param_{t-1} + \\alpha * (fast\_param_{t-1} - slow\_param_{t-1})

        fast\_param_t &=  slow\_param_t

    Args:
        inner_optimizer (Optimizer): The optimizer that update fast params step by step.
        alpha (float, optinal): The learning rate of Lookahead. The default value is 0.5.
        k (int, optinal): The slow params is updated every k steps. The default value is 5.
        name (str, optional): Normally there is no need for user to set this property.
            For more information, please refer to :ref:`api_guide_Name`.
            The default value is None.

    Examples:

        .. code-block:: python

            >>> import numpy as np
            >>> import paddle
            >>> import paddle.nn as nn

            >>> BATCH_SIZE = 16
            >>> BATCH_NUM = 4
            >>> EPOCH_NUM = 4

            >>> IMAGE_SIZE = 784
            >>> CLASS_NUM = 10
            >>> # define a random dataset
            >>> class RandomDataset(paddle.io.Dataset):
            ...     def __init__(self, num_samples):
            ...         self.num_samples = num_samples
            ...     def __getitem__(self, idx):
            ...         image = np.random.random([IMAGE_SIZE]).astype('float32')
            ...         label = np.random.randint(0, CLASS_NUM - 1,
            ...                                 (1, )).astype('int64')
            ...         return image, label
            ...     def __len__(self):
            ...         return self.num_samples

            >>> class LinearNet(nn.Layer):
            ...     def __init__(self):
            ...         super().__init__()
            ...         self._linear = nn.Linear(IMAGE_SIZE, CLASS_NUM)
            ...         self.bias = self._linear.bias
            ...     @paddle.jit.to_static
            ...     def forward(self, x):
            ...         return self._linear(x)

            >>> def train(layer, loader, loss_fn, opt):
            ...     for epoch_id in range(EPOCH_NUM):
            ...         for batch_id, (image, label) in enumerate(loader()):
            ...             out = layer(image)
            ...             loss = loss_fn(out, label)
            ...             loss.backward()
            ...             opt.step()
            ...             opt.clear_grad()
            ...             print("Train Epoch {} batch {}: loss = {}".format(
            ...                 epoch_id, batch_id, np.mean(loss.numpy())))
            >>> layer = LinearNet()
            >>> loss_fn = nn.CrossEntropyLoss()
            >>> optimizer = paddle.optimizer.SGD(learning_rate=0.1, parameters=layer.parameters())
            >>> lookahead = paddle.incubate.LookAhead(optimizer, alpha=0.2, k=5)

            >>> # create data loader
            >>> dataset = RandomDataset(BATCH_NUM * BATCH_SIZE)
            >>> loader = paddle.io.DataLoader(
            ...     dataset,
            ...     batch_size=BATCH_SIZE,
            ...     shuffle=True,
            ...     drop_last=True,
            ...     num_workers=2)

            >>> # doctest: +SKIP('The run time is too long to pass the CI check.')
            >>> train(layer, loader, loss_fn, lookahead)

    Úslowc                 ó  •— |€J d«       ‚d|cxk  rdk  sJ d«       ‚ J d«       ‚t        |t        «      r|dkD  sJ d«       ‚|| _        | j                  j                  €1t	        j
                  «       j                  «       j                  «       }n| j                  j                  }t        ‰| �%  ||d d |¬«       || _
        || _        d| _        t        | j                  j                  «      | _        d | _        d | _        y )	Nzinner optimizer can not be Noneg        ç      ð?zBalpha should be larger or equal to 0.0, and less or equal than 1.0r   zk should be a positive integer)Úlearning_rateÚ
parametersÚweight_decayÚ	grad_clipÚnameÚ	lookahead)Ú
isinstanceÚintÚinner_optimizerÚ_parameter_listr   Údefault_main_programÚglobal_blockÚall_parametersÚsuperÚ__init__ÚalphaÚkÚtyper   Ú	__class__Ú__name__ÚhelperÚ_global_step_varÚ_k_var)Úselfr   r   r   r   r   r    s         €úlG:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/incubate/optimizer/lookahead.pyr   zLookAhead.__init__r   s  ø€ ØÐ*ÐMÐ,MÓMÐ*à�5Ô˜CÒð	PàOó	PÙð	PàOó	PØä˜!œSÔ! a¨!¢eÐMÐ-MÓMÐ+à.ˆÔØ×Ñ×/Ñ/Ð7ä×.Ñ.Ó0×=Ñ=Ó?×NÑNÓPñ ð ×-Ñ-×=Ñ=ˆJä‰ÑØØ!ØØØð 	ô 	
ð ˆŒ
ØˆŒØˆŒ	Ü! $§.¡.×"9Ñ"9Ó:ˆŒØ $ˆÔØˆ�ó    c                 ó^   •— t         ‰| �  ||«       | j                  j                  ||«       y ©N)r   Ú_set_auxiliary_varr   )r%   ÚkeyÚvalr    s      €r&   r*   zLookAhead._set_auxiliary_var�   s(   ø€ Ü‰Ñ" 3¨Ô,Ø×Ñ×/Ñ/°°SÕ9r'   c                 ó(  — | j                   j                  «        | j                  «        g }| j                  D ]C  }|j                  sŒ|j                  «       €Œ!|j                  «       }|j                  ||f«       ŒE | j                  dd|¬«       y)a«  
        Execute the optimizer and update parameters once.

        Returns:
            None

        Examples:

            .. code-block:: python

                >>> import paddle
                >>> inp = paddle.rand([1,10], dtype="float32")
                >>> linear = paddle.nn.Linear(10, 1)
                >>> out = linear(inp)
                >>> loss = paddle.mean(out)
                >>> sgd = paddle.optimizer.SGD(learning_rate=0.1,parameters=linear.parameters())
                >>> lookahead = paddle.incubate.LookAhead(sgd, alpha=0.2, k=5)
                >>> loss.backward()
                >>> lookahead.step()
                >>> lookahead.clear_grad()

        N)ÚlossÚstartup_programÚparams_grads)r   ÚstepÚ_increment_global_varr   Ú	trainableÚ
_grad_ivarÚappendÚ_apply_optimize)r%   r0   ÚparamÚgrad_vars       r&   r1   zLookAhead.step”   s�   € ð2 	×Ñ×!Ñ!Ô#à×"Ñ"Ô$ØˆØ×)Ô)ˆEØ—?’?ØØ×ÑÓ!Ñ-Ø ×+Ñ+Ó-�Ø×#Ñ# U¨HÐ$5Õ6ð *ð 	×ÑØ t¸,ð 	õ 	
r'   c                 ó‚   — t        |t        j                  «      sJ ‚|D ]  }| j                  | j                  |«       Œ  y r)   )r   r   ÚBlockÚ_add_accumulatorÚ	_slow_str)r%   Úblockr   Úps       r&   Ú_create_accumulatorszLookAhead._create_accumulators¼   s4   € Ü˜%¤§¡Ô1Ð1Ð1ãˆAØ×!Ñ! $§.¡.°!Õ4ñ r'   c                 ó  — | j                   €=t        j                  j                  t	        j
                  d«      dgddd¬«      | _         | j                  j                  dd| j                   gid	| j                   gid
di¬«       y )NÚlookahead_stepé   r   Úint32T©r   ÚshapeÚvalueÚdtypeÚpersistableÚ	incrementÚXÚOutr1   r   )r   ÚinputsÚoutputsÚattrs)r#   ÚpaddleÚstaticÚcreate_global_varr   Úgenerater"   Ú	append_op)r%   s    r&   r2   zLookAhead._increment_global_varÂ   s�   € Ø× Ñ Ð(Ü$*§M¡M×$CÑ$CÜ ×)Ñ)Ð*:Ó;Ø�cØØØ ð %Dó %ˆDÔ!ð 	�‰×ÑØØ˜$×/Ñ/Ð0Ð1Ø˜T×2Ñ2Ð3Ð4Ø˜3�-ð	 	õ 	
r'   c                 óf  — t        j                  dgdd¬«      }t        j                  dgdd¬«      }t         j                  j	                  t        j                  d«      dg| j                  dd¬«      }t        j                  | j                  |«      }t        j                  | j                  |«      }t        j                  |d	¬
«      }t        j                  ||«      }t        j                  |d	¬
«      }| j                  | j                  |d   «      }	||d   z  d|z
  |	z  z   }
t        j                  |
|	«       | j                  |d   z  d| j                  z
  |	z  z   }
||
z  d|z
  |d   z  z   }t        j                  ||d   «       ||
z  d|z
  |	z  z   }t        j                  ||	«       y )NrB   rC   Úlookahead_ones)rE   rG   r   Úlookahead_zerosÚlookahead_kTrD   Úfloat32)rG   r   r   )rO   ÚonesÚzerosrP   rQ   r   rR   r   Ú	remainderr#   ÚequalÚcastÚ_get_accumulatorr<   Úassignr   )r%   r=   Úparam_and_gradÚone_varÚzero_varÚk_varÚmodÚcond_1Úcond_2Úslow_varÚtmp_varÚ	tmp_var_1s               r&   Ú_append_optimize_opzLookAhead._append_optimize_opÓ   sŒ  € Ü—+‘+ Q C¨wÐ=MÔNˆÜ—<‘<Ø�#˜WÐ+<ô
ˆô —‘×/Ñ/Ü×%Ñ% mÓ4Ø�#Ø—&‘&ØØð 0ó 
ˆô ×Ñ˜t×4Ñ4°eÓ<ˆä—‘˜d×3Ñ3°WÓ=ˆÜ—‘˜V¨9Ô5ˆä—‘˜c 8Ó,ˆÜ—‘˜V¨9Ô5ˆà×(Ñ(¨¯©¸ÈÑ9JÓKˆà˜>¨!Ñ,Ñ,°°F±
¸hÑ/FÑFˆÜ�‰�g˜xÔ(à—*‘*˜~¨aÑ0Ñ0°C¸$¿*¹*Ñ4DÈÑ3PÑPˆØ˜WÑ$¨¨F©
°nÀQÑ6GÑ'GÑGˆ	Ü�‰�i °Ñ!2Ô3à˜WÑ$¨¨F©
°hÑ'>Ñ>ˆ	Ü�‰�i Õ*r'   c                 óÄ   — t        |t        «      sJ d«       ‚| j                  j                  ||||¬«      \  }}| j	                  «        | j                  |||¬«      }||fS )a‚  
        Add operations to minimize ``loss`` by updating ``parameters``.

        Args:
            loss (Tensor): A ``Tensor`` containing the value to minimize.
            startup_program (Program, optional): :ref:`api_paddle_static_Program` for
                initializing parameters in ``parameters``. The default value
                is None, at this time :ref:`api_paddle_static_default_startup_program` will be used.
            parameters (list, optional): List of ``Tensor`` or ``Tensor.name`` to update
                to minimize ``loss``. The default value is None, at this time all parameters
                will be updated.
            no_grad_set (set, optional): Set of ``Tensor``  or ``Tensor.name`` that don't need
                to be updated. The default value is None.

        Returns:
            tuple: tuple (optimize_ops, params_grads), A list of operators appended
            by minimize and a list of (param, grad) tensor pairs, param is
            ``Parameter``, grad is the gradient value corresponding to the parameter.
            In static graph mode, the returned tuple can be passed to ``fetch_list`` in ``Executor.run()`` to
            indicate program pruning. If so, the program will be pruned by ``feed`` and
            ``fetch_list`` before run, see details in ``Executor``.

        Examples:

            .. code-block:: python

                >>> import paddle

                >>> inp = paddle.rand([1, 10], dtype="float32")
                >>> linear = paddle.nn.Linear(10, 1)
                >>> out = linear(inp)
                >>> loss = paddle.mean(out)
                >>> sgd = paddle.optimizer.SGD(learning_rate=0.1,parameters=linear.parameters())
                >>> lookahead = paddle.incubate.LookAhead(sgd, alpha=0.2, k=5)
                >>> loss.backward()
                >>> lookahead.minimize(loss)
                >>> lookahead.clear_grad()

        zThe loss should be an Tensor.)r/   r   Úno_grad_set)r/   r0   )r   r   r   Úminimizer2   r6   )r%   r.   r/   r   rl   Úoptimize_opsr0   Ú_s           r&   rm   zLookAhead.minimizeô   s�   € ôV ˜$¤Ô)ÐJÐ+JÓJÐ)ð &*×%9Ñ%9×%BÑ%BØØ+Ø!Ø#ð	 &Có &
Ñ"ˆ�lð 	×"Ñ"Ô$à× Ñ Ø /Àð !ó 
ˆð ˜\Ð)Ð)r'   )g      à?é   N)NNN)r!   Ú
__module__Ú__qualname__Ú__doc__r<   r   r*   r   Údygraph_onlyÚimperative_baseÚno_gradr1   r?   r2   rj   rm   Ú__classcell__)r    s   @r&   r
   r
      sq   ø„ ñUðl €Iõô<:ð ×ÑØ×Ññ$
ó ó ð$
òL5ò
ò"+ðB ×ÑàGKò:*ó ô:*r'   r
   )rO   Úpaddle.baser   r   Úpaddle.base.dygraphr   ru   Úpaddle.base.frameworkr   Úpaddle.base.layer_helperr   Úpaddle.optimizerr   Ú__all__r
   © r'   r&   Ú<module>r      s,   ðó ß .Ý 7Ý *Ý 0Ý &à
€ôV*�	õ V*r'   