Keras进阶实战:函数式API、自定义层与模型训练回调全解析
2026/9/14 22:13:49 网站建设 项目流程

学Keras有一个很微妙的阶段:跟着教程把Sequential API用的滚瓜烂熟,手写几个全连接网络、CNN、RNN都不在话下,但真到了自己的实际任务,很可能会当场卡住——输入不止一路怎么办?不同分支要共享参数怎么办?损失函数里要加一个惩罚项怎么办?训练到一半想把最好的模型捞出来该怎么做?如果你正处于这个位置,那这篇正是为你准备的。所谓“进阶”,不是去背更多API,而是从“会用工具搭积木”变成“能用工具解决复杂结构问题”。本文会围绕Keras/深度学习框架体系里最核心的进阶能力展开:函数式API、自定义组件、回调机制、模型序列化保存,以及一个多输入的完整实战。适合已经会基础深度学习、想独立完成真实项目的读者。

1. 进阶之前,先看清Keras的“版本分岔路”

1.1 tf.keras 与 Keras 3,选哪个学

如果你在2024年之后才接触Keras,大概率会遇到一个疑惑:网上教程一会写from tensorflow import keras,一会写import keras,到底哪个对?这里先把这个搞明白,否则后面所有代码都会栽跟头。

目前的现状是:Keras 3已经成为独立的开源库,支持TensorFlow、PyTorch、JAX三种后端,你可以把Keras当成一套统一的高层接口,后端随便切。而TensorFlow内置的tf.keras,本质上是Keras API与TensorFlow强绑定的版本。对大多数深度学习任务来说,两者在写法上几乎一致,差异主要体现在后端切换和多框架协作上。

我的建议很直接:如果你主要用TensorFlow全家桶,那就老老实实用tf.keras,安装TensorFlow时自带,不用额外装;如果你想保持灵活性,之后可能切换JAX或PyTorch做后端,那就直接用Keras 3。进阶教程里的函数式API、自定义层、回调这些核心机制,两套接口完全通用,不影响学习。

1.2 进阶段需要掌握的四个能力

回顾我带过的不少新人,从基础到进阶,差距往往不体现在“知道多少层”,而是体现在四个方面:

  • 结构自由度:能不能用函数式API描述非线性的图结构,比如多输入、多输出、共享层、残差连接。
  • 自定义能力:框架没有现成的损失函数、指标、层结构时,能不能自己去写。
  • 训练过程的控制力:不会只是model.fit(x,y)就跑,而是会在合适时机保存模型、调整学习率、提前停止。
  • 工程化意识:模型训练完怎么存、怎么加载恢复、怎么导出部署,心里要有谱。

这四项正好对应本文后面四个大章节。我不打算罗列API文档,而是按“真实项目里你怎么一步步拆解问题”的顺序来展开。

2. 函数式API:从“顺序堆层”到“构建一张计算图”

2.1 Sequential的天花板在哪

先看一个非常典型的场景:我们要做一个用户购买意愿预测模型,输入有两条线,一是文本评论(比如“东西不错,物流很快”),二是用户历史行为统计特征(比如近30天访问次数、平均停留时长、历史购买率)。这两类特征性质完全不同,应该分别处理之后再融合。

如果用Sequential,你只能被迫把所有特征拼成一个长向量,让全连接层自己去学文本和统计特征之间的交互。这种做法的弊端很明显:文本是需要embedding再进循环网络或卷积网络的结构化序列,统计特征是稠密数值,把二者直接concat再喂同一组全连接层,模型很难学出各自合适的表征。说白了就是任务本来要求“分而治之”,你偏要“大锅炖”。

这就是Sequential的天花板——它只能描述“上一层输出就是下一层输入”这种线性链式结构,表达能力极为有限。真实世界的深度学习模型几乎都不是一条直线走到底的。

2.2 函数式API的核心写法

函数式API非常直接:你可以把每个层当成一个函数,给它输入张量,它返回张量,再手动决定这些张量流到哪里去。所以你可以任意分叉、合并、跳跃。

