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

Method set_model

tensorflow/python/keras/callbacks_v1.py:232–310  ·  view source on GitHub ↗

Sets Keras model and creates summary ops.

(self, model)

Source from the content-addressed store, hash-verified

230 tf_summary.histogram('{}_out'.format(layer.name), layer.output)
231
232 def set_model(self, model):
233 """Sets Keras model and creates summary ops."""
234
235 self.model = model
236 self._init_writer(model)
237 # histogram summaries only enabled in graph mode
238 if not context.executing_eagerly():
239 self._make_histogram_ops(model)
240 self.merged = tf_summary.merge_all()
241
242 # If both embedding_freq and embeddings_data are available, we will
243 # visualize embeddings.
244 if self.embeddings_freq and self.embeddings_data is not None:
245 # Avoid circular dependency.
246 from tensorflow.python.keras.engine import training_utils # pylint: disable=g-import-not-at-top
247 self.embeddings_data = training_utils.standardize_input_data(
248 self.embeddings_data, model.input_names)
249
250 # If embedding_layer_names are not provided, get all of the embedding
251 # layers from the model.
252 embeddings_layer_names = self.embeddings_layer_names
253 if not embeddings_layer_names:
254 embeddings_layer_names = [
255 layer.name
256 for layer in self.model.layers
257 if type(layer).__name__ == 'Embedding'
258 ]
259
260 self.assign_embeddings = []
261 embeddings_vars = {}
262
263 self.batch_id = batch_id = array_ops.placeholder(dtypes.int32)
264 self.step = step = array_ops.placeholder(dtypes.int32)
265
266 for layer in self.model.layers:
267 if layer.name in embeddings_layer_names:
268 embedding_input = self.model.get_layer(layer.name).output
269 embedding_size = np.prod(embedding_input.shape[1:])
270 embedding_input = array_ops.reshape(embedding_input,
271 (step, int(embedding_size)))
272 shape = (self.embeddings_data[0].shape[0], int(embedding_size))
273 embedding = variables.Variable(
274 array_ops.zeros(shape), name=layer.name + '_embedding')
275 embeddings_vars[layer.name] = embedding
276 batch = state_ops.assign(embedding[batch_id:batch_id + step],
277 embedding_input)
278 self.assign_embeddings.append(batch)
279
280 self.saver = saver.Saver(list(embeddings_vars.values()))
281
282 # Create embeddings_metadata dictionary
283 if isinstance(self.embeddings_metadata, str):
284 embeddings_metadata = {
285 layer_name: self.embeddings_metadata
286 for layer_name in embeddings_vars.keys()
287 }
288 else:
289 # If embedding_metadata is already a dictionary

Callers

nothing calls this directly

Calls 13

_init_writerMethod · 0.95
_make_histogram_opsMethod · 0.95
typeFunction · 0.85
executing_eagerlyMethod · 0.80
reshapeMethod · 0.80
VariableMethod · 0.80
placeholderMethod · 0.45
get_layerMethod · 0.45
assignMethod · 0.45
appendMethod · 0.45
valuesMethod · 0.45
keysMethod · 0.45

Tested by

no test coverage detected