""" Scoring logic for the Audio2Tool leaderboard. Mirrors the normalization and metrics of the benchmark code (github.com/audio2tool/Audio2Tool, audio_benchmark/evaluation/metrics.py) so leaderboard scores are directly comparable to local runs. """ import json import re from typing import Any, Dict, List, Tuple CALL_PATTERN = re.compile(r"(\w+)\s*\((.*?)\)") PARAM_PATTERN = re.compile(r"(\w+)\s*=\s*['\"]?([^,'\"]+)['\"]?") def parse_tool_calls(raw: str) -> List[Tuple[str, Dict[str, str]]]: """Parse a prediction string into [(tool_name, params), ...]. Same regex as the benchmark's BaseAudioModel.parse_tool_call, extended to capture multiple calls (for multi-intent tiers). The first call is the primary tool. """ calls = [] for m in CALL_PATTERN.finditer(raw or ""): name = m.group(1) params = {} params_str = m.group(2) if params_str.strip(): for pm in PARAM_PATTERN.finditer(params_str): params[pm.group(1)] = pm.group(2).strip() calls.append((name, params)) if not calls: # Fallback: bare tool name(s) without parentheses for word in (raw or "").split(): token = re.fullmatch(r"\w+", word.strip()) if token: calls.append((token.group(0), {})) break return calls def normalize_tool_name(name: str) -> str: if not name: return "" name = re.sub(r"\(.*\)", "", name) return name.strip().lower() def normalize_parameters(params: Dict[str, Any]) -> Dict[str, str]: if not params: return {} normalized = {} for key, value in params.items(): norm_key = key.strip().lower() if value is None: norm_value = "" elif isinstance(value, bool): norm_value = "true" if value else "false" elif isinstance(value, (int, float)): norm_value = str(value) else: norm_value = str(value).strip().lower() normalized[norm_key] = norm_value return normalized def param_f1(predicted: Dict[str, Any], ground_truth: Dict[str, Any]) -> Tuple[float, bool]: """Return (f1, exact_match) for parameter dicts.""" pred = normalize_parameters(predicted) gt = normalize_parameters(ground_truth) if not gt and not pred: return 1.0, True if not gt or not pred: return 0.0, False correct = sum(1 for k, v in gt.items() if pred.get(k) == v) precision = correct / len(pred) recall = correct / len(gt) f1 = 2 * precision * recall / (precision + recall) if (precision + recall) else 0.0 exact = pred == gt return f1, exact def score_predictions( predictions: Dict[str, str], gt_rows: List[Dict[str, Any]], ) -> Dict[str, Any]: """ Score {sample_id: prediction string} against ground-truth rows. Returns per-tier and overall metrics plus coverage information. """ by_tier: Dict[str, List[Dict[str, float]]] = {} tier_totals: Dict[str, int] = {} missing = [] for gt in gt_rows: tier = gt["tier"] tier_totals[tier] = tier_totals.get(tier, 0) + 1 raw = predictions.get(gt["sample_id"]) if raw is None: missing.append(gt["sample_id"]) continue calls = parse_tool_calls(raw) pred_tool, pred_params = calls[0] if calls else ("", {}) tool_ok = normalize_tool_name(pred_tool) == normalize_tool_name(gt["tool_name"]) f1, exact_params = param_f1(pred_params, gt.get("extracted_params") or {}) by_tier.setdefault(tier, []).append({ "tool": float(tool_ok), "exact": float(tool_ok and exact_params), "f1": f1 if tool_ok else 0.0, }) tiers = {} for tier, scored in sorted(by_tier.items()): n = len(scored) tiers[tier] = { "n": n, "coverage": round(n / tier_totals[tier], 3), "tool_accuracy": round(sum(s["tool"] for s in scored) / n, 4), "exact_match": round(sum(s["exact"] for s in scored) / n, 4), "param_f1": round(sum(s["f1"] for s in scored) / n, 4), } if not tiers: raise ValueError("No predictions matched any eval sample_id.") avg = lambda k: round(sum(t[k] for t in tiers.values()) / len(tiers), 4) return { "tiers": tiers, "overall": { "tool_accuracy": avg("tool_accuracy"), "exact_match": avg("exact_match"), "param_f1": avg("param_f1"), "n_tiers": len(tiers), }, "n_scored": sum(t["n"] for t in tiers.values()), "n_missing": len(missing), } def load_predictions_jsonl(path: str) -> Dict[str, str]: """Load and validate a predictions JSONL file.""" predictions = {} with open(path) as f: for i, line in enumerate(f, 1): line = line.strip() if not line: continue try: row = json.loads(line) except json.JSONDecodeError as e: raise ValueError(f"Line {i}: invalid JSON ({e})") if "sample_id" not in row or "prediction" not in row: raise ValueError(f"Line {i}: rows need 'sample_id' and 'prediction' keys") if not isinstance(row["prediction"], str): raise ValueError(f"Line {i}: 'prediction' must be a string") predictions[row["sample_id"]] = row["prediction"] if not predictions: raise ValueError("File contains no predictions.") return predictions