Класс UnitNormalization
Класс UnitNormalization выполняет L2-нормализацию входного тензора
по последней оси. Каждый вектор вдоль последней оси делится на свою
L2-норму, в результате чего норма каждого вектора становится равной
единице. Слой часто применяется в задачах, где важна только
направленность вектора, а не его длина, например в эмбеддингах и
метрическом обучении. Первым параметром передается целое число
axis - ось, по которой вычисляется норма, по умолчанию
равное -1. Вторым параметром можно передать словарь
kwargs с дополнительными именованными аргументами базового
слоя, такими как name и dtype.
Синтаксис
tf.keras.layers.UnitNormalization(axis=-1, **kwargs)
Пример
Давайте создадим слой UnitNormalization и применим его к
двумерному тензору:
import tensorflow as tf
layer = tf.keras.layers.UnitNormalization()
t = tf.constant([[1, 2, 3], [4, 5, 6]], dtype=tf.float32)
res = layer(t)
print(res)
Результат выполнения кода:
tf.Tensor(
[[0.26726124 0.5345225 0.8017837 ]
[0.45584232 0.5698029 0.68376344]], shape=(2, 3), dtype=float32)
Пример
Давайте проверим, что норма каждой строки после нормализации равна единице:
Результат выполнения кода:
tf.Tensor([1. 1.], shape=(2,), dtype=float32)
Пример
Давайте применим слой к трехмерному тензору и укажем ось нормализации явно:
import tensorflow as tf
layer = tf.keras.layers.UnitNormalization(axis=-1)
t = tf.constant([
[[1, 2, 3], [4, 5, 6]],
[[7, 8, 9], [10, 11, 12]]
], dtype=tf.float32)
res = layer(t)
print(res)
Результат выполнения кода:
tf.Tensor(
[[[0.26726124 0.5345225 0.8017837 ]
[0.45584232 0.5698029 0.68376344]]
[[0.5025707 0.5743665 0.64616233]
[0.5455447 0.6000992 0.65465367]]], shape=(2, 2, 3), dtype=float32)
Пример
Давайте встроим слой UnitNormalization в модель Keras:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Input(shape=(3,)),
tf.keras.layers.UnitNormalization(),
tf.keras.layers.Dense(2)
])
t = tf.constant([[1, 2, 3]], dtype=tf.float32)
res = model(t)
print(res)
Результат выполнения кода:
tf.Tensor([[-0.38623512 0.14737517]], shape=(1, 2), dtype=float32)
Смотрите также
-
класс
BatchNormalization,
который нормализует данные по батчу -
класс
LayerNormalization,
который нормализует данные по признакам -
класс
GroupNormalization,
который нормализует данные по группам каналов -
класс
Normalization,
который выполняет адаптивную нормализацию данных