import tensorflow as tf from tensorflow.keras import layers, Model # 两个独立的输入 text_input = layers.Input(shape=(200,), name="text_input") # 文本序列 feature_input = layers.Input(shape=(64,), name="feature_input") # 数值特征 # 文本分支:Embedding + BiLSTM x_text = layers.Embedding(20000, 128, mask_zero=True)(text_input) x_text = layers.Bidirectional(layers.LSTM(64))(x_text) # 数值分支:两层全连接 x_feat = layers.Dense(32, activation="relu")(feature_input) x_feat = layers.Dense(16, activation="relu")(x_feat) # 融合:拼接之后接输出 combined = layers.concatenate([x_text, x_feat]) combined = layers.Dense(64, activation="relu")(combined) combined = layers.Dropout(0.3)(combined) output = layers.Dense(1, activation="sigmoid", name="output")(combined) model = Model(inputs=[text_input, feature_input], outputs=output) model.summary()

这段代码就是函数式API的一个缩影。注意几个核心点:

  • Input负责定义输入张量,一定要给它一个有意义的name。后续无论是组织训练数据、排查错误,还是做服务化部署,这个name都会成为你定位问题的线索。
  • layers.Xxx(...)(previous_tensor)这种连续调用的方式,左边是层的构造函数,右边是这一层对应张量,这种“洋葱式”写法熟悉之后会很顺手。
  • 合并用layers.concatenate,它也是个层,负责把多个张量在某个维度上拼起来;而Model(inputs=..., outputs=...)最终把这些张量流“冻结”成一个真正的模型对象。

2.3 多输入、多输出模型的完整搭建

2.2的例子只演示了多输入单输出。实际业务里多输出的情况同样常见——比如一个电商场景,你想同时预测“用户是否购买”和“用户可能要买的商品类目”。前者是二分类,后者是多分类。

多输出的写法和多输入完全同构:在模型中间分叉出两个不同的输出头即可。

# 继续沿用上面的文本+特征融合结构 shared = layers.Dense(64, activation="relu")(combined) output_cls = layers.Dense(1, activation="sigmoid", name="click")(shared) output_cate = layers.Dense(10, activation="softmax", name="category")(shared) model = Model(inputs=[text_input, feature_input], outputs=[output_cls, output_cate])

编译时注意,多输出需要给每个输出指定自己的损失函数,甚至每个输出可以有不同的loss权重:

model.compile( optimizer="adam", loss={"click": "binary_crossentropy", "category": "categorical_crossentropy"}, loss_weights={"click": 1.0, "category": 2.0}, metrics={"click": "accuracy", "category": "accuracy"}, )

loss_weights这个参数特别容易被忽略。它表示不同任务的loss在总loss里占多大比重。如果两个任务的重要程度不一样,或者某个任务的loss数值天然比较大,就会在反向传播时“带走”大部分梯度。我在项目里就遇到过主任务指标一直上不去,后来发现是辅助任务的loss权重太高导致梯度被带偏了,调低权重之后主任务立刻回暖。

2.4 共享层:让不同分支用同一套参数

函数式API还有一个杀手锏——共享层。所谓共享层,就是同一个层的实例被多个输入路径同时使用。最经典的场景是孪生网络:判断两张图片是否相似,两个分支完全共享同一套卷积参数,而不是各自学一套。

from tensorflow.keras import layers, Model input_a = layers.Input(shape=(28, 28, 1), name="image_a") input_b = layers.Input(shape=(28, 28, 1), name="image_b") feature_extractor = tf.keras.Sequential([ layers.Conv2D(32, (3, 3), activation="relu"), layers.MaxPooling2D(), layers.Conv2D(64, (3, 3), activation="relu"), layers.GlobalAvgPool2D(), ]) out_a = feature_extractor(input_a) out_b = feature_extractor(input_b) merged = layers.concatenate([out_a, out_b]) output = layers.Dense(1, activation="sigmoid")(merged) similarity_model = Model(inputs=[input_a, input_b], outputs=output)

这里的feature_extractor被调用了两次,但参数是同一份。这一点非常关键——它让模型可以处理“同一对象的不同视图”这类任务,还能大大减少参数量。共享层的思想不止用于孪生网络,多任务学习中让不同任务共享底层特征表示,也是同样的套路。

