Ë
    Ž\;jŠ  ã                   ój   — 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
mZ d dlmZ  G d„ de«      Zy)é    N)ÚcoreÚ	frameworkÚunique_name)Úappend_backward)ÚVariableÚin_dygraph_modeÚprogram_guard)Ú	Optimizerc                   óÈ   — e Zd ZdZd„ Zd„ Zd„ Zej                  d„ «       Z	d„ Z
d„ Zd„ Zd	„ Zd
„ Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zdd„Z	 	 	 	 dd„Zd„ Z	 dd„Zy)ÚRecomputeOptimizera   
        :api_attr: Static Graph

    Recompute Optimizer Wrapper

    Normally, a training step contains three sub-steps: first, run forward
    Operators to calculate the loss; second, run backward Operators to
    calculate gradient of the parameters; third, apply optimization method
    to update the value of the parameters.

    In the forward computation process, all variables that are needed by
    backward computation process will be kept in memory, which occupy a great
    amount of memory when the network becomes very deep.

    Recompute split the network to k segments. In each segment, It will
    recompute the forward Operators, before running backward operators. It is
    very helpful for saving memory.

    The Variables that separate a network to segments are called as checkpoints,
    and users should set it manually. The usage is very simple:

    Args:
        optimizer (Optimizer): The optimizer that is applied to parameters.

    Examples:
        .. code-block:: python

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

            >>> paddle.enable_static()

            >>> def gen_data():
            ...     return {"x": np.random.random(size=(32, 32)).astype('float32'),
            ...     "y": np.random.randint(2, size=(32, 1)).astype('int64')}
            >>> def mlp(input_x, input_y, hid_dim=128, label_dim=2):
            ...     print(input_x)
            ...     fc_1 = paddle.static.nn.fc(x=input_x, size=hid_dim)
            ...     prediction = paddle.static.nn.fc(x=[fc_1], size=label_dim, activation='softmax')
            ...     cost = paddle.nn.functional.cross_entropy(
            ...         input=prediction, label=input_y,
            ...         reduction='none', use_softmax=False
            ...     )
            ...     sum_cost = paddle.mean(cost)
            ...     return sum_cost, fc_1, prediction
            >>> input_x = paddle.static.data(name="x", shape=[-1,32], dtype='float32')
            >>> input_y = paddle.static.data(name="y", shape=[-1,1], dtype='int64')
            >>> cost, fc_1, pred = mlp(input_x, input_y)

            >>> sgd = paddle.optimizer.Adam(learning_rate=0.01)
            >>> sgd = paddle.incubate.optimizer.RecomputeOptimizer(sgd)
            >>> sgd._set_checkpoints([fc_1, pred])
            >>> sgd.minimize(cost)

            >>> print("Finished optimize")
            Finished optimize
            >>> place = paddle.CPUPlace()
            >>> exe = paddle.static.Executor(place)
            >>> exe.run(paddle.static.default_startup_program())
            >>> step = 10

            >>> for i in range(step):
            ...     cost_val = exe.run(feed=gen_data(),
            ...             program=paddle.static.default_main_program(),
            ...             fetch_list=[cost.name])
            ...     print("step=%d cost=%f" % (i, cost_val[0]))
            var x : LOD_TENSOR.shape(-1, 32).dtype(float32).stop_gradient(True)
            Finished optimize
            step=0 cost=0.737203
            step=1 cost=1.308077
            step=2 cost=0.768422
            step=3 cost=1.239475
            step=4 cost=0.882643
            step=5 cost=0.738027
            step=6 cost=0.819374
            step=7 cost=0.818534
            step=8 cost=0.753692
            step=9 cost=0.787448

    c                 óÄ   — t        «       rt        d«      ‚|| _        d | _        | j                  j                  | _        | j                  j
                  | _        d| _        y )Nz-In dygraph, don't support RecomputeOptimizer.F)r   Ú	ExceptionÚ
_optimizerÚ_checkpointsÚ_learning_rateÚ_learning_rate_mapÚenable_offload)ÚselfÚ	optimizers     úlG:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/incubate/optimizer/recompute.pyÚ__init__zRecomputeOptimizer.__init__j   sO   € ÜÔÜÐKÓLÐLØ#ˆŒØ ˆÔØ"Ÿo™o×<Ñ<ˆÔØ"&§/¡/×"DÑ"DˆÔØ#ˆÕó    c                 óˆ   — t        |t        «      sJ d«       ‚|D ]  }t        |t        t        f«      rŒJ d«       ‚ || _        y)zR
        Args:
            checkpoints (list): List of Variable or string
        z=_checkpoints should be a list of Variable or a list of StringN)Ú
isinstanceÚlistr   Ústrr   )r   ÚcheckpointsÚckpts      r   Ú_set_checkpointsz#RecomputeOptimizer._set_checkpointss   s`   € ô
 Øœô
