"""Local HSS metric evaluation for S23DR 2026.""" import numpy as np from typing import Dict, List, Tuple, Optional import time def compute_hss(pred_vertices, pred_edges, gt_vertices, gt_edges, vert_thresh=0.5, edge_thresh=0.5): try: from hoho2025.metric_helper import hss t0 = time.time() result = hss(pred_vertices, pred_edges, gt_vertices, gt_edges, vert_thresh=vert_thresh, edge_thresh=edge_thresh) elapsed = time.time() - t0 return {'hss': float(result.hss), 'f1': float(result.f1), 'iou': float(result.iou), 'time': elapsed} except Exception as e: return {'hss': 0.0, 'f1': 0.0, 'iou': 0.0, 'time': 0.0, 'error': str(e)} def compute_corner_f1(pred_vertices, gt_vertices, thresh=0.5): from scipy.optimize import linear_sum_assignment if len(pred_vertices) == 0 or len(gt_vertices) == 0: return {'f1': 0.0, 'precision': 0.0, 'recall': 0.0, 'tp': 0} diff = pred_vertices[:, None, :] - gt_vertices[None, :, :] dists = np.sqrt((diff ** 2).sum(axis=-1)) row_ind, col_ind = linear_sum_assignment(dists) tp = (dists[row_ind, col_ind] <= thresh).sum() precision = tp / len(pred_vertices) recall = tp / len(gt_vertices) f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0 return {'f1': float(f1), 'precision': float(precision), 'recall': float(recall), 'tp': int(tp), 'num_pred': len(pred_vertices), 'num_gt': len(gt_vertices)} def evaluate_batch(predictions, ground_truths, verbose=True): hss_scores, f1_scores, iou_scores, times, errors = [], [], [], [], [] for i, ((pred_v, pred_e), (gt_v, gt_e)) in enumerate(zip(predictions, ground_truths)): result = compute_hss(pred_v, pred_e, gt_v, gt_e) hss_scores.append(result['hss']) f1_scores.append(result['f1']) iou_scores.append(result['iou']) times.append(result['time']) if 'error' in result: errors.append((i, result['error'])) if verbose and (i + 1) % 10 == 0: print(f" [{i+1}/{len(predictions)}] HSS={np.mean(hss_scores):.4f} F1={np.mean(f1_scores):.4f} IoU={np.mean(iou_scores):.4f}") return { 'hss_mean': float(np.mean(hss_scores)), 'hss_std': float(np.std(hss_scores)), 'hss_median': float(np.median(hss_scores)), 'f1_mean': float(np.mean(f1_scores)), 'iou_mean': float(np.mean(iou_scores)), 'total_time': float(sum(times)), 'avg_time': float(np.mean(times)), 'num_samples': len(predictions), 'num_errors': len(errors), 'errors': errors[:10], } def quick_vertex_eval(pred_vertices, gt_vertices): return compute_corner_f1(pred_vertices, gt_vertices)