顺带说一句,有些朋友会把Model再当成子模块嵌到另一个Model里用,这也是完全合法的。函数式API支持模型嵌套,一个Model实例可以作为另一个更大模型的“层”来调用。

3. 自定义组件:把Keras从工具箱变成你的专属车间

3.1 为什么要自己写损失函数

Keras内置的损失函数确实覆盖了大部分常规需求,但实际项目中总会出现“内置函数无法直接表达”的局面。举一个很经典的例子:类别不平衡的二分类问题,负样本是正样本的100倍,直接交叉熵会让模型把所有样本都预测为负样本,因为这样loss已经很低了。Focal Loss就是为了解决这个问题而出现的。

Focal Loss的核心思想是:让模型把注意力集中在难分类的样本上,对置信度高、容易分类的样本降低权重。公式长这样:

FL(p_t) = -alpha_t * (1 - p_t)^gamma * log(p_t)

其中p_t是模型对真实类别的预测概率,alpha调节正负样本权重,gamma调节困难样本的聚焦程度。当gamma=0时就退化成普通的带权重交叉熵。

用Keras实现Focal Loss非常直接:

def focal_loss(alpha=0.25, gamma=2.0): def loss(y_true, y_pred): epsilon = tf.keras.backend.epsilon() # 截断预测值,避免log(0) y_pred = tf.clip_by_value(y_pred, epsilon, 1.0 - epsilon) # 交叉熵的两种取值 ce = -y_true * tf.math.log(y_pred) - (1 - y_true) * tf.math.log(1 - y_pred) # 调制系数 p_t = y_true * y_pred + (1 - y_true) * (1 - y_pred) alpha_t = y_true * alpha + (1 - y_true) * (1 - alpha) loss_value = alpha_t * tf.pow(1.0 - p_t, gamma) * ce return tf.reduce_mean(loss_value) return loss model.compile(optimizer="adam", loss=focal_loss(gamma=2.0), metrics=["accuracy"])

这里用了闭包(函数内部套函数)的写法,目的是让alphagamma作为外部参数传入后,返回真正计算loss的函数。Keras在训练时会自动调用这个函数,传入y_truey_pred两个张量,返回一个标量损失即可。

注意一个细节:y_pred必须截断,原因很简单,如果某个样本的预测概率接近0或1,log会趋向于负无穷,数值计算就会出NaN。这类细节处理充分说明自定义损失函数不只是“把公式翻译成代码”,还需要考虑数值稳定性。

3.2 自定义评估指标

损失函数是用来优化的,而评估指标是给人看的,两者经常不是同一个东西。比如你训练一个目标检测模型,损失可能是Smooth L1加交叉熵,但上线后你更关心mAP或者IOU超过某个阈值的准确率。

自定义指标在Keras里推荐的方式是继承tf.keras.metrics.Metric,实现三个方法:update_state在每次batch后更新内部状态,result返回当前指标值,reset_state在epoch开始时清零状态。

举一个实际例子——我想监控“预测概率落在0.5到0.8之间的样本比例”,纯粹为了观察模型输出的分布:

class MidConfidenceMetric(tf.keras.metrics.Metric): def __init__(self, name="mid_confidence", **kwargs): super().__init__(name=name, **kwargs) self.mid_count = self.add_weight(name="mid_count", initializer="zeros") self.total = self.add_weight(name="total", initializer="zeros") def update_state(self, y_true, y_pred, sample_weight=None): pred_classes = tf.cast(tf.argmax(y_pred, axis=-1), tf.float32) # 假设二分类,取正类的概率 pos_prob = y_pred[:, 1] is_mid = tf.cast((pos_prob >= 0.5) & (pos_prob <= 0.8), tf.float32) self.mid_count.assign_add(tf.reduce_sum(is_mid)) self.total.assign_add(tf.cast(tf.size(pos_prob), tf.float32)) def result(self): return self.mid_count / self.total def reset_state(self): self.mid_count.assign(0.0) self.total.assign(0.0)

