如何在TrOCR模型中获取行图像中每个字符的边界框?
在TrOCR中获取行图像的字符边界框
TrOCR是序列到序列模型,原生不直接输出字符级边界框。要实现这个需求,需要结合模型注意力映射或视觉后处理工具,以下是两种可行方案:
方法一:利用Decoder注意力关联字符与图像区域
TrOCR的Decoder在生成每个字符时,会对Encoder输出的视觉特征图产生注意力权重,我们可以通过这些权重将字符映射回原图的对应区域。
代码实现
import requests import torch from PIL import Image, ImageDraw from transformers import TrOCRProcessor, VisionEncoderDecoderModel # 加载模型和处理器 processor = TrOCRProcessor.from_pretrained("microsoft/trocr-base-handwritten") model = VisionEncoderDecoderModel.from_pretrained("microsoft/trocr-base-handwritten") # 加载目标图像 url = "https://fki.tic.heia-fr.ch/static/img/a01-122-02.jpg" image = Image.open(requests.get(url, stream=True).raw).convert("RGB") pixel_values = processor(image, return_tensors="pt").pixel_values # 生成文本并获取注意力权重 outputs = model.generate( pixel_values, output_attentions=True, return_dict_in_generate=True ) generated_ids = outputs.sequences generated_text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0] # 提取最后一层Decoder的注意力权重 decoder_attentions = outputs.decoder_attentions[-1] # ViT编码器输出的特征图尺寸:384x384输入对应12x12特征图(patch size=32) feature_map_size = (12, 12) img_w, img_h = image.size # 计算每个字符的边界框 char_bboxes = [] for idx, char in enumerate(generated_text): # 跳过<BOS> token,取当前字符对应的注意力权重 attn_weights = decoder_attentions[0, :, idx+1, :].mean(dim=0).reshape(feature_map_size) # 找到权重最高的特征patch _, max_idx = torch.max(attn_weights.flatten(), dim=0) y, x = divmod(max_idx.item(), feature_map_size[1]) # 将patch坐标映射回原图 patch_w = img_w / feature_map_size[1] patch_h = img_h / feature_map_size[0] left = x * patch_w top = y * patch_h right = (x + 1) * patch_w bottom = (y + 1) * patch_h char_bboxes.append((char, (left, top, right, bottom))) # 可视化边界框 draw = ImageDraw.Draw(image) for char, bbox in char_bboxes: draw.rectangle(bbox, outline="red", width=2) draw.text((bbox[0], bbox[1]-10), char, fill="red") image.save("trocr_attn_bboxes.png")
注意事项
这种方法是近似映射,注意力权重可能分散在多个patch上,对于连笔手写体的边界框精度会受影响。如果需要更准确的结果,可以对注意力权重进行加权平均计算区域中心。
方法二:结合OpenCV进行字符分割
对于规整的印刷体或手写体,可以先用TrOCR识别文本内容,再用OpenCV对行图像做字符分割,匹配字符与轮廓。
代码实现
import requests import cv2 import numpy as np from PIL import Image, ImageDraw from transformers import TrOCRProcessor, VisionEncoderDecoderModel # 加载模型和图像 processor = TrOCRProcessor.from_pretrained("microsoft/trocr-base-handwritten") model = VisionEncoderDecoderModel.from_pretrained("microsoft/trocr-base-handwritten") url = "https://fki.tic.heia-fr.ch/static/img/a01-122-02.jpg" image = Image.open(requests.get(url, stream=True).raw).convert("RGB") pixel_values = processor(image, return_tensors="pt").pixel_values generated_text = processor.batch_decode(model.generate(pixel_values), skip_special_tokens=True)[0] # 转换为OpenCV格式处理 img_cv = cv2.cvtColor(np.array(image), cv2.COLOR_RGB2BGR) gray = cv2.cvtColor(img_cv, cv2.COLOR_BGR2GRAY) # 自适应二值化,适配手写体亮度不均的情况 thresh = cv2.adaptiveThreshold(gray, 255, cv2.ADAPTIVE_THRESH_GAUSSIAN_C, cv2.THRESH_BINARY_INV, 11, 2) # 提取字符轮廓并按左到右排序 contours, _ = cv2.findContours(thresh, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) contours = sorted(contours, key=lambda c: cv2.boundingRect(c)[0]) # 匹配轮廓与识别字符 char_bboxes = [] for i, contour in enumerate(contours): if i >= len(generated_text): break x, y, w, h = cv2.boundingRect(contour) # 过滤噪声小轮廓 if w < 5 or h < 5: continue char_bboxes.append((generated_text[i], (x, y, x+w, y+h))) # 可视化 draw = ImageDraw.Draw(image) for char, bbox in char_bboxes: draw.rectangle(bbox, outline="blue", width=2) draw.text((bbox[0], bbox[1]-10), char, fill="blue") image.save("trocr_opencv_bboxes.png")
注意事项
这种方法对连笔手写体的分割效果较差,需要根据实际场景调整二值化参数或使用更复杂的分割算法(如基于投影的分割)。
关于你提到的Pillow代码
你看到的Pillow代码属于文本渲染场景:已知字体和文本内容,计算渲染后字符的坐标。而TrOCR是图像识别场景:从图像中识别未知文本,两者逻辑完全不同,无法直接适配。
内容的提问来源于stack exchange,提问作者stats_residue
相关产品推荐
相关产品推荐

