o
    i eD+                     @   s   d Z ddlZddlZddlZddlmZmZ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 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ae  Z!dZ"de#fddZ$G dd dej%Z&dS )zHook for asynchronous checkpointing.

This hook dispatches checkpoint writing operations in a separate thread to
allow execution to continue on the main thread.
    N)AnyListOptionalText)	event_pb2)session)
meta_graph)ops)
tf_logging)metrics)basic_session_run_hooks)monitored_session)saver)session_run_hook)training_util)SummaryWriterCacheasync_checkpoint_v1returnc                 C   s   t t||  d dS )z@Returns the duration between start and end time in microseconds.i@B r   )maxint)start_time_secondsend_time_seconds r   U/var/www/myenv/lib/python3.10/site-packages/tensorflow/python/tpu/async_checkpoint.py_get_duration_microseconds1   s   r   c                   @   s   e Zd ZdZ						d"dedee dee deej ded	ee	j
 d
eeej  fddZdd Zdd ZdejdefddZdefddZdejdefddZdejfddZd#ddZd d! ZdS )$AsyncCheckpointSaverHookz+Saves checkpoints every N steps or seconds.N
model.ckptcheckpoint_dir	save_secs
save_stepsr   checkpoint_basenamescaffold	listenersc           	      C   s   t j||}td| |rtdt| |dur#|dur#td|| _d| _d| _	|| _
|| _|| _tj||d| _|p@g | _d| _d| _d| _d| _t tdu rat aW d   dS W d   dS 1 slw   Y  dS )a  Initializes a `CheckpointSaverHook`.

    Args:
      checkpoint_dir: `str`, base directory for the checkpoint files.
      save_secs: `int`, save every N secs.
      save_steps: `int`, save every N steps.
      saver: `Saver` object, used for saving.
      checkpoint_basename: `str`, base name for the checkpoint files.
      scaffold: `Scaffold`, use to get saver object.
      listeners: List of `CheckpointSaverListener` subclass instances. Used for
        callbacks that run immediately before or after this hook saves the
        checkpoint.

    Raises:
      ValueError: One of `save_steps` or `save_secs` should be set.
      ValueError: At most one of `saver` or `scaffold` should be set.
    z1Create AsyncCheckpointSaverHook saving to path
%sz with %d listener(s).Nz+You cannot provide both saver and scaffold.)
every_secsevery_steps   )ospathjoinlogginginfolen
ValueError_saver_save_thread_write_graph_thread_checkpoint_dir
_save_path	_scaffoldr   SecondOrStepTimer_timer
_listeners_steps_per_run_summary_writer_global_step_tensor_last_checkpoint_step_END_TIME_OF_LAST_WRITE_LOCK_END_TIME_OF_LAST_WRITEtime)	selfr   r   r   r   r    r!   r"   	save_pathr   r   r   __init__9   s8   

"z!AsyncCheckpointSaverHook.__init__c                 C   s
   || _ d S N)r6   )r=   steps_per_runr   r   r   _set_steps_per_runo   s   
z+AsyncCheckpointSaverHook._set_steps_per_runc                 C   sB   t | j| _t | _| jd u rtd| jD ]}|	  qd S )Nz9Global step should be created to use CheckpointSaverHook.)
r   getr0   r7   r   _get_or_create_global_step_readr8   RuntimeErrorr5   begin)r=   lr   r   r   rF   r   s   



