4.1 Keras三种构建方式


4.1 Keras 三种构建方式

本节摘要:Keras 是 TensorFlow 的官方高层 API,用层对象拼模型。Sequential 按顺序堆叠,适合入门的一条链:Flatten → Dense(128, relu) → Dense(10, softmax)。函数式 API 用输入张量接到各层输出,可表达多输入拼接。子类化 Modelcall 里写逻辑,更接近 PyTorch 的 forward。构建之后必须 compile 绑定优化器、损失、指标,才能 fit。原文还演示了 summary、多输入模型、以及保存入口,这些构成 Keras 路径的完整最小闭环。

阅读收获

阅读完本节,你应当能够:

  1. 用 Sequential 声明原文入门网络并解释每一层的形状变化
  2. 说明函数式相对 Sequential 多出来的能力:共享层、多输入多输出
  3. 列出 compile 的三个必填意图:优化器、损失、指标
  4. 判断何时该放弃 Sequential,而不是硬用它塞旁路

Keras 想让你少写的是循环,不是层

层还是那些层。Keras 的贡献是:用 Python 列表或函数调用把层连成 Model,再把训练循环收进 fit。对对照课,这意味着 Keras 节重点在「图怎么声明」,循环细节放到 4.3。但声明和编译是连着的:没 compile 的模型可以打印摘要,不能 fit。有人以为 Sequential 构造函数已经「建完」,漏掉 compile,报错信息会提到 compile,不要回头改层数。

Sequential 的心智是队列。第一层要用 input_shape=(28, 28) 或单独 Input 告诉 Flatten 空间尺寸。随后 Dense 只写输出单元数。summary() 会列出每层输出形状与参数量,原文把它当作构建后的第一检查。参数量应接近 2.3 节的心算。若 Flatten 输出不是 784,input_shape 写错了。

函数式的心智是接线。inputs = keras.Input(shape=(784,))x = Dense(...)(inputs),最后 Model(inputs, outputs)。原文多层感知机例子就是这样。多输入例子:两个 Input 分支各自 Dense,再合并,再输出。Sequential 做不了这件事,除非你先自己把特征拼成一个数组——那就不是模型表达多输入,是你在管线里偷懒拼接。真正多源(图像+表格)应用函数式或子类化。

子类化在 Keras 里存在,写法接近 nn.Module:在 __init__ 里建层,在 call 里用。入门不是必须。它的价值是对照:当你觉得「PT 才能写复杂前向」时,Keras 也能,只是默认教材不走这条路。代价是 summary 有时需要先 build,序列化比 Sequential 挑剔。能 Sequential 就 Sequential,能函数式就别上子类——这是减少保存事故的经验,不是能力歧视。

💡 关键直觉:三种构建是表达力阶梯,不是高级阶梯。选刚刚够用的那一档,保存和协作都更省事。

compile 与随后的三板斧

原文 compile:optimizer='adam'loss='sparse_categorical_crossentropy'metrics=['accuracy']。字符串形式够用。若要指定学习率,改成优化器对象。多输入模型 compile 时损失可以是一个,也可以按输出名字给字典——入门单输出不要过度设计。

fit 吃数组或 Dataset。原文 epochs=10,可加 validation_data。返回 history,里面是按 epoch 的损失与指标列表,5.3 节画曲线用它。evaluate 返回标量损失与指标。predict 返回输出张量,入门 Softmax 模型上是概率。np.argmax 得到类别。这三板斧与 4.3 的展开表对应,本节只要求你会调用。

构建方式 擅长 不擅长 序列化
Sequential 单链分类回归 残差、多输入 最省心
函数式 DAG、多输入输出、共享层 动态变化的层数 通常良好
子类化 Python 控制流、自定义训练逻辑 第一次接触 Keras 的人 需额外注意 get_config

常用层原文清单:Dense、Conv2D、MaxPooling2D、Dropout、Flatten、Embedding 等。对照 PT 时按 2.3 表翻译。Embedding 本课不用。卷积在第 5 章出现。Dropout 的 rate 是丢弃比例,0.2 表示保留 80%,不要和「保留比例」参数搞反——PT 的 p 同样是丢弃比例,这点一致。

input_shape 不含批次维。写成 (28, 28) 而不是 (None, 28, 28)(64, 28, 28)。批次维由 fit 或 Dataset 提供。写成 64 会把批次锁死进层定义,换 batch_size 就错。这是从 PyTorch Linear(784, 128) 转过来的人常犯的反向错误:PT 不写批次,Keras 也不写批次,但 Keras 把空间形状写在第一层,PT 把特征维写在 Linear 的 in_features。