自定义指标一个容易踩的坑是:如果在update_state里忘了调用assign系列方法去更新状态变量,指标值根本不会变。我最初写自定义指标时总喜欢把中间结果用一个Python浮点到累计,后来才意识到Keras的Metric体系要求所有状态必须是tf.Variable,而且必须是add_weight注册的变量。跨batch累积时这个区别尤其明显,Python浮点数会在tf.function图执行时被冻结,导致指标越跑越奇怪。

3.3 自定义层:让网络拥有“外科手术级”的操控能力

自从Keras 2开始,自定义层的门槛已经大大降低了。你只需要继承tf.keras.layers.Layer,在__init__里定义子层和初始状态,在build方法里根据输入形状声明可训练参数,在call里写前向计算逻辑。

下面是一个带可学习衰减系数的自定义Dense层,它会在标准全连接输出的基础上乘一个可训练标量:

class ScaledDense(layers.Layer): def __init__(self, units, activation=None, **kwargs): super().__init__(**kwargs) self.units = units self.activation = tf.keras.activations.get(activation) def build(self, input_shape): # 输入特征维度 in_dim = input_shape[-1] # 可训练权重 self.w = self.add_weight( shape=(in_dim, self.units), initializer="glorot_uniform", trainable=True, name="kernel", ) self.bias = self.add_weight( shape=(self.units,), initializer="zeros", trainable=True, name="bias" ) self.scale = self.add_weight( shape=(1,), initializer=initializers.Constant(1.0), trainable=True, name="scale" ) def call(self, inputs): output = tf.matmul(inputs, self.w) + self.bias output = output * self.scale return self.activation(output)

关于build,有一个新手常犯的错误:不管三七二十一,把所有参数都塞在__init__里定义。这就导致你如果不看输入维度就没法确定参数形状,比如权重矩阵的行数必须等于输入特征数。把参数定义放到build里面,Keras会在第一次执行时自动传入输入形状并调用build,这种“延迟创建参数”的机制是Keras层的标准范式。

还有一点,如果你的层在训练和推理阶段行为不同——比如包含Dropout或BatchNormalization——一定要在call方法里接受training参数并传给子层:

def call(self, inputs, training=None): x = self.dense(inputs) x = self.dropout(x, training=training) return x

忘记传training是自定义层最隐蔽的错误之一。因为Dropout在训练和推理时行为完全不同,如果你在call内部调用了一个含Dropout的子层却不传training,它默认按推理模式处理,相当于Dropout根本没生效,但训练日志完全看不出来,只会模型收敛变慢、过拟合变严重。

3.4 关于自定义组件的最佳实践建议

自定义组件虽然自由,但也要克制。我给的建议是:能用内置层组合解决的就不要自己写层;自己写层时尽量保持call方法里的运算简单清晰;写完后先用一行随机数据做一次前向传播测试,再进入训练流程。

# 一秒钟的冒烟测试 dummy_input = tf.random.normal((4, 16)) scaled_dense = ScaledDense(8, activation="relu") out = scaled_dense(dummy_input) assert out.shape == (4, 8)

代码能跑通,再安心去训练。这个习惯能帮你省下大量排查时间。

4. 回调函数:让训练过程自动且可控

4.1 ModelCheckpoint:别让最好的模型从你手中溜走

训练深度学习模型最大的痛点之一是:模型在训练后期会震荡,最后一个epoch的权重往往不是val_loss最低的那个点,而且如果你没有及时保存,崩了就只能重来。ModelCheckpoint回调就是专门解决这个问题的。

from tensorflow.keras.callbacks import ModelCheckpoint checkpoint = ModelCheckpoint( filepath="models/epoch_{epoch:02d}_val_loss_{val_loss:.4f}.keras", monitor="val_loss", save_best_only=True, mode="min", save_weights_only=False, verbose=1, )

这里filepath支持{}格式的占位符,会自动用实际数值填充。monitor决定监控哪个指标,mode告诉回调该指标是越小越好还是越大越好。save_best_only=True配合monitor="val_loss"的含义是:只有当当前epoch的val_loss比历史上所有epoch都好时,才会覆盖保存。

