Audio2ToolLeaderboard / scoring.py
Ramit13's picture
Replace template with Audio2Tool leaderboard app
9bc4c6b verified
Raw History Blame Contribute Delete
5.59 kB
"""
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