保存入口原文在 3.2 末尾就提到,细节在 4.4。构建节只要知道:model.save 存在,且与构建方式有关——子类化需要更多配合。先用 Sequential 完成第 5 章,再炫耀子类化。

形状在 Sequential 里怎么流

以原文入门网络为例。输入 (batch, 28, 28) → Flatten → (batch, 784) → Dense128 → (batch, 128) → Dense10 softmax → (batch, 10)。没有通道维。若输入误成 (batch, 28, 28, 1),Flatten 变成 784 仍可能碰巧对,若再叠 Conv2D 就不是碰巧了。打印 summary 的 Output Shape 列,比在脑子里乘更可靠。

函数式多输入原文:两个分支各有 Input,合并后再 Dense 输出。合并是拼接还是相加,决定后面 Dense 的 in 维。PT 里你在 forwardtorch.cat。Keras 用 concatenate 层。对照到这一步,会发现「动态图好写」的优势在单链 MLP 上几乎为零,在多分支上才开始出现——而函数式已经能覆盖静态多分支。真正需要动态的是「根据输入长度改变层数」这类,入门遇不到。

⚠️ 常见坑:在 Sequential 列表里写了 input_shape 在第二层而不是第一层;或 Flatten 之后仍按图像维去想问题。另一坑:compile 使用 categorical_crossentropy 配整数标签,fit 第一轮数字看起来像在学,其实语义错了。

history.history 的键名带 val_ 前缀区分验证。没有 validation_data 时没有这些键。对照 PT 手写记录时,自己用列表 append,键名不统一会在 5.3 画图时抓狂。建议两侧都用「train_loss / val_loss / train_acc / val_acc」这四个名字记账。

Keras 还有回调:早停、学习率衰减、ModelCheckpoint。构建节不展开,4.4 用检查点回调作为保存的自动化。你现在只要知道 fit(..., callbacks=[...]) 是插入循环的钩子,相当于 PT 循环里你自己写的 if epoch % n。

与 PyTorch 对照的一句话收尾:Keras 把「层列表 + 编译配置」当成模型的对外合同;PT 把「模块树 + 你写的循环」当成合同。读别人代码时先找合同在哪。Keras 仓库先看 Sequential 定义和 compile 行;PT 仓库先看 class 和 train_one_epoch 函数。

compile 之后还能改什么

改学习率要拿优化器对象,不能指望改字符串再 compile 一次却沿用旧动量缓冲。重新 compile 会重建优化器状态,接近「换了新的 Adam」。这在调 lr 时有时是故意的,有时是事故。PT 改 optimizer.param_groups[0]['lr'] 更局部。对照时若一侧重新 compile、一侧只改 param_group,续训曲线会分叉,不要归到框架。

trainable=False 冻结层,在迁移学习时常用。本课从随机初始化开始,不必冻。一旦冻了,trainable_variables 变短,Tape 循环必须用当前列表,不能缓存旧列表。PT 把 requires_grad=False 再把冻结参数从优化器里拿掉(或仍留着但梯度为 None)。漏拿掉时 Adam 仍可能给它们建缓冲,浪费。入门不冻层,少一条事故链。

问题:summary 里的 None 是什么?

批次维未确定时 Keras 显示 None,这是正常的。空间维若也是 None,说明动态形状,入门 MLP 不应出现。Flatten 输出若不是 784,input_shape 写错或输入其实带通道。把 summary 当作单元测试:每次改第一层都重新打印。PT 没有同等一页纸,用一次假前向 model(torch.zeros(2,1,28,28)) 看输出形状 (2,10)。假前向是 PT 的 summary。

函数式模型的 plot_model 能画图,本课用 mermaid 代替,避免依赖额外绘图后端。多输入时在纸上画接线比在 Sequential 里硬拼更不容易错。能画出来再写代码,写完用 summary 对输出维。

compile 使用字符串 'adam' 时,学习率走实现默认值,通常是 1e-3 量级,与原文 PT 的 lr=0.001 接近但不保证每一个小版本都印在你眼前。公平对照请写成优化器对象并显式学习率。metrics=['accuracy'] 对稀疏整数标签会走分类准确率;若你误把标签做成 one-hot 却仍用 accuracy 的默认假设,数字会怪。损失名字与标签编码绑定,指标也一样。改损失时连指标假设一起复查。

