如何在Keras中自定义多输出模型的测试推理逻辑?
实现多输出模型的提前终止推理逻辑
由于Keras默认的predict方法会一次性计算所有输出,要实现第一个输出置信度达标时停止后续层计算的需求,你可以通过拆分训练好的模型、分步执行推理的方式实现,完全不需要改动已有的训练流程(model.compile()/model.fit()部分保持不变)。
核心思路
根据你的模型结构,将其拆分为两部分:
- 从输入到第一个输出的计算链路
- 从共享层(或第一个输出的前置层)到其他输出的计算链路
推理时先计算第一个输出,判断置信度是否达标:
- 达标则跳过后续计算
- 不达标则继续执行剩余输出的计算
代码示例(分支式多输出模型)
假设你的模型是典型的分支结构:输入 -> 共享特征层 -> 分支1(output1)/ 分支2(output2)
import numpy as np from tensorflow.keras.models import Model # 假设训练好的模型名为`trained_model`,共享层名称为`shared_feature`,输出层名称为`output1`、`output2` # 1. 拆分出各部分子模型 shared_model = Model(inputs=trained_model.input, outputs=trained_model.get_layer('shared_feature').output) output1_model = Model(inputs=shared_model.output, outputs=trained_model.get_layer('output1').output) output2_model = Model(inputs=shared_model.output, outputs=trained_model.get_layer('output2').output) # 2. 自定义推理函数(单样本处理) def custom_inference(test_samples, confidence_threshold=0.9): results = [] for sample in test_samples: # 扩展维度以匹配模型输入格式(单样本) sample_input = sample[np.newaxis, ...] # 计算共享特征 shared_feat = shared_model.predict(sample_input, verbose=0) # 计算第一个输出 pred_output1 = output1_model.predict(shared_feat, verbose=0)[0] if np.max(pred_output1) >= confidence_threshold: # 置信度达标,跳过后续计算 results.append({"output1": pred_output1, "output2": None}) else: # 置信度不达标,计算第二个输出 pred_output2 = output2_model.predict(shared_feat, verbose=0)[0] results.append({"output1": pred_output1, "output2": pred_output2}) return results # 3. 批量推理优化版(处理批量数据更高效) def custom_batch_inference(test_batch, confidence_threshold=0.9): # 计算批量共享特征 shared_feats = shared_model.predict(test_batch, verbose=0) # 计算批量第一个输出 batch_output1 = output1_model.predict(shared_feats, verbose=0) # 筛选出置信度不达标的样本索引 low_conf_indices = np.where(np.max(batch_output1, axis=1) < confidence_threshold)[0] # 初始化第二个输出为占位符 batch_output2 = [None] * len(batch_output1) if len(low_conf_indices) > 0: # 仅对低置信度样本计算第二个输出 filtered_feats = shared_feats[low_conf_indices] low_conf_output2 = output2_model.predict(filtered_feats, verbose=0) for idx, pred in zip(low_conf_indices, low_conf_output2): batch_output2[idx] = pred return {"output1": batch_output1, "output2": batch_output2}
适配串联式多输出模型
如果你的模型是串联结构(输入 -> 层1 -> ... -> output1 -> 层N -> ... -> output2),只需调整子模型的拆分逻辑:
# 拆分到第一个输出的子模型 output1_submodel = Model(inputs=trained_model.input, outputs=trained_model.get_layer('output1').output) # 拆分从output1到output2的子模型(假设output1的下一层是`post_output1_layer`) output2_submodel = Model(inputs=trained_model.get_layer('post_output1_layer').input, outputs=trained_model.get_layer('output2').output)
推理逻辑与分支式一致:先算output1,达标则停止,否则继续计算后续层。
关键优势
- 完全复用已训练好的模型参数,无需重新编写训练代码
- 精准控制计算流程,避免不必要的后续层计算,节省推理资源
- 支持单样本和批量数据处理,适配不同推理场景
内容的提问来源于stack exchange,提问作者Andrea
相关产品推荐
相关产品推荐