ð 	KàJó	Kð 
ó  ˆDÜØ”x¤�oõð OàNóOð ð  ð (ˆÕr   c                 ó   — d| _         y )NT)r   ©r   s    r   Ú_enable_offloadz"RecomputeOptimizer._enable_offload‚   s
   € Ø"ˆÕr   c                 ó   — t        d«      ‚)a·  
            :api_attr: Static Graph

        load function is not supported by Recompute Optimizer for now.
        :return: None

        Args:
            state_dict: the dict load by load_persistable method

        Examples:
            .. code-block:: python

                >>> import paddle

                >>> paddle.enable_static()
                >>> def mlp(input_x, input_y, hid_dim=128, label_dim=2):
                ...     fc_1 = paddle.static.nn.fc(x=input_x, size=hid_dim)
                ...     prediction = paddle.static.nn.fc(x=[fc_1], size=label_dim, activation='softmax')
                ...     cost = paddle.nn.functional.cross_entropy(
                ...         input=prediction, label=input_y,
                ...         reduction='none', use_softmax=False
                ...     )
                ...     sum_cost = paddle.mean(cost)
                ...     return sum_cost, fc_1, prediction

                >>> input_x = paddle.static.data(name="x", shape=[-1,32], dtype='float32')
                >>> input_y = paddle.static.data(name="y", shape=[-1,1], dtype='int64')
                >>> cost, fc_1, pred = mlp(input_x, input_y)
                >>> print("Finished FF")
                Finished FF

                >>> sgd = paddle.optimizer.Adam(learning_rate=0.01)
                >>> sgd = paddle.incubate.optimizer.RecomputeOptimizer(sgd)
                >>> sgd._set_checkpoints([fc_1, pred])
                >>> try:
                ...     state_dict = {}
                ...     sgd.load(state_dict)
                >>> except NotImplementedError as e:
                ...     print(e)
                load function is not supported by Recompute Optimizer for now
        z=load function is not supported by Recompute Optimizer for now)ÚNotImplementedError)r   Ú
state_dicts     r   ÚloadzRecomputeOptimizer.load…   s   € ôV "ØKó
ð 	
r   c                 ó:   — | j                   j                  |¬«      S )aô  
        call apply_gradients function of self._optimizer.

        Args:
            params_grads (list): list of (param, grad) pair to do optimization.

        Returns:
            list: A list of operators appended to the current program.

        Examples:
            .. code-block:: python

                >>> import paddle
                >>> import paddle.base.framework as framework

                >>> paddle.enable_static()

                >>> def mlp(input_x, input_y, hid_dim=128, label_dim=2):
                ...     fc_1 = paddle.static.nn.fc(x=input_x, size=hid_dim)
                ...     prediction = paddle.static.nn.fc(x=[fc_1], size=label_dim, activation='softmax')
                ...     cost = paddle.nn.functional.cross_entropy(
                ...         input=prediction, label=input_y,
                ...         reduction='none', use_softmax=False
                ...     )
                ...     sum_cost = paddle.mean(cost)
                ...     return sum_cost, fc_1, prediction

                >>> input_x = paddle.static.data(name="x", shape=[-1,32], dtype='float32')
                >>> input_y = paddle.static.data(name="y", shape=[-1,1], dtype='int64')
                >>> cost, fc_1, pred = mlp(input_x, input_y)
                >>> print("Finished FF")
                Finished FF

                >>> sgd = paddle.optimizer.Adam(learning_rate=0.01)
                >>> sgd = paddle.incubate.optimizer.RecomputeOptimizer(sgd)
                >>> sgd._set_checkpoints([fc_1, pred])
                >>> params_grads = sgd.backward(
                ...     cost,
                ...     startup_program=None,
                ...     parameter_list=None,
                ...     no_grad_set=None)

                >>> program = cost.block.program
                >>> with framework.program_guard(program, None):
                ...     optimize_ops = sgd.apply_gradients(params_grads)

                >>> print("Finished apply gradients")
                Finished apply gradients
        )Úparams_grads)r   Úapply_gradients)r   r(   s     r   r)   z"RecomputeOptimizer.apply_gradients´   s   € ðf �‰×.Ñ.¸LÐ.ÓIÐIr   c                 ó  — t        j                  |dz   «      }t        j                  |dz   «      }| j                  j                  «       j	                  || j
                  | j                  j                  «       j                  |«      j                  dd¬«      }| j                  j                  «       j	                  || j
                  | j                  j                  «       j                  |«      j                  dd¬«      }||fS )Nz@Pinnedz@FetchFT©ÚnameÚshapeÚdtypeÚpersistableÚstop_gradient)r   ÚgenerateÚ_main_programÚglobal_blockÚ
