Keras进阶实战:函数式API、自定义层与模型训练回调全解析
发布时间:2026/9/14 22:13:52来源:尧图网络
学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,), nametext_input) # 文本序列 feature_input layers.Input(shape(64,), namefeature_input) # 数值特征 # 文本分支Embedding BiLSTM x_text layers.Embedding(20000, 128, mask_zeroTrue)(text_input) x_text layers.Bidirectional(layers.LSTM(64))(x_text) # 数值分支两层全连接 x_feat layers.Dense(32, activationrelu)(feature_input) x_feat layers.Dense(16, activationrelu)(x_feat) # 融合拼接之后接输出 combined layers.concatenate([x_text, x_feat]) combined layers.Dense(64, activationrelu)(combined) combined layers.Dropout(0.3)(combined) output layers.Dense(1, activationsigmoid, nameoutput)(combined) model Model(inputs[text_input, feature_input], outputsoutput) model.summary()这段代码就是函数式API的一个缩影。注意几个核心点Input负责定义输入张量一定要给它一个有意义的name。后续无论是组织训练数据、排查错误还是做服务化部署这个name都会成为你定位问题的线索。layers.Xxx(...)(previous_tensor)这种连续调用的方式左边是层的构造函数右边是这一层对应张量这种“洋葱式”写法熟悉之后会很顺手。合并用layers.concatenate它也是个层负责把多个张量在某个维度上拼起来而Model(inputs..., outputs...)最终把这些张量流“冻结”成一个真正的模型对象。2.3 多输入、多输出模型的完整搭建2.2的例子只演示了多输入单输出。实际业务里多输出的情况同样常见——比如一个电商场景你想同时预测“用户是否购买”和“用户可能要买的商品类目”。前者是二分类后者是多分类。多输出的写法和多输入完全同构在模型中间分叉出两个不同的输出头即可。# 继续沿用上面的文本特征融合结构 shared layers.Dense(64, activationrelu)(combined) output_cls layers.Dense(1, activationsigmoid, nameclick)(shared) output_cate layers.Dense(10, activationsoftmax, namecategory)(shared) model Model(inputs[text_input, feature_input], outputs[output_cls, output_cate])编译时注意多输出需要给每个输出指定自己的损失函数甚至每个输出可以有不同的loss权重model.compile( optimizeradam, 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), nameimage_a) input_b layers.Input(shape(28, 28, 1), nameimage_b) feature_extractor tf.keras.Sequential([ layers.Conv2D(32, (3, 3), activationrelu), layers.MaxPooling2D(), layers.Conv2D(64, (3, 3), activationrelu), 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, activationsigmoid)(merged) similarity_model Model(inputs[input_a, input_b], outputsoutput)这里的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调节困难样本的聚焦程度。当gamma0时就退化成普通的带权重交叉熵。用Keras实现Focal Loss非常直接def focal_loss(alpha0.25, gamma2.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(optimizeradam, lossfocal_loss(gamma2.0), metrics[accuracy])这里用了闭包函数内部套函数的写法目的是让alpha、gamma作为外部参数传入后返回真正计算loss的函数。Keras在训练时会自动调用这个函数传入y_true和y_pred两个张量返回一个标量损失即可。注意一个细节y_pred必须截断原因很简单如果某个样本的预测概率接近0或1log会趋向于负无穷数值计算就会出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, namemid_confidence, **kwargs): super().__init__(namename, **kwargs) self.mid_count self.add_weight(namemid_count, initializerzeros) self.total self.add_weight(nametotal, initializerzeros) def update_state(self, y_true, y_pred, sample_weightNone): 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, activationNone, **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), initializerglorot_uniform, trainableTrue, namekernel, ) self.bias self.add_weight( shape(self.units,), initializerzeros, trainableTrue, namebias ) self.scale self.add_weight( shape(1,), initializerinitializers.Constant(1.0), trainableTrue, namescale ) 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, trainingNone): x self.dense(inputs) x self.dropout(x, trainingtraining) return x忘记传training是自定义层最隐蔽的错误之一。因为Dropout在训练和推理时行为完全不同如果你在call内部调用了一个含Dropout的子层却不传training它默认按推理模式处理相当于Dropout根本没生效但训练日志完全看不出来只会模型收敛变慢、过拟合变严重。3.4 关于自定义组件的最佳实践建议自定义组件虽然自由但也要克制。我给的建议是能用内置层组合解决的就不要自己写层自己写层时尽量保持call方法里的运算简单清晰写完后先用一行随机数据做一次前向传播测试再进入训练流程。# 一秒钟的冒烟测试 dummy_input tf.random.normal((4, 16)) scaled_dense ScaledDense(8, activationrelu) 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( filepathmodels/epoch_{epoch:02d}_val_loss_{val_loss:.4f}.keras, monitorval_loss, save_best_onlyTrue, modemin, save_weights_onlyFalse, verbose1, )这里filepath支持{}格式的占位符会自动用实际数值填充。monitor决定监控哪个指标mode告诉回调该指标是越小越好还是越大越好。save_best_onlyTrue配合monitorval_loss的含义是只有当当前epoch的val_loss比历史上所有epoch都好时才会覆盖保存。这个回调最重要的一点是save_weights_only这个参数。如果设为True只保存权重文件小恢复时必须重新构建模型结构代码如果设为False保存完整模型包括网络结构、优化器状态、损失函数配置加载后可以直接继续训练。常规训练过程中我推荐save_weights_onlyFalse虽然文件大一点但“开箱即用”的感觉太好了根本不需要再去拼结构代码。4.2 EarlyStopping 与 ReduceLROnPlateau训练的黄金搭档EarlyStopping用来防止过拟合如果连续多个epoch验证集指标不再变好就提前终止训练。from tensorflow.keras.callbacks import EarlyStopping, ReduceLROnPlateau early_stop EarlyStopping( monitorval_loss, patience5, restore_best_weightsTrue, ) reduce_lr ReduceLROnPlateau( monitorval_loss, factor0.5, patience3, min_lr1e-6, )一个非常实用也容易被忽略的参数是restore_best_weightsTrue。如果训练在epoch 20被提前终止而最好的val_loss出现在epoch 15那么当restore_best_weightsTrue时训练结束后模型会自动回滚到epoch 15的权重。如果是在项目里做模型对比这一点特别重要——否则你拿到的模型是epoch 20的权重并不是验证集表现最好的那个。ReduceLROnPlateau则是等验证指标停滞时把学习率降一半让模型在更小的步长下继续精细搜索。我习惯把patience设置成EarlyStopping的一半让学习率先面临缩减的“警告”如果连续两三次降学习率还止不住颓势再触发早停。4.3 TensorBoard训练过程的可视化仪表盘TensorBoard的作用不只是画loss曲线。它会自动记录训练过程中的很多信息包括计算图、梯度分布、权重分布、样本图像、文本嵌入等。from tensorflow.keras.callbacks import TensorBoard tensorboard TensorBoard( log_dirlogs/run_001, histogram_freq1, write_graphTrue, write_imagesFalse, )在日志目录中积累几次训练之后运行tensorboard --logdir logs就能在浏览器里同时对比不同实验的曲线。我在实际调参时最常看的是Scalars面板里的epoch_loss和epoch_accuracy如果发现在某一步loss突然冒出尖峰再切到Distributions面板看权重分布是否出现异常。很多情况下loss曲线上的异常尖峰都对应着某些层的梯度爆炸提前发现就避免了训练白费。这里有个实用技巧给每次训练单独建一个目录比如logs/run_001、logs/run_002目录名里带上你这次实验想验证的变量比如logs/lr_1e-3_bs_64。这样TensorBoard能把这些实验放在同一个对比视图里一目了然。如果什么都不管全往一个目录写TensorBoard会把它们视为同一份实验曲线画在一起反而乱。4.4 自定义回调把业务逻辑插进训练流程内置回调解决不了所有问题Keras的解决方案是允许你写自定义回调。继承tf.keras.callbacks.Callback重写几个关键方法on_epoch_end、on_batch_end、on_train_begin等。一个非常实用的场景我需要在每个epoch结束后向训练集上重新做一次预测计算模型在某个业务指标上的表现这个指标不在loss里。于是写一个自定义回调class EvaluateBusinessMetric(tf.keras.callbacks.Callback): def __init__(self, validation_data, threshold0.5): super().__init__() self.validation_data validation_data self.threshold threshold def on_epoch_end(self, epoch, logsNone): x_val, y_val self.validation_data y_pred self.model.predict(x_val, verbose0) y_pred_label (y_pred self.threshold).astype(int32) # 计算业务自定义的覆盖率指标 coverage (y_pred_label.sum() / len(y_pred_label)) * 100 print(fepoch {epoch 1} - coverage under threshold: {coverage:.2f}%)自定义回调的要求是不要修改logs字典的keyKeras会用这些名称记录日志也不要自己调用model.fit但model.predict完全可以使用。这样做的好处是你在每个epoch结束的时候能拿到肉眼可见的业务指标反馈而不是只盯着一堆loss数值。5. 模型保存、加载与恢复工程化基本功5.1 三种保存方式怎么选Keras里保存模型有几种常见方式选择的原则其实很简单取决于你想实现什么目的。表格对比一下方式保存内容文件格式典型场景model.save_weights()只有权重.h5/.weights.h5临时存档、迁移学习model.save()完整模型结构权重优化器状态.keras/.h5常规训练完存档、恢复训练TFSavedModel完整模型且带推理接口SavedModel目录上线到TensorFlow Servingmodel.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(optimizeradam, lossbinary_crossentropy, metrics[accuracy]) restored.fit(train_dataset, epochs40, initial_epoch20)注意initial_epoch20这个参数它告诉fit这是从第20个epoch继续跑而不是从头开始。如果你写epochs40且不写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_wordsvocab_size ) x_text_train tf.keras.preprocessing.sequence.pad_sequences(x_text_train, maxlenmax_len) x_text_test tf.keras.preprocessing.sequence.pad_sequences(x_text_test, maxlenmax_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,), nametext) x_text layers.Embedding(vocab_size, 128, mask_zeroTrue)(text_input) x_text layers.Bidirectional(layers.LSTM(64, dropout0.2))(x_text) # 数值输入 num_input layers.Input(shape(num_features,), namenumeric) x_num layers.Dense(32, activationrelu)(num_input) x_num layers.Dense(16, activationrelu)(x_num) # 融合层 combined layers.concatenate([x_text, x_num]) combined layers.Dense(64, activationrelu)(combined) combined layers.Dropout(0.3)(combined) output layers.Dense(1, activationsigmoid, namepurchase)(combined) model Model(inputs[text_input, num_input], outputsoutput) model.compile( optimizertf.keras.optimizers.Adam(learning_rate1e-3), lossbinary_crossentropy, metrics[accuracy], ) model.summary()这里我使用了mask_zeroTrue。它的作用是告诉Embedding层输入中ID为0的位置是填充位在后面的LSTM中这些位置也会被跳过。如果文本长度参差不齐pad之后填入的0会干扰LSTM的语义计算mask_zero能有效屏蔽它们。6.3 训练配置与回调配合真实训练中我会把第4节讲到的回调全部接上checkpoint tf.keras.callbacks.ModelCheckpoint( best_model.keras, monitorval_accuracy, modemax, save_best_onlyTrue, ) early_stop tf.keras.callbacks.EarlyStopping( monitorval_loss, patience3, restore_best_weightsTrue, ) reduce_lr tf.keras.callbacks.ReduceLROnPlateau( monitorval_loss, factor0.5, patience2, min_lr1e-6 ) history model.fit( [x_text_train, x_num_train], y_train, validation_data([x_text_test, x_num_test], y_test), epochs30, batch_size64, 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(fTest accuracy: {acc:.4f})这个实战虽然数据是模拟的但结构和流程完全可以套用到真实项目。从函数式API定义多输入、用内置和自定义层组合提取特征、再到回调机制配合训练、最后保存加载完整模型——这一套操作就是日常深度学习项目的完整闭环。7. 训练效率与稳定性优化的几条实践经验7.1 学习率不要一条路走到黑固定学习率跑完全程基本不是最优解。我常用的策略是配合ReduceLROnPlateau动态降学习率更精细一点的做法是使用余弦退火调度器lr_schedule tf.keras.optimizers.schedules.CosineDecay( initial_learning_rate1e-3, decay_steps5000, alpha1e-5, ) optimizer tf.keras.optimizers.Adam(learning_ratelr_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.Datasetdataset 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_callstf.data.AUTOTUNE并行度会更高dataset dataset.map(preprocess_function, num_parallel_callstf.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才真正变成了你自己的工具。接下来再遇到自己的业务模型就只剩下把数据喂进去这一件事了。
网站建设高端定制企业官网