Функция applications.ResNet50
Функция applications.ResNet50 создает и возвращает модель нейронной сети ResNet50. Первым параметром передается форма входных данных input_shape. Вторым параметром можно указать количество классов classes. Третьим параметром задается использование предобученных весов на наборе данных ImageNet через weights. По умолчанию модель загружается с весами 'imagenet' и 1000 классами.
Синтаксис
tf.keras.applications.ResNet50(
input_shape=None,
classes=1000,
weights='imagenet',
include_top=True,
pooling=None
)
Пример
Давайте загрузим модель ResNet50 с предобученными весами без верхнего классификационного слоя:
import tensorflow as tf
model = tf.keras.applications.ResNet50(
input_shape=(224, 224, 3),
include_top=False,
weights='imagenet'
)
print(model.input_shape)
print(model.output_shape)
Результат выполнения кода:
(None, 224, 224, 3)
(None, 7, 7, 2048)
Пример
Давайте создадим модель ResNet50 с собственным количеством классов и проверим форму выходного тензора:
import tensorflow as tf
model = tf.keras.applications.ResNet50(
input_shape=(224, 224, 3),
classes=10,
weights=None
)
print(model.output_shape)
Результат выполнения кода:
(None, 10)
Пример
Давайте выполним предсказание для случайного изображения с помощью загруженной модели:
import tensorflow as tf
import numpy as np
tf.random.set_seed(0)
model = tf.keras.applications.ResNet50(
input_shape=(224, 224, 3),
weights='imagenet'
)
img = tf.random.uniform((1, 224, 224, 3))
preds = model.predict(img)
print(preds.shape)
Результат выполнения кода:
(1, 1000)
Смотрите также
-
функцию
ResNet101,
которая загружает более глубокую версию сети ResNet -
функцию
MobileNetV2,
которая загружает легковесную модель для мобильных устройств -
функцию
VGG16,
которая загружает классическую сверточную сеть VGG16 -
функцию
EfficientNetB0,
которая загружает эффективную модель семейства EfficientNet