You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何初始化模型参数?补全人脸识别模型训练代码缺失部分

补全《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)

关键说明

  1. 核心缺失补全:model = tflearn.DNN(convnet)是作者遗漏的关键语句,它将定义好的网络结构封装成可训练、可调用的模型实例,没有这行代码,后续的model.fit()就会报错。
  2. 语法错误修正:
    • 原代码中test = [x:]是无效写法,需改为train_data[x:]才能正确截取测试样本
    • 生成numpy数组时,列表推导式必须加[],否则会生成生成器对象导致reshape失败
    • fit方法的epoch参数需改为epochs(复数),符合tflearn的API规范

内容的提问来源于stack exchange,提问作者Shashwata Shastri

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.08 12:55:23