这个回调最重要的一点是save_weights_only这个参数。如果设为True,只保存权重,文件小,恢复时必须重新构建模型结构代码;如果设为False,保存完整模型,包括网络结构、优化器状态、损失函数配置,加载后可以直接继续训练。常规训练过程中我推荐save_weights_only=False,虽然文件大一点,但“开箱即用”的感觉太好了,根本不需要再去拼结构代码。

4.2 EarlyStopping 与 ReduceLROnPlateau:训练的黄金搭档

EarlyStopping用来防止过拟合:如果连续多个epoch验证集指标不再变好,就提前终止训练。

from tensorflow.keras.callbacks import EarlyStopping, ReduceLROnPlateau early_stop = EarlyStopping( monitor="val_loss", patience=5, restore_best_weights=True, ) reduce_lr = ReduceLROnPlateau( monitor="val_loss", factor=0.5, patience=3, min_lr=1e-6, )

一个非常实用也容易被忽略的参数是restore_best_weights=True。如果训练在epoch 20被提前终止,而最好的val_loss出现在epoch 15,那么当restore_best_weights=True时,训练结束后模型会自动回滚到epoch 15的权重。如果是在项目里做模型对比,这一点特别重要——否则你拿到的模型是epoch 20的权重,并不是验证集表现最好的那个。

ReduceLROnPlateau则是等验证指标停滞时把学习率降一半,让模型在更小的步长下继续精细搜索。我习惯把patience设置成EarlyStopping的一半,让学习率先面临缩减的“警告”,如果连续两三次降学习率还止不住颓势,再触发早停。

4.3 TensorBoard:训练过程的可视化仪表盘

TensorBoard的作用不只是画loss曲线。它会自动记录训练过程中的很多信息,包括计算图、梯度分布、权重分布、样本图像、文本嵌入等。

from tensorflow.keras.callbacks import TensorBoard tensorboard = TensorBoard( log_dir="logs/run_001", histogram_freq=1, write_graph=True, write_images=False, )

在日志目录中积累几次训练之后,运行tensorboard --logdir logs,就能在浏览器里同时对比不同实验的曲线。我在实际调参时最常看的是Scalars面板里的epoch_lossepoch_accuracy,如果发现在某一步loss突然冒出尖峰,再切到Distributions面板看权重分布是否出现异常。很多情况下,loss曲线上的异常尖峰都对应着某些层的梯度爆炸,提前发现就避免了训练白费。

这里有个实用技巧:给每次训练单独建一个目录,比如logs/run_001logs/run_002,目录名里带上你这次实验想验证的变量,比如logs/lr_1e-3_bs_64。这样TensorBoard能把这些实验放在同一个对比视图里,一目了然。如果什么都不管全往一个目录写,TensorBoard会把它们视为同一份实验,曲线画在一起,反而乱。

4.4 自定义回调,把业务逻辑插进训练流程

内置回调解决不了所有问题,Keras的解决方案是允许你写自定义回调。继承tf.keras.callbacks.Callback,重写几个关键方法:on_epoch_endon_batch_endon_train_begin等。

一个非常实用的场景:我需要在每个epoch结束后向训练集上重新做一次预测,计算模型在某个业务指标上的表现,这个指标不在loss里。于是写一个自定义回调:

class EvaluateBusinessMetric(tf.keras.callbacks.Callback): def __init__(self, validation_data, threshold=0.5): super().__init__() self.validation_data = validation_data self.threshold = threshold def on_epoch_end(self, epoch, logs=None): x_val, y_val = self.validation_data y_pred = self.model.predict(x_val, verbose=0) y_pred_label = (y_pred > self.threshold).astype("int32") # 计算业务自定义的覆盖率指标 coverage = (y_pred_label.sum() / len(y_pred_label)) * 100 print(f"epoch {epoch + 1} - coverage under threshold: {coverage:.2f}%")

自定义回调的要求是:不要修改logs字典的key(Keras会用这些名称记录日志),也不要自己调用model.fit,但model.predict完全可以使用。这样做的好处是,你在每个epoch结束的时候能拿到肉眼可见的业务指标反馈,而不是只盯着一堆loss数值。

5. 模型保存、加载与恢复:工程化基本功

