Функция linalg.band_part
Функция linalg.band_part применяется к тензорам и позволяет выделить центральную диагональную полосу, обнуляя все остальные элементы. Первым параметром передается входной тензор. Вторым параметром задается количество сохраняемых поддиагоналей (ниже главной диагонали). Третьим параметром указывается количество сохраняемых наддиагоналей (выше главной диагонали). Если передать значение -1, то соответствующая часть полосы сохраняется полностью.
Синтаксис
tf.linalg.band_part(input, num_lower, num_upper)
Пример
Давайте создадим квадратную матрицу и выделим в ней главную диагональ, обнулив все остальные элементы:
import tensorflow as tf
t = tf.constant([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
res = tf.linalg.band_part(t, 0, 0)
print(res)
Результат выполнения кода:
tf.Tensor(
[[1 0 0]
[0 5 0]
[0 0 9]], shape=(3, 3), dtype=int32)
Пример
Давайте сохраним одну поддиагональ и одну наддиагональ, чтобы получить трехдиагональную матрицу:
import tensorflow as tf
t = tf.constant([[1, 2, 3], [4, 5, 6], [7, 8, 9]])
res = tf.linalg.band_part(t, 1, 1)
print(res)
Результат выполнения кода:
tf.Tensor(
[[1 2 0]
[4 5 6]
[0 8 9]], shape=(3, 3), dtype=int32)
Пример
Давайте сохраним все нижние диагонали, но обнулим все элементы выше главной диагонали:
Результат выполнения кода:
tf.Tensor(
[[1 0 0]
[4 5 0]
[7 8 9]], shape=(3, 3), dtype=int32)