咱们在做深度学习开发的时候,总免不了和框架打交道,以前用TensorFlow 1.x的人都懂,搭建模型要处理很多底层的会话、占位符这些麻烦事,还得单独导入Keras的API,两套API混着记特别容易搞混。但到了TensorFlow 2.x的时候,官方直接把Keras设成了默认的高层API,相当于把好用的Keras和TensorFlow的底层能力打通了,模型构建的流程直接简化了不止一点,今天就给大家好好拆解一下这个简化的过程,还有实际用的时候要注意的点。
一、TF2.x默认Keras API的核心简化逻辑
1.1 不用再记两套API,统一接口
以前用TF1.x的话,你得同时记得TensorFlow本身的API和Keras的API,比如要定义层的时候,可能得用tf.layers,还得用keras.layers,或者模型编译的时候,要分开处理,很容易出错。但到了TF2.x,所有Keras的核心功能都整合到了tf.keras下面,你只要导入TensorFlow就行,比如写tf.keras.layers,不用再单独import keras,省了好多记忆成本,而且API的命名也更统一,比如加载数据集、定义层、训练模型的方法,都有一致性,刚接触的人也能快速上手。
1.2 自动封装底层细节,不用管复杂配置
以前用TF1.x的时候,要训练模型得手动定义占位符、会话,还要写梯度下降的代码,每一步都要自己处理,稍微写错一点就跑不起来。但TF2.x把这些底层的求导、计算图的构建都封装到Keras的高层API里了,你只要写好模型的结构,调用compile和fit这两个方法,就能自动完成训练,不用管那些复杂的底层配置,相当于把后台的脏活累活都包了,你只需要专注写模型的逻辑就行。
二、用默认Keras API构建模型的完整示例(技术栈:TensorFlow 2.x 默认 Keras API)
下面这个示例是用最常见的MNIST手写数字识别,从数据加载到模型训练评估,全程用TF2的默认Keras API,每一步都有注释,大家可以直接复制运行:
# 导入TensorFlow,自动加载默认的Keras API
import tensorflow as tf
# 1. 加载内置的MNIST手写数字数据集,这个数据集是TF官方准备好的,方便测试模型
(x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data()
# 2. 数据预处理:把像素值从0-255归一化到0-1之间,这样模型训练的时候收敛更快
x_train = x_train / 255.0
x_test = x_test / 255.0
# 3. 把28x28的二维图像打平成784维的一维向量,适配后面的全连接层输入
x_train = x_train.reshape(-1, 28 * 28)
x_test = x_test.reshape(-1, 28 * 28)
# 4. 构建序贯模型(Sequential),这是Keras里最简单的模型结构,就是一层接一层的线性堆叠
model = tf.keras.Sequential([
# 第一层:全连接层,有256个神经元,激活函数用ReLU(避免梯度消失的常用函数)
tf.keras.layers.Dense(256, activation='relu', input_shape=(28 * 28,)),
# Dropout层:防止过拟合,训练的时候随机丢弃20%的神经元,让模型不会太依赖某几个特征
tf.keras.layers.Dropout(0.2),
# 输出层:10个神经元对应0-9这10个数字,激活函数用Softmax,把输出转成0-1之间的概率
tf.keras.layers.Dense(10, activation='softmax')
])
# 5. 编译模型:指定训练需要的优化器、损失函数和评价指标
model.compile(
optimizer='adam', # Adam是常用的优化器,自适应学习率,不用手动调参数,新手友好
loss='sparse_categorical_crossentropy', # 稀疏分类交叉熵,适合标签是整数的分类任务
metrics=['accuracy'] # 用准确率作为评价指标,方便看模型好坏
)
# 6. 训练模型:用训练集,跑5轮(epochs),每批处理32个样本,留10%的训练数据做验证
history = model.fit(
x_train, y_train,
epochs=5,
batch_size=32,
validation_split=0.1
)
# 7. 评估模型在测试集上的表现,看实际的准确率
test_loss, test_acc = model.evaluate(x_test, y_test)
print(f"测试集上的准确率:{test_acc:.4f}")
这个示例跑下来,大概能达到97%以上的准确率,全程不到20行核心代码,不用处理任何底层的配置,这就是TF2默认Keras API简化模型构建的直观体现。
三、应用场景分析
用TF2默认Keras API搭建模型,适合绝大多数开发者的场景,尤其是这几种情况: 第一种是深度学习入门和课程设计,比如大学生做毕业设计的原型、课程作业,这种时候要求快速出结果,不用纠结底层细节,用这个流程,几个小时就能搭好一个图像分类或者文本分类的模型,把精力放在模型的效果优化上,而不是框架的配置上。 第二种是快速验证想法,比如你突然想到一个新的模型结构,想测一下效果,用Keras的Sequential或者Functional API,能快速搭建起来,不用花时间处理底层的计算图和梯度,几分钟就能跑通测试。 第三种是小项目的快速开发,比如公司内部的小工具,或者个人的AI小项目,比如图片分类的小应用,用这个流程,几天就能上线,不用管太复杂的分布式或者底层优化,高层API足够用。
四、技术优缺点梳理
任何技术都有两面性,TF2默认Keras API也不例外,咱们说一下优缺点: 优点主要有四个:第一,API统一,不用记两套,减少出错的概率,不管是导入、定义层还是训练,都是一套逻辑,刚接触的人也能快速适应;第二,简化流程,封装了底层的求导、计算图构建,不用处理session、占位符这些麻烦事,写更少的代码就能搭建好模型;第三,文档和资源多,Keras是目前最流行的高层深度学习API,网上的教程、案例特别多,遇到问题很容易找到解决办法;第四,灵活性足够,除了最简单的Sequential模型,还有Functional API可以搭建复杂的模型(比如多输入多输出的模型),不用放弃太多灵活性。 缺点也有两个:第一,对于极端定制化的需求,比如要自己改写梯度计算的底层逻辑,或者做特别大规模的分布式训练,Keras的高层封装可能不够灵活,这时候需要用到TensorFlow的原生底层API;第二,性能上,有时候底层定制的模型会比Keras封装的模型快一点,对于对性能要求极高的场景,可能需要结合原生API来优化。
五、注意事项
用TF2默认Keras API搭建模型的时候,有几个点要注意,不然容易踩坑: 第一,版本问题,必须用TensorFlow 2.x及以上的版本,TF1.x的时候Keras是单独的模块,不是默认的,所以如果用旧版本的话,很多API都用不了,升级到TF2.x才能享受到简化的好处; 第二,自定义层的规范,如果要写自己的层,必须按照Keras的规范来,比如要继承tf.keras.layers.Layer类,实现build和call方法,不然模型会报错,比如你写了一个自定义的层,却没实现build方法,训练的时候会出问题; 第三,数据维度的问题,TF2默认的是channels_last的维度顺序,也就是对于图像数据,形状是(样本数,高度,宽度,通道数),比如MNIST的图像是(60000,28,28),如果是RGB图像的话,就是(样本数,28,28,3),如果搞错了维度,模型训练的时候会出错,比如把通道数放在前面,就会得到错误的形状; 第四,训练参数的设置,比如epochs、batch_size这些,要根据数据集的大小来调整,比如数据集小的话,batch_size可以设小一点,epochs不用太多,不然容易过拟合;数据集大的话,batch_size可以设大一点,加快训练速度; 第五,标签类型的问题,如果你的标签是整数(比如MNIST的标签是0-9的整数),要用sparse_categorical_crossentropy损失函数,如果标签是one-hot编码的形式,要用categorical_crossentropy,不然会出错,比如用错损失函数的话,训练的时候loss会一直很高,准确率上不去。
六、文章总结
总的来说,TensorFlow 2.x把Keras设为默认API,是深度学习框架的一次重要优化,它把以前复杂的底层配置都封装了起来,统一了API,让开发者能快速搭建模型,不用花太多精力在框架的细节上,而是专注于模型的逻辑和效果。不管你是刚入门的新手,还是做小项目的开发者,用这个简化流程都能大大提高开发效率,快速看到模型的效果。当然,它也有自己的适用范围,对于极端定制和高性能的场景,可能需要用到原生API,但对于绝大多数常见的开发场景,默认Keras API已经足够好用,能帮你少走很多弯路。
评论
围绕“以Keras为默认API,TensorFlow 2.x如何简化模型构建流程”参与讨论