5.1 三种保存方式怎么选

Keras里保存模型有几种常见方式,选择的原则其实很简单,取决于你想实现什么目的。

表格对比一下:

方式保存内容文件格式典型场景
model.save_weights()只有权重.h5/.weights.h5临时存档、迁移学习
model.save()完整模型结构+权重+优化器状态.keras/.h5常规训练完存档、恢复训练
TFSavedModel完整模型,且带推理接口SavedModel目录上线到TensorFlow Serving

model.save_weights是三者里最轻量的,但它的前提是你已经有模型结构代码。比如做Fine-tuning时,你想把预训练模型换到新任务上,只需要A模型的结构+B模型的权重,这种情况下save_weights再合适不过。

5.2 完整模型的保存和加载

正确做法是保存完整模型:

# 训练结束后 model.save("my_model.keras") # 部署或加载时 from tensorflow.keras.models import load_model restored_model = load_model("my_model.keras")

.keras格式是Keras 3引入的新格式,往早期的.h5格式更安全、可扩展性更好。如果你用的还是tf.keras,文件后缀写成.h5也能用,但如果你装的是Keras 3,我更推荐.keras

加载完整模型最大的诱惑是:你连自定义层和自定义损失函数都不需要重新写代码了,Keras会序列化它们的类名。但是这里有个前提——当你加载时有自定义组件,必须在load_model调用时显式传入custom_objects参数:

restored_model = load_model( "model_with_custom_layer.keras", custom_objects={"ScaledDense": ScaledDense}, )

5.3 恢复训练的关键:优化器状态

想“断点续训”,只保存权重是不够的——因为Adam这类优化器会维护一阶矩和二阶矩的估计值(也就是它的动量信息)。如果只加载权重,优化器的内部状态会从零重新开始,相当于学习率调度中断了,模型收敛过程会被打乱。保存完整模型时,Keras会把优化器状态一并保存,所以恢复后的训练曲线是连续衔接的。

# 保存完整模型(包含优化器状态) model.save("checkpoint_epoch_20.keras") # 恢复训练 restored = load_model("checkpoint_epoch_20.keras") restored.compile(optimizer="adam", loss="binary_crossentropy", metrics=["accuracy"]) restored.fit(train_dataset, epochs=40, initial_epoch=20)

注意initial_epoch=20这个参数,它告诉fit这是从第20个epoch继续跑,而不是从头开始。如果你写epochs=40且不写initial_epoch,模型会重新从epoch 0跑到40,之前的训练记录不会自动衔接。这个问题在真机上很常见——很多新手保存了断点,却不会用initial_epoch恢复,结果白跑了。

5.4 部署场景下的导出选择

如果你要上线模型做推理,我建议导出为SavedModel格式:

model.export("saved_model_dir")

SavedModel目录里面包含了完整的推理图,用户不需要知道任何Keras或TensorFlow的API细节,直接用tf.saved_model.load就能跑。如果你的上线环境是TensorFlow Serving,这个格式更是原生支持的。别用model.save()导出的格式直接上线,我说的不是不能用,而是SavedModel在推理优化的支持上更全面,QuanTization、Optimization、Serving都认它。

6. 完整实战:文本与数值特征融合的推荐模型

6.1 任务设定与数据准备

这部分我们把前面的知识点串起来做一个相对完整的项目:影评购买预测。输入包括一段影评文本和一组用户行为统计特征,目标是预测用户是否会实际购买该电影。这类多输入场景在真实推荐系统、广告点击率预估里非常常见。

我直接用IMDB影评数据集作为文本来源,同时随机生成一个形状为(N, 64)的数值特征来模拟用户行为统计量。两者在同一个样本中一一对应。