zAsyncCheckpointSaverHook.beginr   coordc                 C   s   | | j}dd }tj|| gd| _| j  |  r!|  jnd }t	 }t
j|jdd|d}| jd u r;td| j| | j| | || | j| d S )Nc                 S   s    t t jdd| jd d S )NT
add_shapeszgraph.pbtxt)r   write_graphr	   get_default_graphas_graph_defr0   )r=   r   r   r   _write_graph_fn   s   zFAsyncCheckpointSaverHook.after_create_session.<locals>._write_graph_fn)targetargsTrI   )	graph_def	saver_def!Summary writer is not initialised)runr8   	threadingThreadr/   start
_get_saverrR   r	   rL   r   create_meta_graph_defrM   r7   r,   	add_graphadd_meta_graph_saver4   update_last_triggered_step)r=   r   rH   global_steprN   rR   graphmeta_graph_defr   r   r   after_create_session{   s"   

z-AsyncCheckpointSaverHook.after_create_sessionrun_contextc                 C   s   t | jS r@   )r   SessionRunArgsr8   )r=   rb   r   r   r   
before_run   s   z#AsyncCheckpointSaverHook.before_run
run_valuesc                 C   sT   |j | j}| j|r&| j| td| | |j |r(|	  d S d S d S )NzTriggering checkpoint. %s)
r   rT   r8   r4   should_trigger_for_stepr]   r)   r*   r\   request_stop)r=   rb   re   r^   r   r   r   	after_run   s   z"AsyncCheckpointSaverHook.after_runc                 C   sv   | j rtd | j   | jrtd | j  || j}| j|kr-| j||dd | j	D ]}|
|| q0d S )Nz.Waiting for any pending checkpoints to finish.z.Waiting for any pending write_graph to finish.F)asynchronous)r.   r)   r*   r(   r/   rT   r8   r9   r\   r5   end)r=   r   	last_steprG   r   r   r   rj      s   





zAsyncCheckpointSaverHook.endTc                    s   fdd}t     fdd}|s_|  |  dS jdur:jjdd j r:td |  dS _tj|d	_j	  |  dS )
z1Saves the latest checkpoint, returns should_stop.c                     s
  t d j t }  jD ]}| q  j jd  jdu r,t	d j
tjtjj jd  jD ]}| q>t }tjtt| |d t tjttt| d W d   n1 slw   Y  | at d||   t d j dS )	zRun the saver process.z"Saving checkpoints for %d into %s.)r^   NrS   )statuscheckpoint_path	api_labelmicrosecondsz*Checkpoint actual writing time: (%.3f sec)z#Checkpoint finished for %d into %s.)r)   r*   r1   r<   r5   before_saverX   saver7   r,   add_session_logr   
SessionLog
CHECKPOINT
after_saver   AddAsyncCheckpointWriteDuration_ASYNC_CHECKPOINT_V1r   r:   AddTrainingTimeSavedr;   )
start_timerG   end_time)r=   r   stepr   r   _save_fn   sD   



z0AsyncCheckpointSaverHook._save.<locals>._save_fnc                     s    t   } tjtt | d d S )Nrn   )r<   r   AddCheckpointWriteDurationrx   r   )blocking_end_time)blocking_start_timer   r   end_of_blocking_time   s   
z<AsyncCheckpointSaverHook._save.<locals>.end_of_blocking_timeNg?)timeoutz4Saver thread still in progress, skipping checkpoint.)rO   )
r<   r9   r.   r(   is_aliver)   r*   rU   rV   rW   )r=   r   r|   ri   r}   r   r   )r   r=   r   r|   r   r\      s$   +




zAsyncCheckpointSaverHook._savec                 C   sr   | j d ur| j S | jd ur| jjS tjj}t|}|s#td|t	|dkr0td||d | _ |d S )Nz_No items in collection {}. Please add a saver to the collection or provide a saver or scaffold.r%   zgMore than one item in collection {}. Please indicate which one to use by passing it to the constructor.r   )
r-   r2   r   r	   	GraphKeysSAVERSget_collectionrE   formatr+   )r=   collection_keysaversr   r   r   rX      s$   



z#AsyncCheckpointSaverHook._get_saver)NNNr   NN)T)__name__
__module____qualname____doc__r   r   r   	saver_libSaverr   Scaffoldr   r   CheckpointSaverListenerr?   rB   rF   session_libSessionr   ra   rd   r   SessionRunContextrh   rj   r\   rX   r   r   r   r   r   6   sF    
6	
	
Hr   )'r   r&   rU   r<   typingr   r   r   r   tensorflow.core.utilr   tensorflow.python.clientr   r   tensorflow.python.frameworkr   r	   tensorflow.python.platformr
   r)   0tensorflow.python.saved_model.pywrap_saved_modelr   tensorflow.python.trainingr   r   r   r   r   r   %tensorflow.python.training.summary_ior   r;   Lockr:   rx   r   r   CheckpointSaverHookr   r   r   r   r   <module>   s,   