create_varÚcheckpoint_shapeÚvarr.   )r   ÚvarnameÚpinned_var_nameÚfetched_var_nameÚ
pinned_varÚ	fetch_vars         r   Ú_creat_varszRecomputeOptimizer._creat_varsé   sñ   € Ü%×.Ñ.¨w¸Ñ/BÓCˆÜ&×/Ñ/°¸(Ñ0BÓCÐà×'Ñ'×4Ñ4Ó6×AÑAØ Ø×'Ñ'Ø×$Ñ$×1Ñ1Ó3×7Ñ7¸Ó@×FÑFØØð Bó 
ˆ
ð ×&Ñ&×3Ñ3Ó5×@Ñ@Ø!Ø×'Ñ'Ø×$Ñ$×1Ñ1Ó3×7Ñ7¸Ó@×FÑFØØð Aó 
ˆ	ð Ð 0Ð0Ð0r   c                 ó  — d}|j                  «       }| j                  j                  «       }t        j                  j                  «       }|D ]º  }| j                  j                  «       j                  |«      }|j                  || j                  | j                  j                  «       j                  |j                  «      j                  dd¬«      }|j                  dd|id|j                  d|j                  d	d
dd||i¬«       Œ¼ y)a/  
        add fill_constant_ops to the end of the prog

        we should fill the pinned vars before runing the main_prog
        to instantiate their tensor hold_, which could tell us whether
        the host memory could hold all the checkpoints from all the
        GPU devices in this node.
        r   FTr+   Úfill_constantÚOutr-   r.   Úvalueg        Ú
place_typeé   )ÚtypeÚoutputsÚattrsN)r3   Úcheckpoint_name2pinned_nameÚvaluesr   Úop_proto_and_checker_makerÚkOpRoleAttrNamer2   r6   r4   r5   r,   r.   Ú	append_opr-   )	r   Ústartup_programÚop_roleÚblockÚfill_constant_varsÚOP_ROLE_KEYr7   r6   r:   s	            r   Ú_append_fill_constant_opsz,RecomputeOptimizer._append_fill_constant_opsÿ   sû   € ð ˆØ×,Ñ,Ó.ˆØ!×=Ñ=×DÑDÓFÐÜ×5Ñ5×EÑEÓGˆÛ)ˆGØ×$Ñ$×1Ñ1Ó3×7Ñ7¸Ó@ˆCà×)Ñ)ØØ×+Ñ+Ø×(Ñ(×5Ñ5Ó7×;Ñ;¸C¿H¹HÓE×KÑKØ!Ø"ð *ó ˆJð �O‰OØ$Ø Ð(à˜SŸY™YØ˜SŸY™YØ˜SØ  !Ø ðð õ 
ñ *r   c           
      óB  — t         j                  j                  «       }| j                  j	                  |dd| j
                  j                  «       j                  |«      gid| j
                  j                  «       j                  |«      gidt        |«      ||i¬«       y )NÚmemcpyÚXr?   Údst_place_type)rC   ÚinputsrD   rE   )	r   rH   rI   rM   Ú_insert_op_without_syncr2   r3   r6   Úint)r   Ú
insert_idxÚsrc_varnameÚdst_varnamerL   rT   rO   s          r   Ú_insert_async_memcpy_opz*RecomputeOptimizer._insert_async_memcpy_op"  s”   € ô ×5Ñ5×EÑEÓGˆØ�
‰
×*Ñ*ØØØ˜$×,Ñ,×9Ñ9Ó;×?Ñ?ÀÓLÐMÐNà˜×*Ñ*×7Ñ7Ó9×=Ñ=¸kÓJÐKðð $¤S¨Ó%8¸+ÀwÐOð 	+õ 	
r   c                 óœ   — || j                   v sJ d|› d�«       ‚| j                   |   }| j                  |   }| j                  |||dd«       y )NzTry to fetch z/ from Pinned Memory, but it is NOT a checkpointé   )rF   Úcheckpoint_name2fetch_namer[   )r   Úidxr7   Úpinned_varnameÚfetch_varnames        r   Ú_insert_fetch_opz#RecomputeOptimizer._insert_fetch_op0  sd   € à�t×7Ñ7Ñ7ð	Tà˜7˜)Ð#RÐSó	TØ7ð ×9Ñ9¸'ÑBˆØ×7Ñ7¸Ñ@ˆØ×$Ñ$ S¨.¸-ÈÈAÕNr   c                 ó~   — || j                   v sJ d|› d�«       ‚| j                   |   }| j                  |||dd«       y )NzTry to offload z- to Pinned Memory, but it is NOT a checkpointr   rB   )rF   r[   )r   r_   r7   r`   s       r   Ú_insert_offload_opz%RecomputeOptimizer._insert_offload_op9  sR   € à�t×7Ñ7Ñ7ð	Tà˜W˜IÐ%RÐSó	TØ7à×9Ñ9¸'ÑBˆØ×$Ñ$ S¨'°>À1ÀaÕHr   c                  ó   — y ©N© )r   Úop_idxÚcheckpoint_names      r   Ú_insert_sync_opz"RecomputeOptimizer._insert_sync_op@  ó   € àr   c                 óÎ   — t        | j                  «      dkD  sJ d«       ‚| j                  j                  d«      }t        j                  d|› d�«       d|f| j
                  |<   |S )Nr   z#Could NOT found checkpoint to fetchéÿÿÿÿzRecord fetch [Ú]Úfetch)ÚlenÚun_fetch_checkpoint_namesÚpopÚloggingÚdebugÚidx2insertions©r   r_   ri   s      r   Ú_record_fetch_opz#RecomputeOptimizer._record_fetch_opD  sl   € ä�×.Ñ.Ó/°!Ò3ð	1à0ó	1Ø3à×8Ñ8×<Ñ<¸RÓ@ˆÜ�‰˜ Ð&7°qÐ9Ô:Ø$+¨_Ð#=ˆ×Ñ˜CÑ àÐr   c                 óÆ   — | j                   j                  d«      }||k(  sJ dj                  ||«      «       ‚t        j                  d|› d�«       d|f| j
                  |<   y )Nr   z%expected to offload [{}] but got [{}]zRecord offload [rn   Úoffload)Úun_offload_checkpoint_namesrr   Úformatrs   rt   ru   )r   r_   ri   Úexpected_checkpoint_names       r   Ú_record_offload_opz%RecomputeOptimizer._record_offload_opN  sp   € Ø#'×#CÑ#C×#GÑ#GÈÓ#JÐ àÐ7Ò7ð	
