Функция split
Функция split разделяет тензор на несколько
тензоров вдоль заданной оси. Первым параметром
функция принимает тензор, который нужно разделить.
Вторым параметром передается количество частей
num_or_size_splits, которое может быть целым
числом или списком размеров. Третьим параметром
указывается ось axis, вдоль которой выполняется
разделение.
Синтаксис
tf.split(value, num_or_size_splits, [axis])
Пример
Давайте разделим тензор из пяти чисел на пять
отдельных частей вдоль оси 0:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
res = tf.split(t, 5)
print(res)
Результат выполнения кода:
[<tf.Tensor: shape=(1,), dtype=int32, numpy=array([1])>, <tf.Tensor: shape=(1,), dtype=int32, numpy=array([2])>, <tf.Tensor: shape=(1,), dtype=int32, numpy=array([3])>, <tf.Tensor: shape=(1,), dtype=int32, numpy=array([4])>, <tf.Tensor: shape=(1,), dtype=int32, numpy=array([5])>]
Пример
Давайте разделим тензор из пяти чисел на две
неравные части с размерами 2 и 3:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
res = tf.split(t, [2, 3])
print(res)
Результат выполнения кода:
[<tf.Tensor: shape=(2,), dtype=int32, numpy=array([1, 2])>, <tf.Tensor: shape=(3,), dtype=int32, numpy=array([3, 4, 5])>]
Пример
Давайте разделим двумерный тензор на две части
вдоль оси 1:
import tensorflow as tf
t = tf.constant([[1, 2, 3], [4, 5, 6]])
res = tf.split(t, 3, axis=1)
print(res)
Результат выполнения кода:
[<tf.Tensor: shape=(2, 1), dtype=int32, numpy=array([[1], [4]])>, <tf.Tensor: shape=(2, 1), dtype=int32, numpy=array([[2], [5]])>, <tf.Tensor: shape=(2, 1), dtype=int32, numpy=array([[3], [6]])>]