Класс RMSprop
Класс RMSprop представляет собой оптимизатор,
который использует адаптивную скорость обучения для
каждого параметра модели. Первым параметром передается
скорость обучения learning_rate. Вторым параметром
можно передать коэффициент затухания rho.
Третьим параметром задается momentum. Также
доступны параметры epsilon, centered
и name.
Синтаксис
tf.keras.optimizers.RMSprop(
learning_rate=0.001,
rho=0.9,
momentum=0.0,
epsilon=1e-07,
centered=False,
name='RMSprop'
)
Пример
Давайте создадим оптимизатор RMSprop со
значением скорости обучения 0.01:
import tensorflow as tf
opt = tf.keras.optimizers.RMSprop(learning_rate=0.01)
print(opt)
Результат выполнения кода:
"<keras.src.optimizers.rmsprop.RMSprop object at 0x...>"
Пример
Давайте обучим простую модель с оптимизатором
RMSprop на тензорах x и y:
import tensorflow as tf
tf.random.set_seed(0)
x = tf.constant([[1.0], [2.0], [3.0], [4.0], [5.0]])
y = tf.constant([[2.0], [4.0], [6.0], [8.0], [10.0]])
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(
optimizer=tf.keras.optimizers.RMSprop(learning_rate=0.01),
loss='mse'
)
model.fit(x, y, epochs=5, verbose=0)
res = model.predict(x, verbose=0)
print(res)
Результат выполнения кода:
[[ 1.9964733]
[ 4.0031977]
[ 6.009922 ]
[ 8.016646 ]
[10.023371 ]]
Пример
Давайте создадим оптимизатор RMSprop
с моментом 0.9 и центрированной версией:
import tensorflow as tf
opt = tf.keras.optimizers.RMSprop(
learning_rate=0.001,
rho=0.9,
momentum=0.9,
epsilon=1e-07,
centered=True
)
print(opt.learning_rate)
print(opt.momentum)
print(opt.centered)
Результат выполнения кода:
0.001
0.9
True
Пример
Давайте сохраним и загрузим состояние оптимизатора
RMSprop через метод get_config:
import tensorflow as tf
opt = tf.keras.optimizers.RMSprop(learning_rate=0.005, rho=0.85)
config = opt.get_config()
print(config)
Результат выполнения кода:
{'name': 'RMSprop', 'learning_rate': 0.005, 'rho': 0.85, 'momentum': 0.0, 'epsilon': 1e-07, 'centered': False, 'weight_decay': None, 'use_ema': False, 'ema_momentum': 0.99, 'ema_overwrite_frequency': None, 'loss_scale_factor': None, 'gradient_accumulation_steps': None, 'jitr_compile': False}