à2×9Ñ9Ø$ oó
ó	
Ø7ô 	�‰Ð(¨Ð(9¸Ð;Ô<Ø$-¨Ð#?ˆ×Ñ˜CÒ r   c                 óÀ   — || j                   vsJ d|› d�«       ‚| j                   j                  |«       t        j                  d|› d�«       d|f| j                  |<   y )NzTry to sync the checkpoint [z] twicezRecord offload sync [rn   Úsync)Úsynced_checkpointsÚaddrs   rt   ru   rv   s      r   Ú_record_sync_opz"RecomputeOptimizer._record_sync_opX  sl   € à 4×#:Ñ#:Ñ:ð	Cà)¨/Ð):¸'ÐBó	CØ:à×Ñ×#Ñ# OÔ4Ü�‰Ð-¨oÐ->¸aÐ@ÔAØ$*¨OÐ#<ˆ×Ñ˜CÒ r   c                 ó  — i | _         | j                  d d  | _        | j                  j                  d«       | j                  d d  }i | _        | j                  D ]  }d| j                  |<   Œ t        | j                  j                  «      | _        t        | j                  j                  «      D ]5  \  }}t        |j                  j                  d«      «      dk(  sŒ.|| _         n | j                  t        | j                  j                  «      k  sJ d«       ‚| j                  | j                  «      }d }t        | j                  j                  | j                  d  «      D ]÷  \  }}| j                  |z   }|j                  j                  «       }|D ]Ä  }	|	|v sŒ|	| j                  vr¡| j                  |	   dk(  r%|}
|	| j                  d   k7  r| j                  |«      }
|	k(  sJ dj                  |
|	«      «       ‚| j                  j                  |   j!                  |	| j"                  |	   «       | j                  |	xx   dz  cc<   Œ·t%        d|	› d�«      ‚ Œù t        | j                  «      dk(  sJ | j                  › d	�«       ‚y )
Nrm   r   rL   r]   z#Could NOT found backword op in progz6Current recompute segment should use [{}] BUT got [{}]zuse checkpoint [z] before fetch in BWú# checkpoints have NOT been Recorded)ru   Úsorted_checkpoint_namesrq   rr   Úcheckpoint_usage_countrp   rM   ÚopsÚbw_strart_op_idxÚ	enumeraterW   ÚdescÚattrrw   Úinput_arg_namesr{   Ú_rename_inputr^   Ú
ValueError)r   Úneed_fetch_checkpoint_namesri   r_   ÚopÚfetched_checkpoint_varnameÚlast_last_fetch_checkpointÚiÚ
input_varsÚ	input_varÚsecond_to_last_fetch_checkpoints              r   Ú_parse_backwardz"RecomputeOptimizer._parse_backward`  sŒ  € Ø ˆÔà)-×)EÑ)EÁaÐ)HˆÔ&Ø×&Ñ&×*Ñ*¨2Ô.Ø&*×&DÑ&DÁQÐ&GÐ#Ø&(ˆÔ#Ø#×=Ô=ˆOØ;<ˆD×'Ñ'¨Ò8ð  >ô !$ D§J¡J§N¡NÓ 3ˆÔÜ  §¡§¡Ö0‰GˆC�Ü�2—7‘7—<‘< 	Ó*Ó+¨qÓ0Ø(+�Ô%Ùð 1ð
 ×$Ñ$¤sØ�J‰J�N‰Nó(
ò 
ð 	1à0ó	1ð 
ð
 &*×%:Ñ%:Ø×!Ñ!ó&
Ð"ð &*Ð"ä˜tŸz™zŸ~™~¨d×.CÑ.CÐ.EÐFÖG‰EˆAˆrØ×'Ñ'¨!Ñ+ˆCØŸ™×0Ñ0Ó2ˆJã'�	ØÐ ;Ò;Ø ¨×(FÑ(FÑFà×6Ñ6°yÑAÀQÒFð !;ð <ð  )¨D×,HÑ,HÈÑ,KÒKà$(×$9Ñ$9¸#Ó$>ð !;ð <¸yÒHðàS×ZÑZØ;¸YóóØHð
 Ÿ
™
Ÿ™ sÑ+×9Ñ9Ø%Ø ×;Ñ;¸IÑFôð ×3Ñ3°IÓ>À!ÑCÔ>ä(Ø.¨y¨kÐ9MÐNóð ñ9 (ð	 HôJ �×.Ñ.Ó/°1Ò4ð	Rà×,Ñ,Ð-Ð-PÐQó	RÙ4r   c                 óÈ  — t        | j                  «      dk(  ry t        | j                  j                  «      }t	        t        | j                  |«      «      D ]’  }|| j                  v sŒ| j                  |   \  }}|dk(  r9| j                  ||«       t        j                  d|› d�«       | j                  |= Œb|dk(  sŒh| j                  ||«       t        j                  d|› d�«       Œ” | j                  j                  «        t        | j                  «      dk(  s?J dj                  | j                  j                  «       D �cg c]  }|d   ‘Œ	 c}«      «       ‚y c c}w )	Nr   ro   úInsert [z] fetch op.r   zSync [z{} checkpoints left un-Fecthedr]   )rp   ru   rM   r‡   ÚreversedÚrangerˆ   rb   rs   rt   rj   Ú_sync_with_cppr{   rG   )r   Útotal_oprh   Ú	operationri   Úeles         r   Ú_update_backwardz#RecomputeOptimizer._update_backward¢  s>  € Üˆt×"Ñ"Ó# qÒ(ØÜ�t—z‘z—~‘~Ó&ˆÜœu T×%:Ñ%:¸HÓEÖFˆFØ˜×,Ñ,Ò,Ø-1×-@Ñ-@ÀÑ-HÑ*�	˜?Ø Ò'Ø×)Ñ)¨&°/ÔBÜ—M‘M H¨_Ð,=¸[Ð"IÔJØ×+Ñ+¨FÑ3Ø &Ó(Ø×(Ñ(¨°ÔAÜ—M‘M F¨?Ð*;¸;Ð"GÕHð Gð 	�
‰
×!Ñ!Ô#ä�×#Ñ#Ó$¨Ò)ð	
à+×2Ñ2Ø#×2Ñ2×9Ñ9Ô;Ó<Ñ;˜ˆS�‹VÐ;Ñ<ó
ó	
Ù)ùâ<s   ÅE
c                 ób  — i | _         | j                  d d  | _        | j                  j                  d«      }| j                  d d  }i | _        | j                  D ]  }dddœ| j                  |<   Œ t        «       | _        t        | j                  j                  «      | _
        t        | j                  j                  «      D ]5  \  }}t        |j                  j                  d«      «      dk(  sŒ.|| _
         n | j                  t        | j                  j                  «      k  sJ d«       ‚d }t        | j                  j                  | j                  | j                   «      D �]E  \  }}| j                  |z   }|j                  j!                  «       }|j                  j#                  «       }	|D �]¥  }
