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

如何让TensorFlow图像识别模型仅参考传入的指定关键词?

问题解决思路与代码修改

核心问题分析

你当前的filter_classes函数仅过滤了类别名称列表,但模型是基于灯塔+无人机两类训练的,输出的置信度维度固定为2,因此无法通过过滤类列表改变输出的分数数量。要只获取指定类别的分数,需要直接从模型输出中提取对应类别的索引值。

代码修改方案

1. 添加类别索引获取函数

替换原filter_classes函数,改为获取目标类别在原始类别列表中的索引:

def get_target_class_index(target_class, all_classes):
    # 返回目标类别在原始列表中的位置索引
    return all_classes.index(target_class)

2. 修改主函数逻辑

在get_correct_images中,获取目标类别索引后直接提取对应分数:

import tensorflow
import numpy
import urllib.request
from io import BytesIO
from PIL import Image
from utils.n import n
from utils.c import c

def get_target_class_index(target_class, all_classes):
    return all_classes.index(target_class)

def get_correct_images(question, images):
    model = tensorflow.saved_model.load('./')
    # 原始类别列表,需与训练时完全一致
    all_classes = ["lighthouse" , "drone"]
    # 获取目标类别的索引
    target_idx = get_target_class_index(question, all_classes)
    
    for url in images:
        req = urllib.request.Request(url, headers={'User-Agent': 'Mozilla/5.0'})
        with urllib.request.urlopen(req) as response:
            img_data = response.read()
        img = Image.open(BytesIO(img_data)).convert('RGB')
        img = img.resize((300, 300 * img.size[1] // img.size[0]), Image.ANTIALIAS)
        inp_numpy = numpy.array(img)[None]
        inp = tensorflow.constant(inp_numpy, dtype='float32')
        
        class_scores = model(inp)[0].numpy()
        # 提取指定类别的分数
        target_score = [class_scores[target_idx]]
        
        print("")
        print("Target score:", target_score)
        print("Class : ", question)
        print("Url: ", url)

额外优化建议(解决梯子误识别问题)

当前模型是二分类逻辑(非灯塔即无人机),遇到未训练的梯子时会强行归类到两者之一。若要实现单类别独立判定(仅判断是否属于指定类别,而非二选一),可以:

  • 将模型改为多标签分类架构,为每个类别单独训练二分类器(分别判断是/不是灯塔、是/不是无人机)
  • 在训练数据中加入负样本(如梯子、其他无关图像),让模型学习区分目标类与非目标类

内容的提问来源于stack exchange,提问作者Bungfro

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 14:52:32