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

如何在TrOCR模型中获取置信度值并设置自定义阈值?

TrOCR置信度获取与阈值设置方案

一、从TrOCR模型提取置信度值

TrOCR基于Transformer架构,推理时可通过两种方式获取置信度:

1. 字符级置信度(单字符预测概率)

直接调用模型并获取token级的预测概率,对应到每个识别字符:

from transformers import TrOCRProcessor, VisionEncoderDecoderModel
import torch

# 加载微调后的模型和处理器
processor = TrOCRProcessor.from_pretrained("你的微调模型路径")
model = VisionEncoderDecoderModel.from_pretrained("你的微调模型路径")

# 输入图像(PIL格式)
pixel_values = processor(images=image, return_tensors="pt").pixel_values

# 推理时返回分数和生成结果
outputs = model.generate(
    pixel_values,
    return_dict_in_generate=True,
    output_scores=True
)

# 将token分数转换为概率,提取每个字符的最大置信度
scores = outputs.scores
probs = [torch.softmax(score, dim=-1) for score in scores]
char_confidences = [prob.max().item() for prob in probs]

# 解码生成文本
generated_text = processor.batch_decode(outputs.sequences, skip_special_tokens=True)[0]

此时char_confidences列表中的值与generated_text的每个字符一一对应,代表单字符的预测置信度。

2. 文本级置信度(整句聚合值)

若需要整段文本的整体置信度,可对字符级置信度做聚合:

  • 平均置信度:avg_confidence = sum(char_confidences) / len(char_confidences)
  • 最小置信度:min_confidence = min(char_confidences)(适合严格场景,只要有一个字符置信度低就判定整体不可靠)

二、设置阈值的方法

1. 基于验证集统计确定阈值

  • 用标注好的验证集跑一遍模型,收集所有样本的文本级置信度(平均/最小)和对应识别准确率。
  • 绘制置信度与准确率的关系曲线,找到业务需求的平衡点:比如当置信度≥0.8时,识别准确率达95%以上,即可将0.8设为阈值。
  • 也可计算不同阈值下的精确率、召回率,根据业务倾向选择:优先精确率选较高阈值,优先召回率选较低阈值。

2. 动态阈值(可选)

如果不同场景的文本难度差异大,可按文本长度、字体类型等分类设置不同阈值:比如短文本(≤5字符)设0.75,长文本(≥10字符)设0.85。

三、构建自定义阈值模型

可以在TrOCR推理流程外添加判断逻辑,或封装成带阈值的自定义Pipeline:

1. 简单阈值过滤函数

def ocr_with_threshold(image, threshold=0.8):
    # 复用前面的推理代码,获取generated_text和avg_confidence
    pixel_values = processor(images=image, return_tensors="pt").pixel_values
    outputs = model.generate(pixel_values, return_dict_in_generate=True, output_scores=True)
    generated_text = processor.batch_decode(outputs.sequences, skip_special_tokens=True)[0]
    scores = outputs.scores
    probs = [torch.softmax(score, dim=-1) for score in scores]
    char_confidences = [prob.max().item() for prob in probs]
    avg_confidence = sum(char_confidences)/len(char_confidences)
    
    # 阈值判断
    status = "可靠" if avg_confidence >= threshold else "不可靠"
    return {"text": generated_text, "confidence": avg_confidence, "status": status}

2. 封装自定义Pipeline

如果需要更通用的调用方式,可继承Pipeline类整合阈值逻辑:

from transformers import Pipeline

class TrOCRWithThresholdPipeline(Pipeline):
    def __init__(self, model, processor, threshold=0.8, **kwargs):
        super().__init__(
            model=model,
            tokenizer=processor.tokenizer,
            feature_extractor=processor.feature_extractor,
            **kwargs
        )
        self.threshold = threshold
        self.processor = processor

    def _sanitize_parameters(self, **kwargs):
        params = {}
        if "threshold" in kwargs:
            params["threshold"] = kwargs["threshold"]
        return {}, params, {}

    def preprocess(self, image):
        return self.processor(images=image, return_tensors="pt")

    def _forward(self, model_inputs, threshold=None):
        threshold = threshold or self.threshold
        outputs = self.model.generate(**model_inputs, return_dict_in_generate=True, output_scores=True)
        generated_text = self.processor.batch_decode(outputs.sequences, skip_special_tokens=True)[0]
        
        # 计算置信度
        scores = outputs.scores
        probs = [torch.softmax(score, dim=-1) for score in scores]
        char_confidences = [prob.max().item() for prob in probs]
        avg_confidence = sum(char_confidences)/len(char_confidences)
        
        # 阈值判断
        status = "可靠" if avg_confidence >= threshold else "不可靠"
        return {"text": generated_text, "confidence": avg_confidence, "status": status}

    def postprocess(self, model_outputs):
        return model_outputs

# 使用示例
pipe = TrOCRWithThresholdPipeline(model=model, processor=processor, threshold=0.8)
result = pipe(image)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 14:18:18