fitvalidation_split 能从训练数组末尾切一块,注意它按顺序切,不是分层。数据若按类别排过序,末尾可能全是同一类,验证会骗人。Fashion 官方数组通常已打乱,风险较低,但仍不如自己分层切。Dataset 路径没有 validation_split,必须另建验证 Dataset。这是数组 API 与 Dataset API 的又一差异,写进 3.3 的取舍:能偷懒的切分只存在于数组 fit

回调在 4.3、4.4 展开,这里只需记住 fit(..., callbacks=...) 是列表。空列表与不传等价。自定义回调能在 epoch 末做与 PT 循环里相同的事,所以「Keras 不能在每个 epoch 存 best」是假的。能力在,入口叫回调。

要点串联

  • Sequential 单链:原文 Flatten+Dense128+Dense10 softmax,input_shape 不含批次
  • 函数式接线:多输入、共享、DAG;原文有拼接例子
  • 子类化接近 Module:能写但入门非必须,序列化成本更高
  • compile 三件套:优化器、损失、指标;字符串入门,对象才能改 lr
  • 三板斧:fit / evaluate / predict;history 供曲线
  • 形状用 summary 验:784→128→10,对不上先查 Flatten 与 input_shape

Keras 构建验收:summary 上 Flatten 输出 784,Dense 128,Dense 10,参数量接近心算;compile 三件套已写;input_shape 不含批次维。函数式只在多输入或旁路时启用。子类化能写但不作为本课默认,以免保存变脆。fit 之前用零张量假前向一次,输出形状批次乘 10。history 的键名统一成训练损失、验证损失、训练准确率、验证准确率,方便和 PyTorch 的列表对齐。validation_split 按顺序切,类别若排序过会坑;宁可自己分层切。字符串 adam 要改成显式学习率对象才能公平对照。

假前向不只查输出形状,还查是否已 build。有的层在第一次调用前权重为空,summary 可能要求指定 input_shape。Sequential 第一层写了 input_shape 则构造后即可 summary。函数式在 Input 处钉死形状,summary 通常立刻可用。子类化常要先 call 一次。这三种「何时能打印摘要」的差异,会让你误以为某一构建方式坏了。对照 PT 的假前向,Keras 的 summary 就是声明期体检。compile 之后改层结构应重新 compile,否则优化器与权重列表可能对不上。metrics 名称会影响 history 的键,自定义指标时键名要自己记住。入门只用 accuracy,键稳定。不要在构建节引入学习率衰减回调,那是 5.3 的处方,放到这里会让人以为 Sequential 必须配回调才能训练。

把 Sequential 当默认,函数式当多输入的升级,子类化当对照 PyTorch 的选修。升级的理由是结构需要旁路,不是年资。年资用子类化会让保存变脆,协作变劝退。input_shape 不含批次。假前向输出批次乘 10。参数量接近心算。compile 显式学习率。history 键名四件套。validation_split 小心有序数据。改结构后重新 compile。构建节不塞早停回调,避免「不会回调就不会 fit」的错觉。fit 本身就能完成入门训练。回调是窗,不是门。门是 compile 与数据形状。窗以后再开。

审查 Sequential 清单时从输入形状读到输出类别数,中间不要跳层。跳层会漏掉 Flatten,参数量对不上还以为是框架统计口径不同。口径并不不同,是你漏了展平。函数式接线画在纸上再写代码。子类化 call 里不要新建层。三句话够应付构建节剩下的事故。事故处理完,才把时间交给 4.3 的循环翻译。翻译需要一个已经体检过的模型对象,体检在本节完成。体检对象一旦交给循环,就不要在循环里偷偷加层。偷偷加层会让 summary 与实际前向分家,对照从内部破产。

构建节的最后一条纪律:把 summary 文本贴进实验卡片,而不是只看一眼。贴上的数字在 5.2 还能对得上,才说明你没有在循环里改结构。对不上,先回到本节,不要在 Fashion 上解释曲线。曲线解释不了失踪的 Flatten。失踪的 Flatten 只在声明里。声明里找到它,曲线才有资格被解释。找不到,解释权作废,回到 Flatten 那一行。那一行是 Sequential 的第一块积木,不是装饰。装饰性的 Flatten 不存在。存在的 Flatten 必须吃掉空间维。空间维不吃掉,Dense 会接到错误的输入。

下一节用 nn.Module 表达同一件事:层注册在 __init__,数据流写在 forward,没有 compile。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U