保存可视化结果
(self, result: Dict, output_path: str)
| 1205 | f.write(f'{merged_text}\n\n') |
| 1206 | |
| 1207 | def save_visualization(self, result: Dict, output_path: str): |
| 1208 | """保存可视化结果""" |
| 1209 | img_name, img_dir = _get_image_name_and_dir(result, output_path) |
| 1210 | vis_path = os.path.join(img_dir, f'{img_name}_vis.jpg') |
| 1211 | |
| 1212 | if '_page_image' in result: |
| 1213 | image = result['_page_image'].copy() |
| 1214 | else: |
| 1215 | image = cv2.imread(result['input_path']) |
| 1216 | |
| 1217 | for box_info in result['layout_results']['boxes']: |
| 1218 | x1, y1, x2, y2 = map(int, box_info['coordinate']) |
| 1219 | label = box_info['label'] |
| 1220 | score = box_info['score'] |
| 1221 | |
| 1222 | # 提取基础标签名(去除编号后缀,如 text_01 -> text) |
| 1223 | base_label = label.rsplit('_', 1)[0] if '_' in label and label.rsplit('_', 1)[1].isdigit() else label |
| 1224 | |
| 1225 | # 获取颜色,如果没有定义则使用默认红色 |
| 1226 | color = self.colors.get(base_label, (255, 0, 0)) |
| 1227 | |
| 1228 | cv2.rectangle(image, (x1, y1), (x2, y2), color, 2) |
| 1229 | cv2.putText(image, f'{label}: {score: .2f}', (x1, y1 - 10), |
| 1230 | cv2.FONT_HERSHEY_SIMPLEX, 0.5, color, 2) |
| 1231 | |
| 1232 | cv2.imwrite(vis_path, image) |
| 1233 | # logger.info(f" Saved visualization to {vis_path}") |
| 1234 | |
| 1235 | |
| 1236 | # ==================== Main Function ==================== |
no test coverage detected