如何在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
相关产品推荐
相关产品推荐