|
|v rÑt        |«      dk(  sJ dj%                  |
|«      «       ‚|
| j                  v r„|�j| j                  |   d   dk(  r| j'                  ||«       nB| j                  |   d	   }|dkD  sJ d
j%                  |«      «       ‚| j'                  |dz   |«       | j)                  |dz   |
«       |
}nt+        dj%                  |
«      «      ‚|
|k(  sŒßt        |«      dk(  sJ dj%                  |
|«      «       ‚|| j                  d   k(  s%J dj%                  || j                  d   |«      «       ‚| j                  |   d	   dk(  r| j'                  ||«       �Œd| j                  |   d	   }|dkD  sJ d
j%                  |«      «       ‚| j'                  |dz   |«       �Œ¨ |	D ]L  }||v sŒ|| j                  vsJ d|› d�«       ‚| j                  |   dxx   dz  cc<   || j                  |   d	<   ŒN �ŒH t        | j                  «      dk(  sJ | j,                  › d�«       ‚t        | j                  «      t        |«      k(  s5J dj%                  t        |«      t        | j                  «      z
  «      «       ‚y )Nrm   r   )Úcountr_   rL   z"Could NOT found Forward op in progr]   zJchekpoint should be the only Output of a certain op, but [{}] is from [{}]r¢   r_   z5last_usage_idx of checkpoint [{}] should large than 0z7There should be just ONE op that output checkpoint [{}]éþÿÿÿzJthe last offload chekpoint before [{}] is suppose to be [{}], but got [{}]zcheckpoint [z] used after syncr„   z%{} checkpoints have NOT been Recorded)ru   r…   rz   rr   Úcheckpoint_usage_count_and_idxÚsetr€   rp   rM   r‡   Úfw_strart_op_idxr‰   rW   rŠ   r‹   rˆ   Úoutput_arg_namesrŒ   r{   r‚   r}   rŽ   rq   )r   Úlast_checkpointÚneed_offload_checkpoint_namesri   r_   r�   Úlast_offload_checkpointr“   Úoutput_varsr”   Ú
output_varÚlast_usage_idxr•   s                r   Ú_parse_forwardz!RecomputeOptimizer._parse_forward·  s»  € Ø ˆÔà+/×+GÑ+GÉÐ+JˆÔ(Ø×:Ñ:×>Ñ>¸rÓBˆØ(,×(HÑ(HÉÐ(KÐ%Ø.0ˆÔ+Ø#×?Ô?ˆOàØñDˆD×/Ñ/°Ò@ð  @ô
 #&£%ˆÔÜ # D§J¡J§N¡NÓ 3ˆÔÜ  §¡§¡Ö0‰GˆC�Ü�2—7‘7—<‘< 	Ó*Ó+¨qÓ0Ø(+�Ô%Ùð 1ð
 ×$Ñ$¤sØ�J‰J�N‰Nó(
ò 
ð 	0à/ó	0ð 
ð #'ÐäØ�J‰J�N‰N˜4×0Ñ0°4×3HÑ3HÐI÷
‰EˆAˆrð ×'Ñ'¨!Ñ+ˆCØŸ'™'×2Ñ2Ó4ˆKØŸ™×0Ñ0Ó2ˆJä)�
ØÐ!>Ñ>ä˜KÓ(¨AÒ-ðàc×jÑjØ" BóóØ-ð
 " T×%EÑ%EÑEà2Ð>à $× CÑ CØ$;ñ!"à")ñ!+ð $%ò!%ð
 !%× 4Ñ 4Ø$'Ð)@õ!"ð
 %)×$GÑ$GØ(?ñ%&à&+ñ%-ð !/ð %3°QÒ$6ð!"à#Z×#aÑ#aØ$;ó$"ó!"Ø$6ð !%× 4Ñ 4Ø$2°QÑ$6Ð8Oô!"ð ×/Ñ/°°a±¸ÔDØ2<Ñ/ä(ØU×\Ñ\Ø *óóð ð  Ó0ä˜KÓ(¨AÒ-ðàc×jÑjØ" BóóØ-ð
 0Ø×7Ñ7¸Ñ;ò<ðð d×jÑjØ'Ø×4Ñ4°RÑ8Ø/óóð<ð ×;Ñ;Ø3ñàñ!ð òð
 ×,Ñ,¨SÐ2IÖJà)-×)LÑ)LØ3ñ*àñ*!˜ð +¨QÒ.ðàR×YÑYØ3óóØ.ð ×,Ñ,Ø*¨QÑ.Ð0GöðW *ó^ (�	ØÐ =Ò=à!¨×)@Ñ)@Ñ@ðCà% i [Ð0AÐBóCØ@à×7Ñ7¸	ÑBÀ7ÓKÈqÑPÓKØLO�D×7Ñ7¸	ÑBÀ5ÒIò (ðm
ô~ �×0Ñ0Ó1°QÒ6ð	Rà×,Ñ,Ð-Ð-PÐQó	RØ6ä�4×*Ñ*Ó+¬sØ)ó0
ò 
ð 	
à2×9Ñ9ÜÐ-Ó.´°T×5LÑ5LÓ1MÑMó
ó	
ñ 
r   c                 ó¸  — t        | j                  «      dk(  ry t        t        | j                  | j
                  «      «      D ]Ÿ  }|| j                  v sŒ| j                  |   \  }}|dk(  r9| j                  ||«       t        j                  d|› d�«       | j                  |= Œb|dk(  sŒh| j                  ||«       t        j                  d|› d�«       | j                  |= Œ¡ | j                  j                  «        t        | j                  «      dk(  s?J dj                  | j                  j                  «       D �cg c]  }|d   ‘Œ	 c}«      «       ‚y c c}w )	Nr   ry   r™   z] offload op.r   z] offload_sync op.z {} checkpoints left un-Offloadedr]   )rp   ru   rš   r›   r¦   rˆ   rd   rs   rt   rj   rM   rœ   r{   rG   )r   rh   rž   ri   rŸ   s        r   Ú_update_forwardz"RecomputeOptimizer._update_forward6  sJ  € Üˆt×"Ñ"Ó# qÒ(ØÜÜ�$×'Ñ'¨×)>Ñ)>Ó?ö