import numpy as np import tensorflow as tf from tensorflow.keras import layers, Model # IMDB数据 (vocab_size, max_len) = (20000, 200) (x_text_train, y_train), (x_text_test, y_test) = tf.keras.datasets.imdb.load_data( num_words=vocab_size ) x_text_train = tf.keras.preprocessing.sequence.pad_sequences(x_text_train, maxlen=max_len) x_text_test = tf.keras.preprocessing.sequence.pad_sequences(x_text_test, maxlen=max_len) # 随机生成的数值特征,模拟“用户行为统计” num_features = 64 np.random.seed(0) x_num_train = np.random.randn(len(x_text_train), num_features).astype("float32") x_num_test = np.random.randn(len(x_text_test), num_features).astype("float32") # 只看前20000条,模拟中等规模数据 x_text_train = x_text_train[:20000] x_num_train = x_num_train[:20000] y_train = y_train[:20000] x_text_test = x_text_test[:5000] x_num_test = x_num_test[:5000] y_test = y_test[:5000]

注意,这里数值特征完全是随机生成的,模型不可能从中学到真实规律,我们主要看代码流程和结构是否正确——这是很多项目起步阶段的常用验证方法。

6.2 用函数式API搭建融合模型

模型结构设计为“文本分支和数值分支分别提取表征,最后融合二分类”。文本分支使用Embedding加BiLSTM,数值分支使用两层全连接,融合后接一个输出层。

# 文本输入 text_input = layers.Input(shape=(max_len,), name="text") x_text = layers.Embedding(vocab_size, 128, mask_zero=True)(text_input) x_text = layers.Bidirectional(layers.LSTM(64, dropout=0.2))(x_text) # 数值输入 num_input = layers.Input(shape=(num_features,), name="numeric") x_num = layers.Dense(32, activation="relu")(num_input) x_num = layers.Dense(16, activation="relu")(x_num) # 融合层 combined = layers.concatenate([x_text, x_num]) combined = layers.Dense(64, activation="relu")(combined) combined = layers.Dropout(0.3)(combined) output = layers.Dense(1, activation="sigmoid", name="purchase")(combined) model = Model(inputs=[text_input, num_input], outputs=output) model.compile( optimizer=tf.keras.optimizers.Adam(learning_rate=1e-3), loss="binary_crossentropy", metrics=["accuracy"], ) model.summary()

这里我使用了mask_zero=True。它的作用是告诉Embedding层输入中ID为0的位置是填充位,在后面的LSTM中这些位置也会被跳过。如果文本长度参差不齐,pad之后填入的0会干扰LSTM的语义计算,mask_zero能有效屏蔽它们。

6.3 训练配置与回调配合

真实训练中我会把第4节讲到的回调全部接上:

checkpoint = tf.keras.callbacks.ModelCheckpoint( "best_model.keras", monitor="val_accuracy", mode="max", save_best_only=True, ) early_stop = tf.keras.callbacks.EarlyStopping( monitor="val_loss", patience=3, restore_best_weights=True, ) reduce_lr = tf.keras.callbacks.ReduceLROnPlateau( monitor="val_loss", factor=0.5, patience=2, min_lr=1e-6 ) history = model.fit( [x_text_train, x_num_train], y_train, validation_data=([x_text_test, x_num_test], y_test), epochs=30, batch_size=64, callbacks=[checkpoint, early_stop, reduce_lr], )

注意fit的输入数据是以列表形式传入[x_text_train, x_num_train],列表的顺序要和模型输入定义的顺序一致。如果你的模型输入指定了name,还可以用字典传入:

history = model.fit( {"text": x_text_train, "numeric": x_num_train}, y_train, validation_data=({"text": x_text_test, "numeric": x_num_test}, y_test), ... )

字典传参在输入分支很多的时候可读性要高得多,而且能避免“输入顺序写错”这种低级但致命的错误。我在项目里一旦输入超过两个,一律用字典,绝不参数列表——这是踩过一次“顺序颠倒,训练还能跑但效果全无”的大坑之后形成的习惯。

6.4 结果评估与小节

训练结束后加载最佳模型,验证集评估:

best_model = tf.keras.models.load_model("best_model.keras") loss, acc = best_model.evaluate([x_text_test, x_num_test], y_test) print(f"Test accuracy: {acc:.4f}")

这个实战虽然数据是模拟的,但结构和流程完全可以套用到真实项目。从函数式API定义多输入、用内置和自定义层组合提取特征、再到回调机制配合训练、最后保存加载完整模型——这一套操作就是日常深度学习项目的完整闭环。

