tf中自定义层实战(CIFAR10图像识别)
1. 思路
- 预处理函数
- 加载数据集与数据预处理
- 自定义层
- 自定义模型网络
- 训练测试(基于keras高层API)
- 保存模型并加载
2. 代码
import tensorflow as tf
from tensorflow.keras import datasets, layers, optimizers, Sequential, metrics
from tensorflow import keras
import os
os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'
def pre_process(x, y):
x = tf.cast(x, dtype=tf.float32) / 255.
y = tf.cast(y, dtype=tf.int32)
return x, y
batch_size = 128
(x, y), (x_test, y_test) = datasets.cifar10.load_data()
y = tf.squeeze(y)
y = tf.one_hot(y, depth=10)
y_test = tf.squeeze(y_test)
y_test = tf.one_hot(y_test, depth=10)
train_db = tf.data.Dataset.from_tensor_slices((x, y))
train_db = train_db.map(pre_process).shuffle(10000).batch(batch_size)
test_db = tf.data.Dataset.from_tensor_slices((x_test, y_test))
test_db = test_db.map(pre_process).batch(batch_size)
sample = next(iter(train_db))
class MyDense(layers.Layer):
def __init__(self, input_dim, output_dim):
super(MyDense, self).__init__()
self.kernel = self.add_variable('w', [input_dim, output_dim])
def call(self, inputs, training=None):
x = inputs @ self.kernel
return x
class MyModel(keras.Model):
def __init__(self):
super(MyModel, self).__init__()
self.fc1 = MyDense(32*32*3, 256)
self.fc2 = MyDense(256, 128)
self.fc3 = MyDense(128, 64)
self.fc4 = MyDense(64, 32)
self.fc5 = MyDense(32, 10)
def call(self, inputs, training=None):
x = tf.reshape(inputs, [-1, 32*32*3])
x = self.fc1(x)
x = tf.nn.relu(x)
x = self.fc2(x)
x = tf.nn.relu(x)
x = self.fc3(x)
x = tf.nn.relu(x)
x = self.fc4(x)
x = tf.nn.relu(x)
x = self.fc5(x)
return x
Model = MyModel()
Model.compile(
optimizer=optimizers.Adam(lr=1e-3),
loss=tf.losses.CategoricalCrossentropy(from_logits=True),
metrics=['accuracy']
)
Model.fit(train_db, epochs=20, validation_data=test_db, validation_freq=1)
Model.evaluate(test_db)
Model.save_weights('Model_weights.ckpt')
del Model
print('saved')
Model = MyModel()
Model.compile(
optimizer=optimizers.Adam(lr=1e-3),
loss=tf.losses.CategoricalCrossentropy(from_logits=True),
metrics=['accuracy']
)
Model.load_weights('Model_weights.ckpt')
print('loaded')
Model.evaluate(test_db)