You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在Keras中自定义多输出模型的测试推理逻辑?

实现多输出模型的提前终止推理逻辑

由于Keras默认的predict方法会一次性计算所有输出,要实现第一个输出置信度达标时停止后续层计算的需求,你可以通过拆分训练好的模型、分步执行推理的方式实现,完全不需要改动已有的训练流程(model.compile()/model.fit()部分保持不变)。

核心思路

根据你的模型结构,将其拆分为两部分:

  1. 从输入到第一个输出的计算链路
  2. 从共享层(或第一个输出的前置层)到其他输出的计算链路

推理时先计算第一个输出,判断置信度是否达标:

  • 达标则跳过后续计算
  • 不达标则继续执行剩余输出的计算

代码示例(分支式多输出模型)

假设你的模型是典型的分支结构:输入 -> 共享特征层 -> 分支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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.01 16:37:33