o
    i e8                     @   s  d Z ddl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dlmZ edg dZd"ddZdd Zdd Z					d#ddZdd Z			d$ddZdd Z				d%d d!Z dS )&zExtracts tensors for checkpointing while updating a TrackableObjectGraph.

This is labelled "v1" because the methods here use SaveableObject, which will
soon be deprecated.
    N)trackable_object_graph_pb2)saveable_compat)util)constant_op)dtypes)ops)registration)base)python_state)trackable_utils)saveable_object)saveable_object_util)object_identity_CheckpointFactoryDatafactorynamecheckpoint_keyc              	   C   s   t  }tt}|  D ]G\}}t||}t	|}|r%||| |< qg ||< t
| D ]#\}}	t|p:|}
t||
}t sG|
}|| t|	||d q0q||fS )a  Gets a map of saveable factories and corresponding checkpoint keys.

  Args:
    object_names: a dictionary that maps `Trackable` objects to auto-generated
      string names.
    object_map: a dictionary mapping `Trackable` to copied `Trackable` objects.
      The copied objects are generated from `Trackable.
      _export_to_saved_model_graph()` which copies the object into another
      graph. Generally only resource objects (e.g. Variables, Tables) will be
      in this map.

  Returns:
    A tuple of (
      Dictionary mapping trackable -> list of _CheckpointFactoryData,
      Dictionary mapping registered saver name -> {object name -> trackable})
  r   )r   ObjectIdentityDictionarycollectionsdefaultdictdictitemsr   get_mapped_trackabler   get_registered_saver_namer   saveable_objects_from_trackabler   get_saveable_namer   r   #force_checkpoint_conversion_enabledappendr   )object_names
object_mapcheckpoint_factory_mapunmapped_registered_savers	trackableobject_nameobject_to_save
saver_namer   saveable_factory
key_suffixr    r)   X/var/www/myenv/lib/python3.10/site-packages/tensorflow/python/checkpoint/save_util_v1.py!get_checkpoint_factories_and_keys+   s4   


r+   c                 C   sj   t t| |jD ]\}\}}	|| |ksJ qt||\}
}t||||}t|
|||||\}}|||fS )zECreate saveables/savers and corresponding protos in the object graph.)	enumeratezipnodesr+   5_add_attributes_to_object_graph_for_registered_saversgenerate_saveable_objects)trackable_objectsobject_graph_protonode_idsr   r    call_with_mapped_capturessaveables_cachecheckpoint_idr#   unused_object_protor!   r"   registered_saversnamed_saveable_objectsfeed_additionsr)   r)   r*   _add_attributes_to_object_graph`   s   

r;   c                 C   sh   t t}|  D ](\}}| D ]\}}|j||  }	||	j_||	j_t	||}
|
|| |< qq	|S )zCFills the object graph proto with data about the registered savers.)
r   r   r   r   r.   registered_saverr   r$   r   r   )r"   r2   r3   r    r8   r&   
trackablesr$   r#   object_protor%   r)   r)   r*   r/   x   s   
r/   c                 C   s:  g }|du r	d}ni }|   D ]\}}	|duo|du}
|
r%|j||  }t||}|dur6||i }nd}|	D ]}|j}|j}|j}|rL||nd}|durc|D ]}||jvrbd}||=  nqT|du rt	|rtt
||||}n|}t|tjr|f}n	tt
j||d}|D ]}||jvrtd| d|j d| d| d	q|dur|||< t|tjrt|dksJ |d	 }|du r|du sJ | f}n||  || |
sq:t|d	 t
jrt st|du r|d	  D ]\}}|jj||t |d
 qq:|jj||t |d
 q:q||fS )zACreate SaveableObjects and corresponding SerializedTensor protos.N)opr   zThe object z& produced a SaveableObject with name 'z' for attribute 'z'. Expected a name containing 'z'.   r   )r   r   	full_name)!r   r.   r   r   
setdefaultr   r   r   getcallabler   create_saveable_object
isinstancesaveable_object_libSaveableObjecttuplesaveable_objects_for_opAssertionErrorr
   PythonStatelenfreezeupdatefeed_dict_additionsextendTrackableSaveabler   r   r   #get_proto_names_and_checkpoint_keys
attributesaddget_full_name)r!   r2   r3   r    r4   r5   r9   r:   r#   factory_data_listfill_object_protor>   r%   cached_attributesfactory_datar   keyr'   	saveablessaveablemaybe_saveable
local_name	local_keyr)   r)   r*   r0      s   




Jr0   c           	      C   sl   t  }t|D ]+\}}|| |ksJ |jj||dd}| |D ]}|jj||j |j	d q$q|S )z@Name non-slot `Trackable`s and add them to `object_graph_proto`.r)   )slot_variables)node_idr_   )
r   TrackableObjectGraphr,   r.   rU   rC   list_childrenchildrenrefr   )	
graph_viewr1   r3   ra   r2   r6   r#   r>   childr)   r)   r*   _fill_object_graph_proto   s   
ri   c              	   C   s   |   \}}t }| D ]\}}t|||< qt }	t|D ]\}
}|
|	|< q"tj||	|d}t	| ||	|d}t
|||	||||d\}}}t| ||||fS )z7Create SaveableObjects and protos for gathered objects.)r1   r3   r   )rg   r1   r3   ra   )r1   r2   r3   r   r    r4   r5   )breadth_first_traversalr   r   r   r   object_path_to_stringr,   r   serialize_slot_variablesri   r;   add_checkpoint_values_check)rg   r    r4   r5   r1   
node_pathsr   objpathr3   rb   nodera   r2   r9   r:   r8   r)   r)   r*   serialize_gathered_objects  s@   

rr   c                 C   s   t | |dS )zEDetermine checkpoint keys for variables and build a serialized graph.)r5   )rr   )rg   r5   r)   r)   r*   -serialize_object_graph_with_registered_savers&  s   rs   c              	   C   s   |r|j }ntj}| @ t| |||\}}}}	td tj| tj	d}
W d   n1 s2w   Y  |
tj|
tjd W d   ||	fS 1 sOw   Y  ||	fS )zDGenerates SaveableObjects and registered savers in the frozen graph.z/cpu:0)dtypeN)tensorr   )
as_defaultr   NullContextmanagerrr   devicer   constantSerializeToStringr   stringr   r	   NoRestoreSaveableOBJECT_GRAPH_PROTO_KEY)rg   r    to_graphr4   r5   target_contextr9   graph_proto_r8   object_graph_tensorr)   r)   r*   frozen_saveables_and_savers+  s,   




r   )N)NNNNN)NNN)NNNN)!__doc__r   tensorflow.core.protobufr   tensorflow.python.checkpointr   r   tensorflow.python.frameworkr   r   r   tensorflow.python.saved_modelr   tensorflow.python.trackabler	   r
   r   !tensorflow.python.training.savingr   rG   r   tensorflow.python.utilr   
namedtupler   r+   r;   r/   r0   ri   rr   rs   r   r)   r)   r)   r*   <module>   sL   
5
h
%