Функция strided_slice
Функция strided_slice извлекает срез из тензора.
Первым параметром функция принимает исходный тензор.
Вторым параметром передается список начальных индексов begin,
третьим - список конечных индексов end,
четвертым - список шагов strides.
В отличие от обычного среза, функция поддерживает шаг
для каждого измерения и позволяет использовать отрицательные индексы.
Синтаксис
tf.strided_slice(input_, begin, end, [strides], [begin_mask], [end_mask], [ellipsis_mask], [new_axis_mask], [shrink_axis_mask])
Пример
Давайте извлечем первые три элемента из тензора:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
res = tf.strided_slice(t, [0], [3], [1])
print(res)
Результат выполнения кода:
tf.Tensor([1 2 3], shape=(3,), dtype=int32)
Пример
Давайте извлечем элементы с шагом 2:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
res = tf.strided_slice(t, [0], [5], [2])
print(res)
Результат выполнения кода:
tf.Tensor([1 3 5], shape=(3,), dtype=int32)
Пример
Давайте извлечем строки и столбцы из двумерного тензора:
import tensorflow as tf
t = tf.constant([[1, 2, 3], [4, 5, 6]])
res = tf.strided_slice(t, [0, 0], [2, 2], [1, 1])
print(res)
Результат выполнения кода:
tf.Tensor(
[[1 2]
[4 5]], shape=(2, 2), dtype=int32)
Пример
Давайте используем отрицательный шаг для разворота тензора:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
res = tf.strided_slice(t, [4], [-1], [-1])
print(res)
Результат выполнения кода:
tf.Tensor([5 4 3 2 1], shape=(5,), dtype=int32)