ˆFð ˜×,Ñ,Ò,Ø-1×-@Ñ-@ÀÑ-HÑ*�	˜?Ø 	Ò)Ø×+Ñ+¨F°OÔDÜ—M‘M H¨_Ð,=¸]Ð"KÔLØ×+Ñ+¨FÑ3Ø &Ó(Ø×(Ñ(¨°ÔAÜ—M‘MØ" ?Ð"3Ð3EÐFôð ×+Ñ+¨FÑ3ð
ð  	�
‰
×!Ñ!Ô#ä�×#Ñ#Ó$¨Ò)ð	
à-×4Ñ4Ø#×2Ñ2×9Ñ9Ô;Ó<Ñ;˜ˆS�‹VÐ;Ñ<ó
ó	
Ù)ùâ<s   Ä?E
c                  ó   — y rf   rg   r!   s    r   Ú_check_offload_fetchz'RecomputeOptimizer._check_offload_fetchP  rk   r   Nc                 ó>  — |j                   j                  | _        |j                   | _         |€t        j                  j                  «       }t        | j                  |«      5  t        | j                  «      dkD  s J dj                  | j                  «      «       ‚t        d„ | j                  D «       «      s J dj                  | j                  «      «       ‚i | _        i | _        | j                  D ]4  }| j                  |«      \  }}|| j                  |<   || j                  |<   Œ6 | j                  |«       | j!                  «        | j#                  «        | j%                  «        | j'                  «        | j)                  «        ddd«       y# 1 sw Y   yxY w)zó
        core steps for recompute offload
        1. create pinned vars and temp vars
        2. parse & update Forward pass: offload, sync
        3. parse & update Backward pass: rename, fetch, sync
        4. verify the correctness
        Nr   zFcheckpoints shape {} should be an non empty list like: [12, 512, 1024]c              3   ó&   K  — | ]	  }|d kD  –— Œ y­w)r   Nrg   )Ú.0rŸ   s     r   Ú	<genexpr>z.RecomputeOptimizer._offload.<locals>.<genexpr>g  s   è ø€ ð Ù#8˜C��a•Ñ#8ùs   ‚zLall ele in checkpoints shape {} should be a determined integer larger than 0)rM   Úprogramr2   ÚpaddleÚstaticÚdefault_startup_programr	   rp   r5   r{   ÚallrF   r^   r…   r<   rP   r—   r    r®   r°   r²   )r   ÚlossrK   Úcheckpoint_varnamer8   Úfetch_var_names         r   Ú_offloadzRecomputeOptimizer._offloadT  s‹  € ð "ŸZ™Z×/Ñ/ˆÔØ—Z‘ZˆŒ
