Функция segment_sum
Функция segment_sum суммирует элементы тензора
по сегментам вдоль первой оси. Первым параметром
функция принимает тензор с данными. Вторым параметром
передается тензор с идентификаторами сегментов. Третьим
необязательным параметром можно передать количество
сегментов.
Каждый элемент данных относится к сегменту с номером из соответствующей позиции тензора идентификаторов. Все элементы, попавшие в один сегмент, складываются. Результат содержит по одной сумме на каждый сегмент.
Синтаксис
tf.math.segment_sum(data, segment_ids, [num_segments])
Пример
Давайте просуммируем элементы тензора по сегментам
0, 0, 1, 1, 1:
import tensorflow as tf
data = tf.constant([1, 2, 3, 4, 5])
segment_ids = tf.constant([0, 0, 1, 1, 1])
res = tf.math.segment_sum(data, segment_ids)
print(res)
Результат выполнения кода:
tf.Tensor([3 12], shape=(2,), dtype=int32)
Пример
Давайте просуммируем элементы тензора, где сегменты идут не по порядку:
import tensorflow as tf
data = tf.constant([1, 2, 3, 4, 5])
segment_ids = tf.constant([1, 1, 0, 0, 1])
res = tf.math.segment_sum(data, segment_ids)
print(res)
Результат выполнения кода:
tf.Tensor([7 8], shape=(2,), dtype=int32)
Пример
Давайте зададим количество сегментов явно, включив
пустой сегмент с номером 2:
import tensorflow as tf
data = tf.constant([1, 2, 3, 4, 5])
segment_ids = tf.constant([0, 0, 1, 1, 1])
res = tf.math.segment_sum(data, segment_ids, num_segments=3)
print(res)
Результат выполнения кода:
tf.Tensor([3 12 0], shape=(3,), dtype=int32)
Пример
Давайте просуммируем двумерный тензор по сегментам вдоль первой оси:
import tensorflow as tf
data = tf.constant([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
segment_ids = tf.constant([0, 0, 1])
res = tf.math.segment_sum(data, segment_ids)
print(res)
Результат выполнения кода:
tf.Tensor(
[[5 7 9]
[7 8 9]], shape=(2, 3), dtype=int32)
Смотрите также
-
функцию
unsorted_segment_sum,
которая суммирует элементы по неотсортированным сегментам -
функцию
cumsum,
которая вычисляет кумулятивную сумму элементов -
функцию
bincount,
которая подсчитывает количество вхождений значений -
функцию
reduce_sum,
которая суммирует элементы вдоль осей тензора