Функция nn.embedding_lookup_sparse
Функция nn.embedding_lookup_sparse извлекает строки из матрицы эмбеддингов по разреженным индексам, заданным в объекте SparseTensor. Первым параметром передается матрица эмбеддингов (или список матриц), вторым - разреженный тензор с идентификаторами, третьим - способ агрегации значений для одной строки (например, "sum", "mean" или "sqrtn"). Дополнительно можно передать веса элементов и имя операции.
Синтаксис
tf.nn.embedding_lookup_sparse(params, sp_ids, sp_weights=None, combiner=None, max_norm=None, name=None)
Пример
Давайте создадим матрицу эмбеддингов и разреженный тензор с индексами, после чего извлечем и просуммируем строки:
import tensorflow as tf
params = tf.constant([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
sp_ids = tf.sparse.SparseTensor(
indices=[[0, 0], [1, 0]],
values=[0, 2],
dense_shape=[2, 2]
)
res = tf.nn.embedding_lookup_sparse(params, sp_ids, combiner="sum")
print(res)
Результат выполнения кода:
tf.Tensor(
[[1. 2.]
[5. 6.]], shape=(2, 2), dtype=float32)
Пример
Давайте используем комбинатор "mean", чтобы усреднить значения для каждой строки:
import tensorflow as tf
params = tf.constant([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
sp_ids = tf.sparse.SparseTensor(
indices=[[0, 0], [0, 1], [1, 0]],
values=[0, 1, 2],
dense_shape=[2, 2]
)
res = tf.nn.embedding_lookup_sparse(params, sp_ids, combiner="mean")
print(res)
Результат выполнения кода:
tf.Tensor(
[[2. 3.]
[5. 6.]], shape=(2, 2), dtype=float32)
Пример
Давайте передадим веса для каждого идентификатора и используем комбинатор "sum":
import tensorflow as tf
params = tf.constant([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
sp_ids = tf.sparse.SparseTensor(
indices=[[0, 0], [0, 1]],
values=[0, 1],
dense_shape=[1, 2]
)
sp_weights = tf.sparse.SparseTensor(
indices=[[0, 0], [0, 1]],
values=[0.5, 2.0],
dense_shape=[1, 2]
)
res = tf.nn.embedding_lookup_sparse(params, sp_ids, sp_weights, combiner="sum")
print(res)
Результат выполнения кода:
tf.Tensor([[6.5 9. ]], shape=(1, 2), dtype=float32)
Смотрите также
-
функцию
embedding_lookup,
которая извлекает эмбеддинги по плотным индексам -
функцию
safe_embedding_lookup_sparse,
которая выполняет безопасный поиск эмбеддингов по разреженным данным -
функцию
l2_normalize,
которая нормализует тензоры по L2-норме -
функцию
dropout,
которая применяет dropout к элементам тензора