如何初始化模型参数?补全人脸识别模型训练代码缺失部分
补全《Deep Learning Application with Python》中人脸识别模型训练的缺失代码
我在查找图像处理相关资料时接触到Navin Kumar Manaswi所著、2018年出版的《Deep Learning Application with Python (Face detection, Recognition etc.)》一书,该书第12章(第187页)中关于卷积层搭建的段落引起了我的注意,本以为其中包含项目所需的完整代码实现,但发现代码存在缺失:在人脸识别系统训练代码中,作者要求创建并拟合模型,但未给出模型创建或调用语句,请求补全代码缺失部分,或指出是否存在遗漏。
原代码问题分析
这段代码基于tflearn库(TensorFlow的高层封装库)编写,核心缺失是模型实例化语句,同时存在几处语法错误,导致代码无法运行。
修正并补全后的完整代码
# 先导入必要依赖库(原代码可能省略了这部分) import numpy as np import tflearn from tflearn.layers.conv import conv_2d, max_pool_2d from tflearn.layers.core import input_data, fully_connected, dropout from tflearn.layers.estimator import regression # 假设create_training_data()是已实现的自定义函数,返回(图像数组, 标签)格式的列表 train_data = create_training_data() x = -2 train = train_data[:x] test = train_data[x:] # 修正原代码错误:原写法test = [x:]无法获取测试样本 # 修正数组生成语法错误:列表推导式需加[] X = np.array([i[0] for i in train]).reshape(-1, 200, 200, 1) Y = [i[1] for i in train] test_x = np.array([i[0] for i in test]).reshape(-1, 200, 200, 1) test_y = [i[1] for i in test] # 定义卷积网络结构 convnet = input_data(shape=[None, 200, 200, 1], name='input') convnet = conv_2d(convnet, 4, 5, activation='relu') convnet = max_pool_2d(convnet, 5) convnet = conv_2d(convnet, 5, 5, activation='relu') convnet = max_pool_2d(convnet, 5) convnet = conv_2d(convnet, 8, 5, activation='relu') convnet = max_pool_2d(convnet, 5) convnet = fully_connected(convnet, 8, activation='relu') convnet = dropout(convnet, 0.2) convnet = fully_connected(convnet, 2, activation='softmax') convnet = regression(convnet, optimizer='adam', learning_rate=LR, loss='categorical_crossentropy', name='targets') # 补全作者遗漏的模型实例化语句 model = tflearn.DNN(convnet) # 修正epoch参数为epochs(tflearn要求复数形式) model.fit({'input': X}, {'targets': Y}, epochs=1, validation_set=({'input': test_x}, {'targets': test_y}), snapshot_step=500, show_metric=True, run_id=MODEL_NAME)
关键说明
- 核心缺失补全:
model = tflearn.DNN(convnet)是作者遗漏的关键语句,它将定义好的网络结构封装成可训练、可调用的模型实例,没有这行代码,后续的model.fit()就会报错。 - 语法错误修正:
- 原代码中
test = [x:]是无效写法,需改为train_data[x:]才能正确截取测试样本 - 生成numpy数组时,列表推导式必须加
[],否则会生成生成器对象导致reshape失败 fit方法的epoch参数需改为epochs(复数),符合tflearn的API规范
- 原代码中
内容的提问来源于stack exchange,提问作者Shashwata Shastri
相关产品推荐
相关产品推荐

