Process a batch of images and questions
(model, batch_images, batch_questions, id_list, all_outputs, has_bbox)
| 114 | json.dump(all_outputs, f, indent=2, ensure_ascii=False) |
| 115 | |
| 116 | def process_batch(model, batch_images, batch_questions, id_list, all_outputs, has_bbox): |
| 117 | """Process a batch of images and questions""" |
| 118 | batch_results = model.segment_objects_batch(batch_images, batch_questions) |
| 119 | |
| 120 | for i, result in enumerate(batch_results): |
| 121 | try: |
| 122 | thinking = result["thinking"] |
| 123 | bboxes = result["bboxes"] |
| 124 | mask_all = result["masks"] |
| 125 | gt_mask = np.array(id_list[i]["mask"]) |
| 126 | |
| 127 | # visualize_result(batch_images[i], mask_all, gt_mask, thinking, batch_questions[i]) |
| 128 | |
| 129 | bbox_iou = 0.0 |
| 130 | if has_bbox: |
| 131 | try: |
| 132 | gt_bbox = id_list[i]["bbox"] |
| 133 | for pred_bbox in bboxes: |
| 134 | if compute_bbox_iou(pred_bbox, gt_bbox) > 0.5: |
| 135 | bbox_iou = 1.0 |
| 136 | break |
| 137 | except Exception as e: |
| 138 | print(f"Bbox error: {e}, Image ID: {id_list[i]['image_id']}, Ann ID: {id_list[i]['ann_id']}") |
| 139 | bbox_iou = 0.0 |
| 140 | |
| 141 | all_outputs.append({ |
| 142 | "image_id": id_list[i]["image_id"], |
| 143 | "ann_id": id_list[i]["ann_id"], |
| 144 | "think": thinking, |
| 145 | # "mask_pred": mask_all.astype(int).tolist(), |
| 146 | # "mask_gt": gt_mask.astype(int).tolist(), |
| 147 | "bbox_iou": bbox_iou, |
| 148 | "anomaly_pred": float(mask_all.sum() / (id_list[i]["img_height"] * id_list[i]["img_width"])), |
| 149 | "anomaly_label": int(gt_mask.sum() > 0) |
| 150 | }) |
| 151 | |
| 152 | except Exception as e: |
| 153 | print(f"Error processing result: {e}") |
| 154 | # Add penalty in this situation |
| 155 | all_outputs.append({ |
| 156 | "image_id": id_list[i]["image_id"], |
| 157 | "ann_id": id_list[i]["ann_id"], |
| 158 | "think": "", |
| 159 | # "mask_pred": np.zeros_like(id_list[i]["mask"]).astype(int).tolist(), |
| 160 | # "mask_gt": id_list[i]["mask"].astype(int).tolist(), |
| 161 | "bbox_iou": 0.0, |
| 162 | "anomaly_pred": float(id_list[i]["mask"].sum() / (id_list[i]["img_height"] * id_list[i]["img_width"])), |
| 163 | "anomaly_label": int(id_list[i]["mask"].sum() > 0) |
| 164 | }) |
| 165 | |
| 166 | if __name__ == "__main__": |
| 167 | main() |
no test coverage detected