Класс TruncatedNormal
Класс TruncatedNormal представляет собой инициализатор,
который заполняет тензор значениями из усеченного нормального
распределения. В отличие от обычного нормального распределения,
здесь значения, выходящие за пределы двух стандартных отклонений
от среднего, отбрасываются и генерируются заново. Это позволяет
избежать редких, но крайне больших или малых значений, которые
могут дестабилизировать обучение нейронной сети.
Первым параметром передается среднее значение распределения
mean, вторым - стандартное отклонение stddev.
Дополнительно можно указать зерно генератора seed
для воспроизводимости результатов.
Синтаксис
tf.keras.initializers.TruncatedNormal(mean=0.0, stddev=0.05, seed=None)
Пример
Давайте создадим инициализатор усеченного нормального
распределения и сгенерируем тензор формы
3 на 4:
import tensorflow as tf
tf.random.set_seed(0)
init = tf.keras.initializers.TruncatedNormal(mean=0.0, stddev=1.0, seed=0)
t = init(shape=(3, 4))
print(t)
Результат выполнения кода:
tf.Tensor(
[[-0.20176056 1.8500295 -0.6113055 1.3166363 ]
[ 1.3852261 -0.5459486 -0.03916568 -1.0382727 ]
[-0.71656096 -1.3244755 0.4312024 -0.9926889 ]], shape=(3, 4), dtype=float32)
Пример
Давайте используем инициализатор TruncatedNormal
в полносвязном слое Dense:
import tensorflow as tf
tf.random.set_seed(0)
layer = tf.keras.layers.Dense(
units=3,
kernel_initializer=tf.keras.initializers.TruncatedNormal(mean=0.0, stddev=0.05, seed=0)
)
t = tf.constant([[1.0, 2.0, 3.0]])
res = layer(t)
print(res)
Результат выполнения кода:
tf.Tensor([[-0.02112615 0.08439481 0.03521308]], shape=(1, 3), dtype=float32)
Пример
Давайте сравним усеченное нормальное распределение с обычным, посмотрев на минимальное и максимальное значения:
import tensorflow as tf
tf.random.set_seed(0)
trunc = tf.keras.initializers.TruncatedNormal(mean=0.0, stddev=1.0, seed=0)
normal = tf.keras.initializers.RandomNormal(mean=0.0, stddev=1.0, seed=0)
t_trunc = trunc(shape=(1000,))
t_normal = normal(shape=(1000,))
print("Truncated min:", tf.reduce_min(t_trunc).numpy())
print("Truncated max:", tf.reduce_max(t_trunc).numpy())
print("Normal min:", tf.reduce_min(t_normal).numpy())
print("Normal max:", tf.reduce_max(t_normal).numpy())
Результат выполнения кода:
Truncated min: -1.9999999
Truncated max: 1.9999999
Normal min: -3.008108
Normal max: 2.9746094
Смотрите также
-
класс
RandomNormal,
который генерирует значения из нормального распределения -
класс
RandomUniform,
который генерирует значения из равномерного распределения -
класс
GlorotNormal,
который реализует инициализацию Ксавье с нормальным распределением -
класс
HeNormal,
который реализует инициализацию Хе с нормальным распределением