7. 训练效率与稳定性优化的几条实践经验

7.1 学习率不要一条路走到黑

固定学习率跑完全程,基本不是最优解。我常用的策略是配合ReduceLROnPlateau动态降学习率,更精细一点的做法是使用余弦退火调度器:

lr_schedule = tf.keras.optimizers.schedules.CosineDecay( initial_learning_rate=1e-3, decay_steps=5000, alpha=1e-5, ) optimizer = tf.keras.optimizers.Adam(learning_rate=lr_schedule)

余弦退火的思路是让学习率按照余弦曲线从大到小平滑下降,不像阶梯式那样突变。实际效果通常比ReduceLROnPlateau在训练后期更平滑,尤其在训练步数相对固定的任务中。如果你发现模型到了瓶颈期,val_loss上下震荡无法继续收敛,试着把学习率曲线改成余弦退火,经常能再压下去一小截。

7.2 batch size与收敛效果的关系

batch_size太小,梯度噪声大,训练不稳定;batch_size太大,单个epoch时间缩短,但收敛可能变慢,因为每步更新条数少。很多人在固定batch size时忽略了它和学习率的联动。经验法则是:batch size翻倍,学习率也可以尝试翻倍,这样能在保持稳定性的前提下加快收敛。

不过要注意,GPU显存是硬约束。在现网环境里,我通常先把batch size设成模型单batch能塞进显存的最大值,比如序列模型用64或128,再反过来调学习率。拓展数据加载时用tf.data流水线,这一点在第7.3节展开。

7.3 数据加载瓶颈:别让GPU等你

训练速度慢,不一定是模型的问题,更多时候是数据加载跟不上。如果你发现GPU利用率经常在50%以下,多半是数据IO卡住了。解决方案是使用tf.data.Dataset

dataset = tf.data.Dataset.from_tensor_slices(({"text": x_text_train, "numeric": x_num_train}, y_train)) dataset = dataset.shuffle(10000).batch(64).prefetch(tf.data.AUTOTUNE)

prefetch可以在GPU计算当前batch时,CPU提前准备下一个batch,大幅减少了模型等待数据的时间。如果你的任务涉及图片解码、文本tokenize等预处理,把这些操作也放进Dataset的map里,配合num_parallel_calls=tf.data.AUTOTUNE,并行度会更高:

dataset = dataset.map(preprocess_function, num_parallel_calls=tf.data.AUTOTUNE)

7.4 混合精度:白捡的加速

如果你的GPU是支持NVIDIA AMP的版本,比如V100、T4、A100等,开启混合精度训练是一个非常简单的加速手段:

tf.keras.mixed_precision.set_global_policy("mixed_float16")

设置之后,Keras会自动把适合用低精度计算的算子(比如卷积、全连接)转换为float16,而维持一些对精度敏感的计算(如loss)在float32。大部分模型的训练时间能缩短30%到50%。唯一要注意的是,开启混合精度后loss可能波动更大一点,建议配合学习率调度一起使用。

7.5 写过无数次之后,最终沉淀下来的几条心得

写到这里,把我在Keras实战中最想说的几条心得一起放出来:

  • 别迷信验证集loss,多关注业务指标。loss下降不代表模型可用,用回调把业务指标打印出来,它们比loss更接近真实目标。
  • 每次实验保留完整闭环。数据版本、模型结构、训练超参、回调配置,这些信息远比一个检测精度重要。保存模型的同时顺手把实验配置写进JSON放进同一目录,复盘时才不会靠回忆。
  • 优先跑通小规模。不管项目多复杂,先拿几千条数据、几个epoch把整条流水线跑通,再上全量数据。我在真实项目中就是这么做的,省下的时间远远超过“人生得意须尽欢”式直接全量训练所花费的成本。

进阶的过程其实就是一个又一个“为什么”被解答的过程:为什么函数式API能表达复杂结构?为什么自定义层要放在build里?为什么保存时连优化器状态一起存?把这些问题一个个想透了,你手里的Keras才真正变成了你自己的工具。接下来再遇到自己的业务模型,就只剩下把数据喂进去这一件事了。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询