ØÐ"Ü$Ÿm™m×CÑCÓEˆOä˜4×-Ñ-¨Õ?ä�D×)Ñ)Ó*¨QÒ.ðàW×^Ñ^Ø×%Ñ%óóØ.ô ñ Ø#'×#8Ò#8óô ð à]×dÑdØ×%Ñ%óóð ð
 02ˆDÔ,Ø.0ˆDÔ+Ø&*×&BÔ&BÐ"Ø26×2BÑ2BØ&ó3Ñ/� ð
 $ð ×0Ñ0Ø&ñð
 #ð ×/Ñ/Ø&òð 'Cð ×*Ñ*¨?Ô;ð × Ñ Ô"Ø×!Ñ!Ô#à×ÑÔ!Ø× Ñ Ô"à×%Ñ%Ô'÷A @×?Ñ?ús   Á#D'FÆFc                 óP  — | j                   €J d«       ‚t        «       rt        d«      ‚|j                  | _        |j
                  j                  }t        ||«      5  g }| j                   D ]N  }t        |t        «      r|j                  |«       Œ%|j                  |j
                  j                  |«      «       ŒP t        |«      dkD  rt        ||||¬«      \  }	}
nt        ||||¬«      }	ddd«       | j                  r
| _        | j!                  ||¬«       	S # 1 sw Y   Œ1xY w)a^  
        call append_backward with checkpoints.

        Args:
            loss (Variable): loss variable to run optimizations.
            startup_program (Program): startup_program for initializing parameters
                in `parameter_list`.
            parameter_list (list): list of Variables or Variable.names to update.
            no_grad_set (set|None): set of Variables or Variables.names should be ignored.
            callbacks (list|None): list of callables to run when appending backward
                operator for one parameter.
            checkpoints (list): list of Variables as checkpoints

        Examples:
            .. code-block:: python

                >>> import paddle

                >>> paddle.enable_static()

                >>> def mlp(input_x, input_y, hid_dim=128, label_dim=2):
                ...     fc_1 = paddle.static.nn.fc(x=input_x, size=hid_dim)
                ...     prediction = paddle.static.nn.fc(x=[fc_1], size=label_dim, activation='softmax')
                ...     cost = paddle.nn.functional.cross_entropy(
                ...         input=prediction, label=input_y,
                ...         reduction='none', use_softmax=False
                ...     )
                ...     sum_cost = paddle.mean(cost)
                ...     return sum_cost, fc_1, prediction

                >>> input_x = paddle.static.data(name="x", shape=[-1,32], dtype='float32')
                >>> input_y = paddle.static.data(name="y", shape=[-1,1], dtype='int64')
                >>> cost, fc_1, pred = mlp(input_x, input_y)
                >>> print("Finished FF")
                Finished FF

                >>> sgd = paddle.optimizer.Adam(learning_rate=0.01)
                >>> sgd = paddle.incubate.optimizer.RecomputeOptimizer(sgd)
                >>> sgd._set_checkpoints([fc_1, pred])
                >>> params_grads = sgd.backward(
                ...     cost,
                ...     startup_program=None,
                ...     parameter_list=None,
                ...     no_grad_set=None)
                >>> print("Finished backward")
                Finished backward
        Nú&You should call _set_checkpoints firstú*DyGraph current does not support recomputer   )r   )rK   )r   r   r$   r.   Ú_dtyperM   r·   r	   r   r   Úappendr6   rp   r   r   r…   r¿   )r   r¼   rK   Úparameter_listÚno_grad_setÚ	callbacksr·   Úcheckpoint_varsr   r(   r…   s              r   ÚbackwardzRecomputeOptimizer.backwardƒ  s!  € ðp ×ÑÐ)ð	4à3ó	4Ø)ô ÔÜ%Ø<óð ð —j‘jˆŒØ—*‘*×$Ñ$ˆÜ˜7 OÕ4Ø ˆOØ×)Ô)�Ü˜d¤HÔ-Ø#×*Ñ*¨4Õ0à#×*Ñ*¨4¯:©:¯>©>¸$Ó+?Õ@ð	 *ô �?Ó# aÒ'Ü8GØØ"ØØ /ô	9Ñ5�Ñ5ô  /ØØ"ØØ /ô	 �÷# 5ð0 ×ÒØ+BˆDÔ(Ø�M‰M˜$°ˆMÔ@àÐ÷9 5Ð4ús   ÁBDÄD%c                 óœ   — t        | j                  d«      r| j                  j                  n| j                  j                  } ||||¬«      S )aß  
        call the apply_optimize function of self._optimizer
        Args:
            loss (Variable): loss variable to run optimizations.
            startup_program (Program): startup_program for initializing parameters
                in `parameter_list`.
            params_grads (list): list of (param, grad) pair to do optimization.
        Examples:
            .. code-block:: python

                >>> import paddle

                >>> paddle.enable_static()

                >>> def mlp(input_x, input_y, hid_dim=128, label_dim=2):
                ...     fc_1 = paddle.static.nn.fc(x=input_x, size=hid_dim)
                ...     prediction = paddle.static.nn.fc(x=[fc_1], size=label_dim, activation='softmax')
                ...     cost = paddle.nn.functional.cross_entropy(
                ...         input=prediction, label=input_y,
                ...         reduction='none', use_softmax=False
                ...     )
                ...     sum_cost = paddle.mean(cost)
                ...     return sum_cost, fc_1, prediction

                >>> input_x = paddle.static.data(name="x", shape=[-1,32], dtype='float32')
                >>> input_y = paddle.static.data(name="y", shape=[-1,1], dtype='int64')
                >>> cost, fc_1, pred = mlp(input_x, input_y)
                >>> print("Finished FF")
                Finished FF

                >>> sgd = paddle.optimizer.Adam(learning_rate=0.01)
                >>> sgd = paddle.incubate.optimizer.RecomputeOptimizer(sgd)
                >>> sgd._set_checkpoints([fc_1, pred])
                >>> params_grads = sgd.backward(
                ...     cost,
                ...     startup_program=None,
                ...     parameter_list=None,
                ...     no_grad_set=None)

                >>> optimize_ops = sgd.apply_optimize(
                ...     cost, startup_program=None, params_grads=params_grads)

                >>> print("Finished apply_optimize")
                Finished apply_optimize
        Úapply_optimize©rK   r(   )Úhasattrr   rË   Ú_apply_optimize)r   r¼   rK   r(   Úfuncs        r   rË   z!RecomputeOptimizer.apply_optimizeã  sK   € ôb �t—‘Ð(8Ô9ð �O‰O×*Ò*à—‘×0Ñ0ð 	ñ
 Ø /Àô
ð 	
r   c                 óÚ   — t        |t        «      sJ d«       ‚| j                  €J d«       ‚t        «       rt	        d«      ‚| j                  ||||¬«      }| j                  |||¬«      }||fS )NzThe loss should be an Variable.rÁ   rÂ   )rK   rÅ   rÆ   rÌ   )r   r   r   r   r$   rÉ   rË   )r   r¼   rK   rÅ   rÆ   r(   Úoptimize_opss          r   ÚminimizezRecomputeOptimizer.minimize  s˜   € ô ˜$¤Ô)ÐLÐ+LÓLÐ)à×ÑÐ)ð	4à3ó	4Ø)äÔÜ%Ø<óð ð —}‘}ØØ+Ø)Ø#ð	 %ó 
ˆð ×*Ñ*Ø /Àð +ó 
ˆð ˜\Ð)Ð)r   rf   )NNNN)NNN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r"   r   Údeprecate_stat_dictr&   r)   r<   rP   r[   rb   rd   rj   rw   r}   r‚   r—   r    r®   r°   r²   r¿   rÉ   rË   rÒ   rg   r   r   r   r      s¶   „ ñOòb$ò(ò#ð ×"Ñ"ñ,
ó #ð,
ò\3Jòj1ò,!òF
òOòIòòò@ò=ò@RòD
ò*}
ò~
ò4ó-(ðd ØØØó^ò@6
ðr LPô*r   r   )rs   r¸   Úpaddle.baser   r   r   Úpaddle.base.backwardr   Úpaddle.base.frameworkr   r   r	   Úpaddle.optimizerr
   r   rg   r   r   Ú<module>rÜ      s-   ðó ã ß 4Ñ 4Ý 0ß JÑ JÝ &ôY*˜õ Y*r   