Класс Nadam
Класс Nadam создает оптимизатор, объединяющий
алгоритм Adam с импульсом Нестерова. Первым параметром
передается скорость обучения learning_rate, вторым -
коэффициент экспоненциального затухания для первого момента
beta_1, третьим - коэффициент для второго момента
beta_2, четвертым - малое число для стабильности
вычислений epsilon. Оптимизатор применяется к модели
через метод compile.
Синтаксис
tf.keras.optimizers.Nadam(
learning_rate=0.001,
beta_1=0.9,
beta_2=0.999,
epsilon=1e-07,
weight_decay=None,
clipnorm=None,
clipvalue=None,
global_clipnorm=None,
use_ema=False,
ema_momentum=0.99,
ema_overwrite_frequency=None,
loss_scale_factor=None,
gradient_accumulation_steps=None,
name="nadam",
**kwargs
)
Пример
Давайте создадим оптимизатор Nadam и применим его для обучения простой модели:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(
optimizer=tf.keras.optimizers.Nadam(learning_rate=0.01),
loss='mse'
)
x = tf.constant([[1.0], [2.0], [3.0], [4.0]])
y = tf.constant([[2.0], [4.0], [6.0], [8.0]])
history = model.fit(x, y, epochs=50, verbose=0)
print(round(history.history['loss'][-1], 4))
Результат выполнения кода:
0.0
Пример
Давайте создадим оптимизатор Nadam с измененными коэффициентами затухания и выведем их значения:
import tensorflow as tf
opt = tf.keras.optimizers.Nadam(
learning_rate=0.005,
beta_1=0.85,
beta_2=0.995,
epsilon=1e-08
)
print(opt.learning_rate)
print(opt.beta_1)
print(opt.beta_2)
print(opt.epsilon)
Результат выполнения кода:
0.005
0.85
0.995
1e-08
Пример
Давайте применим Nadam для обучения модели на двумерных данных:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(3, activation='relu', input_shape=(3,)),
tf.keras.layers.Dense(1)
])
model.compile(
optimizer=tf.keras.optimizers.Nadam(),
loss='mse'
)
x = tf.constant([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]])
y = tf.constant([[1.0], [2.0]])
history = model.fit(x, y, epochs=10, verbose=0)
print(round(history.history['loss'][-1], 4))
Результат выполнения кода:
0.2053