Функция plot_model
Функция plot_model строит визуальное представление архитектуры модели Keras в виде графа слоев и сохраняет его в файл изображения или PDF. Первым параметром функция принимает саму модель. Вторым параметром передается путь к файлу, в который будет сохранена схема. Третьим параметром можно указать, нужно ли отображать формы тензоров на каждом слое. Дополнительно можно управлять раскладкой графа слева направо и показывать типы слоев.
Синтаксис
tf.keras.utils.plot_model(
model,
to_file="model.png",
show_shapes=False,
show_dtype=False,
show_layer_names=True,
rankdir="TB",
expand_nested=False,
dpi=96,
show_layer_activations=False,
show_trainable=False
)
Пример
Давайте создадим простую последовательную модель и сохраним ее архитектуру в файл:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(64, activation="relu", input_shape=(10,)),
tf.keras.layers.Dense(32, activation="relu"),
tf.keras.layers.Dense(1, activation="sigmoid")
])
tf.keras.utils.plot_model(model, to_file="model.png")
print("model plotted")
Результат выполнения кода:
"model plotted"
Пример
Давайте сохраним схему модели с отображением форм тензоров на каждом слое:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(64, activation="relu", input_shape=(10,)),
tf.keras.layers.Dense(32, activation="relu"),
tf.keras.layers.Dense(1, activation="sigmoid")
])
tf.keras.utils.plot_model(
model,
to_file="model_shapes.png",
show_shapes=True
)
print("model plotted with shapes")
Результат выполнения кода:
"model plotted with shapes"
Пример
Давайте построим горизонтальную схему модели с отображением активаций и типов данных:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(64, activation="relu", input_shape=(10,)),
tf.keras.layers.Dense(32, activation="relu"),
tf.keras.layers.Dense(1, activation="sigmoid")
])
tf.keras.utils.plot_model(
model,
to_file="model_horizontal.png",
show_shapes=True,
show_dtype=True,
show_layer_activations=True,
rankdir="LR"
)
print("model plotted horizontally")
Результат выполнения кода:
"model plotted horizontally"
Смотрите также
-
функцию
save_model,
которая сохраняет модель Keras в файл -
функцию
load_model,
которая загружает модель Keras из файла -
функцию
clone_model,
которая создает копию архитектуры модели -
функцию
model_from_json,
которая создает модель из JSON-описания