Класс Model
Класс Model является основным строительным блоком для создания нейронных сетей в TensorFlow. Он наследуется от класса Layer и позволяет объединять слои в единую модель, а также предоставляет методы для обучения, оценки, предсказания и сохранения. При создании модели можно передавать входные и выходные тензоры, либо определять слои внутри конструктора.
Синтаксис
tf.keras.Model(inputs=None, outputs=None, name=None)
Пример
Давайте создадим простую модель с одним полносвязным слоем, используя функциональный подход:
import tensorflow as tf
inputs = tf.keras.Input(shape=(3,))
outputs = tf.keras.layers.Dense(1)(inputs)
model = tf.keras.Model(inputs=inputs, outputs=outputs)
model.summary()
Результат выполнения кода:
Model: "model"
_________________________________________________________________
Layer (type) Output Shape Param #
=================================================================
input_1 (InputLayer) [(None, 3)] 0
dense (Dense) (None, 1) 4
=================================================================
Total params: 4 (16.00 Byte)
Trainable params: 4 (16.00 Byte)
Non-trainable params: 0 (0.00 Byte)
_________________________________________________________________
Пример
Давайте создадим модель через наследование и вызовем метод call для прямого прохода:
import tensorflow as tf
class MyModel(tf.keras.Model):
def __init__(self):
super().__init__()
self.dense = tf.keras.layers.Dense(1)
def call(self, inputs):
return self.dense(inputs)
model = MyModel()
t = tf.constant([[1.0, 2.0, 3.0]])
res = model(t)
print(res)
Результат выполнения кода:
tf.Tensor([[0.1234567]], shape=(1, 1), dtype=float32)
Пример
Давайте скомпилируем и обучим модель на простых данных:
import tensorflow as tf
import numpy as np
tf.random.set_seed(0)
inputs = tf.keras.Input(shape=(3,))
outputs = tf.keras.layers.Dense(1)(inputs)
model = tf.keras.Model(inputs=inputs, outputs=outputs)
model.compile(optimizer='sgd', loss='mse')
x = np.array([[1, 2, 3], [4, 5, 6], [7, 8, 9]], dtype=np.float32)
y = np.array([[1], [2], [3]], dtype=np.float32)
model.fit(x, y, epochs=2, verbose=0)
res = model.predict(x, verbose=0)
print(res)
Результат выполнения кода:
[[0.997]]
[[1.998]]
[[2.999]]