Класс Lion
Класс Lion представляет оптимизатор Lion,
опубликованный в 2023 году как альтернатива AdamW.
Название расшифровывается как EvoLved Sign Momentum.
В отличие от Adam, Lion хранит только момент импульса
и применяет знак градиента для обновления параметров.
Первым параметром передаётся скорость обучения
learning_rate, вторым - коэффициент
beta_1 для момента импульса, третьим -
beta_2 для второго момента. Также можно
передать weight_decay для регуляризации.
Синтаксис
tf.keras.optimizers.Lion(
learning_rate=0.001,
beta_1=0.9,
beta_2=0.99,
weight_decay=0.0
)
Пример
Давайте создадим оптимизатор Lion со скоростью
обучения 0.01 и применим его к простой
модели:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,))
])
optimizer = tf.keras.optimizers.Lion(
learning_rate=0.01
)
model.compile(
optimizer=optimizer,
loss='mse'
)
x = tf.constant([[1.0, 2.0, 3.0]])
y = tf.constant([[1.0, 0.0]])
loss = model.train_on_batch(x, y)
print(loss)
Результат выполнения кода:
0.8345212
Пример
Давайте настроим параметры beta_1,
beta_2 и weight_decay:
import tensorflow as tf
tf.random.set_seed(0)
optimizer = tf.keras.optimizers.Lion(
learning_rate=0.001,
beta_1=0.95,
beta_2=0.98,
weight_decay=0.01
)
var = tf.Variable([1.0, 2.0, 3.0])
with tf.GradientTape() as tape:
loss = tf.reduce_sum(var ** 2)
grads = tape.gradient(loss, [var])
optimizer.apply_gradients(zip(grads, [var]))
print(var.numpy())
Результат выполнения кода:
[0.999 1.999 2.999]
Пример
Давайте обучим модель на нескольких эпохах с оптимизатором Lion:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(4, activation='relu', input_shape=(3,)),
tf.keras.layers.Dense(1)
])
model.compile(
optimizer=tf.keras.optimizers.Lion(learning_rate=0.01),
loss='mse'
)
x = tf.constant([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]])
y = tf.constant([[1.0], [0.0]])
history = model.fit(x, y, epochs=3, verbose=0)
print(history.history['loss'])
Результат выполнения кода:
[7.8324165, 7.7410927, 7.6503243]