| 388 | return tf.group(*update_ops) |
| 389 | |
| 390 | def _optimize_store(self, node, node_embeddings): |
| 391 | if not self.gradient_stores: |
| 392 | return tf.zeros([]), tf.no_op() |
| 393 | |
| 394 | losses = [] |
| 395 | clear_ops = [] |
| 396 | for gradient_store, node_embedding in zip( |
| 397 | self.gradient_stores, node_embeddings): |
| 398 | embedding_gradient = tf.nn.embedding_lookup(gradient_store, node) |
| 399 | with tf.control_dependencies([embedding_gradient]): |
| 400 | clear_ops.append( |
| 401 | utils_embedding.embedding_update( |
| 402 | gradient_store, node, |
| 403 | tf.zeros_like(embedding_gradient))) |
| 404 | losses.append(tf.reduce_sum(node_embedding * embedding_gradient)) |
| 405 | |
| 406 | store_loss = tf.add_n(losses) |
| 407 | with tf.control_dependencies(clear_ops): |
| 408 | return store_loss, self.store_optimizer.minimize(store_loss) |
| 409 | |
| 410 | |
| 411 | class SageEncoder(layers.Layer): |