Метод fit класса Sequential
Метод fit класса Sequential запускает процесс обучения модели.
Первым параметром передаются входные данные x, вторым - целевые значения y.
Третьим параметром указывается число эпох epochs - сколько раз модель
пройдёт по всему набору данных. Также можно задать размер батча batch_size,
долю данных для валидации validation_split и другие настройки.
Метод возвращает объект History с историей обучения.
Синтаксис
model.fit(x, y, [epochs], [batch_size], [validation_split], [verbose])
Пример
Давайте создадим простую модель и обучим её на данных
1, 2, 3, 4, 5 с целевыми значениями
2, 4, 6, 8, 10:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer='sgd', loss='mse')
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])
history = model.fit(x, y, epochs=10, verbose=0)
print(history.history['loss'][-1])
Результат выполнения кода:
0.032178539
Пример
Давайте обучим модель с размером батча 2 и посмотрим итоговую ошибку:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer='sgd', loss='mse')
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])
history = model.fit(x, y, epochs=20, batch_size=2, verbose=0)
print(history.history['loss'][-1])
Результат выполнения кода:
0.005412271
Пример
Давайте обучим модель с выделением части данных для валидации:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer='sgd', loss='mse')
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])
history = model.fit(x, y, epochs=10, validation_split=0.2, verbose=0)
print(history.history['loss'][-1])
Результат выполнения кода:
0.037845086
Смотрите также
-
класс
Sequential,
который представляет линейный стек слоёв -
метод
compile,
который настраивает модель для обучения -
метод
predict,
который выполняет предсказание на новых данных -
метод
summary,
который выводит структуру модели