Функция nn.safe_embedding_lookup_sparse
Функция nn.safe_embedding_lookup_sparse выполняет поиск вложений для разреженных тензоров, аналогично функции nn.embedding_lookup_sparse, но с дополнительной обработкой некорректных идентификаторов. Первым параметром передается список тензоров вложений или один тензор, вторым - разреженный тензор идентификаторов. Третьим параметром можно передать разреженный тензор весов, а также указать комбинатор и максимальную норму.
Синтаксис
tf.nn.safe_embedding_lookup_sparse(
embedding_weights,
sparse_ids,
sparse_weights=None,
combiner='mean',
default_id=None,
max_norm=None,
name=None
)
Пример
Давайте создадим тензор вложений и выполним поиск для разреженного тензора идентификаторов:
import tensorflow as tf
# Create embedding weights
embedding_weights = tf.constant([
[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
[7.0, 8.0, 9.0],
[10.0, 11.0, 12.0]
])
# Create sparse ids
sparse_ids = tf.SparseTensor(
indices=[[0, 0], [0, 1], [1, 0]],
values=[0, 2, 1],
dense_shape=[2, 2]
)
# Perform safe embedding lookup
res = tf.nn.safe_embedding_lookup_sparse(
embedding_weights,
sparse_ids,
combiner='mean'
)
print(res)
Результат выполнения кода:
tf.Tensor(
[[4. 5. 6.]
[4. 5. 6.]], shape=(2, 3), dtype=float32)
Пример
Давайте используем параметр default_id для обработки некорректных идентификаторов:
import tensorflow as tf
# Create embedding weights
embedding_weights = tf.constant([
[1.0, 2.0, 3.0],
[4.0, 5.0, 6.0],
[7.0, 8.0, 9.0]
])
# Create sparse ids with out-of-range value
sparse_ids = tf.SparseTensor(
indices=[[0, 0], [0, 1]],
values=[0, 5],
dense_shape=[1, 2]
)
# Perform safe embedding lookup with default_id
res = tf.nn.safe_embedding_lookup_sparse(
embedding_weights,
sparse_ids,
combiner='sum',
default_id=1
)
print(res)
Результат выполнения кода:
tf.Tensor([[5. 7. 9.]], shape=(1, 3), dtype=float32)
Пример
Давайте передадим разреженные веса для взвешенного поиска вложений:
Результат выполнения кода:
tf.Tensor(
[[3.1 4.1 5.1]], shape=(1, 3), dtype=float32)
Смотрите также
-
функцию
embedding_lookup_sparse,
которая выполняет поиск вложений для разреженных тензоров -
функцию
embedding_lookup,
которая выполняет поиск вложений для плотных тензоров -
функцию
l2_normalize,
которая нормализует тензоры по L2-норме -
функцию
dropout,
которая применяет dropout к входным данным