Класс SGD
Класс SGD представляет собой оптимизатор, который реализует алгоритм стохастического градиентного спуска. Первым параметром передается скорость обучения learning_rate, вторым - коэффициент импульса momentum, третьим - флаг использования импульса Нестерова nesterov. Оптимизатор применяется к модели через метод compile, где указывается функция потерь и метрики.
Синтаксис
tf.keras.optimizers.SGD(learning_rate=0.01, momentum=0.0, nesterov=False)
Пример
Давайте создадим простую модель и скомпилируем ее с оптимизатором SGD:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
optimizer = tf.keras.optimizers.SGD(learning_rate=0.01)
model.compile(optimizer=optimizer, loss='mse')
print(model.optimizer)
Результат выполнения кода:
<keras.src.optimizers.sgd.SGD object at 0x7f8b8c0b4d30>
Пример
Давайте создадим оптимизатор SGD с импульсом и обучим модель на простых данных:
import tensorflow as tf
import numpy as np
tf.random.set_seed(0)
x = np.array([1, 2, 3, 4, 5], dtype=np.float32)
y = np.array([2, 4, 6, 8, 10], dtype=np.float32)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
optimizer = tf.keras.optimizers.SGD(learning_rate=0.01, momentum=0.9)
model.compile(optimizer=optimizer, loss='mse')
history = model.fit(x, y, epochs=100, verbose=0)
print(f"Final loss: {history.history['loss'][-1]:.6f}")
Результат выполнения кода:
"Final loss: 0.000124"
Пример
Давайте создадим оптимизатор SGD с импульсом Нестерова и выведем его конфигурацию:
import tensorflow as tf
tf.random.set_seed(0)
optimizer = tf.keras.optimizers.SGD(
learning_rate=0.001,
momentum=0.9,
nesterov=True
)
config = optimizer.get_config()
print(config)
Результат выполнения кода:
{'name': 'SGD', 'learning_rate': 0.001, 'momentum': 0.9, 'nesterov': True, 'weight_decay': None, 'clipnorm': None, 'global_clipnorm': None, 'clipvalue': None, 'use_ema': False, 'ema_momentum': 0.99, 'ema_overwrite_frequency': None, 'loss_scale_factor': None, 'gradient_accumulation_steps': None}