基于ALTO标注坐标分割图像行用于TrOCR训练时遇索引越界错误
问题描述
需要基于ALTO格式标注数据训练TrOCR模型,将整页图像裁剪为单行图像,但运行代码时触发索引越界错误:
Traceback (most recent call last): File "/home/incognito/TrOCR-py3.10/LINES/alto2lines.py", line 95, in <module> xys[i]=correct_xy([list(map(int,s.split(','))) for s in xy]) File "/home/incognito/TrOCR-py3.10/LINES/alto2lines.py", line 19, in correct_xy y=np.linspace(xx[i][1],xx[i+1][1],c[i]+2,dtype=int)[1:-1] IndexError: list index out of range
错误原因分析
BASELINE解析格式错误:
ALTO标准中,BASELINE属性的格式是空格分隔的x、y坐标交替值(例如示例中的"627 397 964 397 1713 373 2473 373"对应4个坐标点:(627,397)、(964,397)、(1713,373)、(2473,373))。但当前代码错误地将每个单个数字作为独立元素拆分后用逗号分割,导致每个"坐标"变成了单元素列表(如[627]、[397]),在correct_xy函数中访问xx[i][1](y值)时触发索引越界。correct_xy函数逻辑死代码:
在if len(xy)<21分支中执行return xx后,后续的if len(xy)>21分支永远无法被执行,属于无效代码。
解决方法
1. 修正BASELINE解析逻辑
将拆分后的数字列表按两两分组,生成正确的(x,y)坐标对:
# 替换原代码中解析BASELINE的部分 xy_str = element.get('BASELINE').split(' ') # 将列表按两两分组为坐标对 xy_pairs = [tuple(map(int, xy_str[i:i+2])) for i in range(0, len(xy_str), 2)] xys[i] = correct_xy(xy_pairs)
2. 修复correct_xy函数的逻辑错误
调整条件判断结构,移除死代码,同时增加单点防护:
def correct_xy(xy): xx = sorted(xy) if len(xx) < 21: l = len(xx) missing = 21 - l if l > 1: # 避免只有1个点时无法生成插值 c = Counter(np.random.choice(range(l-1), missing)) for i in range(l-1): x = np.linspace(xx[i][0], xx[i+1][0], c[i]+2, dtype=int)[1:-1] y = np.linspace(xx[i][1], xx[i+1][1], c[i]+2, dtype=int)[1:-1] if len(x): xx += list(map(list, zip(x, y))) return xx elif len(xx) > 21: new_xy = [0]*21 new_xy[0] = xx[0] step = (len(xx)-2)/19 for i in range(1,20): new_xy[i] = xx[int(np.floor(i*step))] new_xy[-1] = xx[-1] return new_xy return xx
3. 增加异常处理(可选)
为避免无效的BASELINE数据(如长度为奇数、空值)导致崩溃,可增加校验:
xy_str = element.get('BASELINE', '').strip().split(' ') if len(xy_str) % 2 != 0 or len(xy_str) < 2: print(f"无效的BASELINE数据,跳过TextLine {i}") to_del.append(i) continue xy_pairs = [tuple(map(int, xy_str[j:j+2])) for j in range(0, len(xy_str), 2)] xys[i] = correct_xy(xy_pairs)
完整修正后的关键代码片段
# 替换原代码中解析TextLine的循环部分 for i,element in enumerate(root.iter(prefix+'TextLine')): #images boxes[i] = tuple([int(element.get(s)) for s in ['HPOS','VPOS','HEIGHT','WIDTH']]) xy_str = element.get('BASELINE', '').strip().split(' ') # 校验BASELINE格式有效性 if len(xy_str) % 2 != 0 or len(xy_str) < 2: print(f"无效的BASELINE数据,跳过TextLine {i}") to_del.append(i) continue # 生成正确的坐标对 xy_pairs = [tuple(map(int, xy_str[j:j+2])) for j in range(0, len(xy_str), 2)] xys[i] = correct_xy(xy_pairs)
内容的提问来源于stack exchange,提问作者bsteo
相关产品推荐
相关产品推荐

