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

Method _get_embedding_input

modelzoo/din/train.py:231–257  ·  view source on GitHub ↗
(self,
                             builder,
                             feature_column,
                             embedding_table,
                             get_seq_len=False)

Source from the content-addressed store, hash-verified

229 return logits
230
231 def _get_embedding_input(self,
232 builder,
233 feature_column,
234 embedding_table,
235 get_seq_len=False):
236 sparse_tensors = feature_column._get_sparse_tensors(builder)
237 sparse_tensors_ids = sparse_tensors.id_tensor
238 sparse_tensors_weights = sparse_tensors.weight_tensor
239
240 if self._emb_fusion and not self.tf:
241 from tensorflow.python.ops.embedding_ops import fused_safe_embedding_lookup_sparse
242 embedding = fused_safe_embedding_lookup_sparse(
243 embedding_weights=embedding_table,
244 sparse_ids=sparse_tensors_ids,
245 sparse_weights=sparse_tensors_weights)
246 else:
247 embedding = tf.nn.safe_embedding_lookup_sparse(
248 embedding_weights=embedding_table,
249 sparse_ids=sparse_tensors_ids,
250 sparse_weights=sparse_tensors_weights)
251
252 if get_seq_len:
253 sequence_length = fc_utils.sequence_length_from_sparse_tensor(
254 sparse_tensors_ids)
255 return embedding, sequence_length
256 else:
257 return embedding
258
259 def _embedding_input_layer(self):
260 for key in SEQ_COLUMNS:

Callers 1

Calls 2

_get_sparse_tensorsMethod · 0.45

Tested by

no test coverage detected