MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / HashTableSaveable

Class HashTableSaveable

tensorflow/python/training/saver.py:97–155  ·  view source on GitHub ↗

SaveableObject implementation that handles HashTable.

Source from the content-addressed store, hash-verified

95 self._handle)
96
97class HashTableSaveable(saveable_object.SaveableObject):
98 """SaveableObject implementation that handles HashTable."""
99 custom_restore = True
100
101 def __init__(self,
102 hash_tables,
103 restore_tensor_names=None):
104 simple_hash_table = hash_tables[0].hash_table
105 ht = [simple_hash_table] + hash_tables
106 inits = [i.false_initializer for i in hash_tables]
107 self._restore_name = []
108 with ops.colocate_with(ht[0]):
109 if restore_tensor_names is None:
110 names = [i.distributed_name for i in ht]
111 distributed_name = names[0]
112 names[0] += "/ids"
113 if not simple_hash_table.children:
114 self._restore_name.append(";".join(names))
115 else:
116 for child in simple_hash_table.children:
117 child_names = [name.replace(distributed_name, child)
118 for name in names]
119 self._restore_name.append(";".join(child_names))
120 else:
121 names = [restore_tensor_names[0]] + restore_tensor_names
122 names[0] += "/ids"
123 self._restore_name.append(";".join(names))
124 self._restore_name.sort()
125 handles = [i.handle for i in ht]
126 handles = array_ops.stack(handles)
127 self._handles = handles
128 slicex = "{} {},{}".format(ht[0].slicer[2], ht[0].slicer[0],
129 ht[0].slicer[1] - ht[0].slicer[0])
130 spec = saveable_object.SaveSpec(
131 handles, slicex, "|".join(self._restore_name))
132 self._restore_slice = slicex
133 self._restore_specs = []
134 self._inits = inits
135 for i in ht:
136 if i is ht[0]:
137 name = i.distributed_name + "/ids"
138 else:
139 name = i.distributed_name
140 handle = i.handle
141 self._restore_specs.append(saveable_object.SaveSpec(handle, slicex, name))
142 super(HashTableSaveable, self).__init__(
143 handles, [spec], ht[0].name)
144
145 def restore_specs(self):
146 return []
147
148 def restore(self, restored_tensors, restored_shapes):
149 from tensorflow.python.ops.hash_table import hash_table
150 restore_clear = hash_table.restore_clear()
151 with ops.colocate_with(self._handles):
152 with ops.control_dependencies(self._inits):
153 return gen_io_ops.restore_hash_table(
154 restored_tensors[0], self._restore_name, self._restore_slice,

Callers 1

_buildMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected