WireFrameDETR / src /metrics.py
StarAtNyte1's picture
Add src/metrics.py
a3f0ec9 verified
Raw
History Blame Contribute Delete
2.75 kB
"""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)