SaveableObject implementation that handles HashTable.
| 95 | self._handle) |
| 96 | |
| 97 | class 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, |