Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
138 changes: 136 additions & 2 deletions benchmarks/bfcl_tool_selection/llm_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,7 @@

DEFAULT_OFFICIAL_MODEL_NAME = "qwen3-32b-FC"
BFCL_RESULT_ARGUMENT_FORMATS = ("json-string", "decoded")
MODEL_CASE_CACHE_VERSION = 8
MODEL_CASE_CACHE_VERSION = 9


@dataclass(frozen=True)
Expand Down Expand Up @@ -90,6 +90,7 @@ class BFCLModelCaseEvaluation:
strict_exact_match: float
official_ast_exact_match: float
evaluator_exact_match: float
equivalence_adjusted_exact_match: float
input_tokens: int
output_tokens: int
latency_ms: float
Expand Down Expand Up @@ -563,6 +564,13 @@ def evaluate_model_case(
if evaluator == "official"
else match["strict_exact_match"]
)
equivalence_adjusted_exact_match = _equivalence_adjusted_exact_match(
expected_calls,
predicted_calls,
tools_by_name=tools_by_name,
category=category,
strict_exact_match=match["strict_exact_match"],
)
error = response.error or official["official_error"] or match["error"]
failure_category = _classify_failure(
expected_calls=expected_calls,
Expand Down Expand Up @@ -600,6 +608,7 @@ def evaluate_model_case(
strict_exact_match=match["strict_exact_match"],
official_ast_exact_match=official["official_ast_exact_match"],
evaluator_exact_match=evaluator_exact_match,
equivalence_adjusted_exact_match=equivalence_adjusted_exact_match,
input_tokens=response.input_tokens,
output_tokens=response.output_tokens,
latency_ms=round(latency_ms, 3),
Expand Down Expand Up @@ -1105,6 +1114,111 @@ def _failure_tags(
return tags


def _equivalence_adjusted_exact_match(
expected_calls: list[ExpectedToolCall],
predicted_calls: list[PredictedToolCall],
*,
tools_by_name: dict[str, dict[str, Any]],
category: str,
strict_exact_match: float,
) -> float:
"""Return exact-match credit after high-confidence equivalent tool surfaces.

BFCL contains near-duplicate tools whose names differ while their function
surface is effectively identical. This metric keeps the strict BFCL-style
exact score intact, but also reports whether a miss is semantically the same
tool surface with matching argument values.
"""

if strict_exact_match == 1.0:
return 1.0
if len(expected_calls) != len(predicted_calls):
return 0.0

if "parallel" in category:
used: set[int] = set()
for expected in expected_calls:
matched = False
for index, predicted in enumerate(predicted_calls):
if index in used:
continue
if _equivalent_call_matches(expected, predicted, tools_by_name=tools_by_name):
used.add(index)
matched = True
break
if not matched:
return 0.0
return 1.0

return float(
all(
_equivalent_call_matches(expected, predicted, tools_by_name=tools_by_name)
for expected, predicted in zip(expected_calls, predicted_calls, strict=True)
)
)


def _equivalent_call_matches(
expected: ExpectedToolCall,
predicted: PredictedToolCall,
*,
tools_by_name: dict[str, dict[str, Any]],
) -> bool:
if expected.name == predicted.name:
return _single_call_matches(expected, predicted)
if not _tool_names_are_equivalent(expected.name, predicted.name, tools_by_name=tools_by_name):
return False
return _semantic_argument_values_match(expected, predicted)


def _tool_names_are_equivalent(
left: str,
right: str,
*,
tools_by_name: dict[str, dict[str, Any]],
) -> bool:
groups = build_tool_equivalence_groups([left, right], tools_by_name)
return any({left, right}.issubset(set(group.get("members") or [])) for group in groups)


def _semantic_argument_values_match(
expected: ExpectedToolCall,
predicted: PredictedToolCall,
) -> bool:
predicted_values = list(predicted.arguments.values())
used: set[int] = set()

for possible_values in expected.arguments.values():
if _allows_missing(possible_values):
continue
index = _find_matching_argument_value(predicted_values, possible_values, used)
if index is None:
return False
used.add(index)

for possible_values in expected.arguments.values():
if not _allows_missing(possible_values):
continue
index = _find_matching_argument_value(predicted_values, possible_values, used)
if index is not None:
used.add(index)

return len(used) == len(predicted_values)


def _find_matching_argument_value(
values: list[Any],
possible_values: list[Any],
used: set[int],
) -> int | None:
for index, value in enumerate(values):
if index in used:
continue
if _value_matches(value, possible_values):
return index
return None


def _has_equivalent_expected_and_predicted_tool(
expected_calls: list[ExpectedToolCall],
predicted_calls: list[PredictedToolCall],
Expand Down Expand Up @@ -1502,6 +1616,10 @@ def _case_from_dict(payload: dict[str, Any]) -> BFCLModelCaseEvaluation:
data.setdefault("official_error_type", "")
data.setdefault("official_error", "")
data.setdefault("failure_tags", [])
data.setdefault(
"equivalence_adjusted_exact_match",
float(data.get("evaluator_exact_match") or 0.0),
)
return BFCLModelCaseEvaluation(**data)


Expand All @@ -1521,6 +1639,15 @@ def _summarize(
"strict_exact_match": _mean(row.strict_exact_match for row in rows),
"official_ast_exact_match": _mean(row.official_ast_exact_match for row in rows),
"evaluator_exact_match": _mean(row.evaluator_exact_match for row in rows),
"equivalence_adjusted_exact_match": _mean(
row.equivalence_adjusted_exact_match for row in rows
),
"equivalence_adjusted_exact_gain": _mean(
row.equivalence_adjusted_exact_match - row.evaluator_exact_match for row in rows
),
"equivalence_adjusted_exact_case_count": sum(
int(row.equivalence_adjusted_exact_match > row.evaluator_exact_match) for row in rows
),
"avg_input_tokens": round(_mean(row.input_tokens for row in rows), 1),
"avg_output_tokens": round(_mean(row.output_tokens for row in rows), 1),
"avg_latency_ms": round(_mean(row.latency_ms for row in rows), 1),
Expand Down Expand Up @@ -1578,6 +1705,9 @@ def bfcl_result_rows(
"failure_category": str(case.get("failure_category") or ""),
"failure_tags": list(case.get("failure_tags") or []),
"evaluator_exact_match": float(case.get("evaluator_exact_match") or 0.0),
"equivalence_adjusted_exact_match": float(
case.get("equivalence_adjusted_exact_match") or 0.0
),
},
}
)
Expand Down Expand Up @@ -1641,6 +1771,7 @@ def print_report(report: dict[str, Any]) -> None:
"retrieval@K={retrieval:.2f} call_rate={call:.2f} "
"func_exact={func:.2f} arg_names={arg_names:.2f} "
"arg_values={arg_values:.2f} strict={strict:.2f} exact={exact:.2f} "
"equiv_exact={equiv_exact:.2f} "
"latency={latency:.1f}ms".format(
retrieval=summary["retrieval_recall_at_k"],
call=summary["model_tool_call_rate"],
Expand All @@ -1649,6 +1780,7 @@ def print_report(report: dict[str, Any]) -> None:
arg_values=summary["argument_value_exact_match"],
strict=summary["strict_exact_match"],
exact=summary["evaluator_exact_match"],
equiv_exact=summary["equivalence_adjusted_exact_match"],
latency=summary["avg_latency_ms"],
)
)
Expand All @@ -1664,13 +1796,15 @@ def print_report(report: dict[str, Any]) -> None:
cat = category["summary"]
print(
" {name}: cases={cases} tools={tools} retrieval@K={retrieval:.2f} "
"strict={strict:.2f} exact={exact:.2f} call_rate={call:.2f} failures={failures}".format(
"strict={strict:.2f} exact={exact:.2f} equiv_exact={equiv_exact:.2f} "
"call_rate={call:.2f} failures={failures}".format(
name=category["category"],
cases=category["case_count"],
tools=category["corpus_tool_count"],
retrieval=cat["retrieval_recall_at_k"],
strict=cat["strict_exact_match"],
exact=cat["evaluator_exact_match"],
equiv_exact=cat["equivalence_adjusted_exact_match"],
call=cat["model_tool_call_rate"],
failures=_format_failure_breakdown(cat.get("failure_breakdown") or {}),
)
Expand Down
41 changes: 39 additions & 2 deletions benchmarks/bfcl_tool_selection/sweep.py
Original file line number Diff line number Diff line change
Expand Up @@ -160,6 +160,10 @@ def _summarize_sweep(
"model_tool_call_rate": summary["model_tool_call_rate"],
"strict_exact_match": summary["strict_exact_match"],
"evaluator_exact_match": summary["evaluator_exact_match"],
"equivalence_adjusted_exact_match": summary.get(
"equivalence_adjusted_exact_match",
summary["evaluator_exact_match"],
),
"avg_latency_ms": summary["avg_latency_ms"],
"failure_breakdown": summary.get("failure_breakdown") or {},
}
Expand Down Expand Up @@ -198,6 +202,10 @@ def _category_summary_rows(run: dict[str, Any]) -> list[dict[str, Any]]:
"model_tool_call_rate": summary.get("model_tool_call_rate", 0.0),
"strict_exact_match": summary.get("strict_exact_match", 0.0),
"evaluator_exact_match": summary.get("evaluator_exact_match", 0.0),
"equivalence_adjusted_exact_match": summary.get(
"equivalence_adjusted_exact_match",
summary.get("evaluator_exact_match", 0.0),
),
"avg_latency_ms": summary.get("avg_latency_ms", 0.0),
"failure_breakdown": summary.get("failure_breakdown") or {},
}
Expand All @@ -213,6 +221,7 @@ def _best_retrieved(rows: list[dict[str, Any]]) -> dict[str, Any]:
retrieved_rows,
key=lambda row: (
row["evaluator_exact_match"],
row["equivalence_adjusted_exact_match"],
row["retrieval_recall_at_k"],
-row["avg_latency_ms"],
),
Expand Down Expand Up @@ -252,12 +261,15 @@ def _row_vs_retrieved_deltas(runs: list[dict[str, Any]]) -> list[dict[str, Any]]
both_fail_retrieved_breakdown: Counter[str] = Counter()
row_pass_retrieved_fail_case_ids: list[str] = []
both_fail_case_ids: list[str] = []
retrieved_adjusted_pass_on_row_pass = 0

for key in paired_keys:
row_case = row_cases[key]
retrieved_case = retrieved_cases[key]
row_pass = _case_passed(row_case)
retrieved_pass = _case_passed(retrieved_case)
if row_pass and _case_equivalence_adjusted_passed(retrieved_case):
retrieved_adjusted_pass_on_row_pass += 1
if row_pass and retrieved_pass:
counts["both_pass"] += 1
elif row_pass and not retrieved_pass:
Expand All @@ -277,6 +289,9 @@ def _row_vs_retrieved_deltas(runs: list[dict[str, Any]]) -> list[dict[str, Any]]
paired_count = len(paired_keys)
row_pass_count = counts["both_pass"] + counts["row_pass_retrieved_fail"]
retrieved_on_row_pass = counts["both_pass"] / row_pass_count if row_pass_count else None
adjusted_on_row_pass = (
retrieved_adjusted_pass_on_row_pass / row_pass_count if row_pass_count else None
)
deltas.append(
{
"repeat": repeat,
Expand All @@ -288,6 +303,9 @@ def _row_vs_retrieved_deltas(runs: list[dict[str, Any]]) -> list[dict[str, Any]]
"both_fail": counts["both_fail"],
"row_pass_count": row_pass_count,
"retrieved_exact_on_row_pass": _round_or_none(retrieved_on_row_pass),
"retrieved_equivalence_adjusted_exact_on_row_pass": _round_or_none(
adjusted_on_row_pass
),
"row_pass_retrieved_fail_rate": _round_or_none(
counts["row_pass_retrieved_fail"] / row_pass_count if row_pass_count else None
),
Expand Down Expand Up @@ -322,6 +340,12 @@ def _case_passed(case: dict[str, Any]) -> bool:
return float(case.get("evaluator_exact_match") or 0.0) >= 1.0


def _case_equivalence_adjusted_passed(case: dict[str, Any]) -> bool:
if _case_passed(case):
return True
return float(case.get("equivalence_adjusted_exact_match") or 0.0) >= 1.0


def _round_or_none(value: float | None) -> float | None:
if value is None:
return None
Expand Down Expand Up @@ -349,6 +373,9 @@ def _summarize_category_repeat_groups(rows: list[dict[str, Any]]) -> list[dict[s
"evaluator_exact_match": _metric_stats(
row["evaluator_exact_match"] for row in group_rows
),
"equivalence_adjusted_exact_match": _metric_stats(
row["equivalence_adjusted_exact_match"] for row in group_rows
),
"strict_exact_match": _metric_stats(
row["strict_exact_match"] for row in group_rows
),
Expand Down Expand Up @@ -378,6 +405,9 @@ def _summarize_repeat_groups(rows: list[dict[str, Any]]) -> list[dict[str, Any]]
"evaluator_exact_match": _metric_stats(
row["evaluator_exact_match"] for row in group_rows
),
"equivalence_adjusted_exact_match": _metric_stats(
row["equivalence_adjusted_exact_match"] for row in group_rows
),
"strict_exact_match": _metric_stats(
row["strict_exact_match"] for row in group_rows
),
Expand Down Expand Up @@ -522,6 +552,9 @@ def _mean_selected_rows(rows: list[dict[str, Any]]) -> dict[str, Any]:
"model_tool_call_rate": _mean(row["model_tool_call_rate"] for row in rows),
"strict_exact_match": _mean(row["strict_exact_match"] for row in rows),
"evaluator_exact_match": _mean(row["evaluator_exact_match"] for row in rows),
"equivalence_adjusted_exact_match": _mean(
row["equivalence_adjusted_exact_match"] for row in rows
),
"avg_latency_ms": _mean(row["avg_latency_ms"] for row in rows),
}

Expand Down Expand Up @@ -577,14 +610,16 @@ def print_report(report: dict[str, Any]) -> None:
failures = _format_failure_breakdown(row.get("failure_breakdown") or {})
print(
"{source:9s} k={top_k:<2} repeat={repeat:<2} cases={cases:<4} "
"retrieval@K={retrieval:.2f} exact={exact:.2f} strict={strict:.2f} "
"retrieval@K={retrieval:.2f} exact={exact:.2f} "
"equiv_exact={equiv_exact:.2f} strict={strict:.2f} "
"latency={latency:.1f}ms failures={failures}".format(
source=row["tool_source"],
top_k=row["top_k"],
repeat=row["repeat"],
cases=row["cases"],
retrieval=row["retrieval_recall_at_k"],
exact=row["evaluator_exact_match"],
equiv_exact=row["equivalence_adjusted_exact_match"],
strict=row["strict_exact_match"],
latency=row["avg_latency_ms"],
failures=failures,
Expand All @@ -593,9 +628,11 @@ def print_report(report: dict[str, Any]) -> None:
best = report["summary"].get("best_retrieved") or {}
if best:
print(
"best_retrieved: k={top_k} exact={exact:.2f} retrieval@K={retrieval:.2f}".format(
"best_retrieved: k={top_k} exact={exact:.2f} equiv_exact={equiv_exact:.2f} "
"retrieval@K={retrieval:.2f}".format(
top_k=best["top_k"],
exact=best["evaluator_exact_match"],
equiv_exact=best["equivalence_adjusted_exact_match"],
retrieval=best["retrieval_recall_at_k"],
)
)
Expand Down
19 changes: 15 additions & 4 deletions docs/research/validation-loop.md
Original file line number Diff line number Diff line change
Expand Up @@ -109,10 +109,12 @@ argument preservation을 우선 보고, row-source에서도 실패한 repeated-c
`row_vs_retrieved_deltas`가 들어간다. 같은 repeat/top-K의 row-source와
retrieved-source를 case-id 기준으로 pair해서 `both_pass`,
`row_pass_retrieved_fail`, `row_fail_retrieved_pass`, `both_fail`,
`retrieved_exact_on_row_pass`, `row_pass_retrieved_fail_breakdown`,
`row_pass_retrieved_fail_tags`, `row_pass_retrieved_fail_case_ids`를 남긴다.
이 값으로 full/smoke 이후 바로 "검색 계층이 실제로 깎은 케이스"만 subset으로
뽑는다.
`retrieved_exact_on_row_pass`,
`retrieved_equivalence_adjusted_exact_on_row_pass`,
`row_pass_retrieved_fail_breakdown`, `row_pass_retrieved_fail_tags`,
`row_pass_retrieved_fail_case_ids`를 남긴다. 이 값으로 full/smoke 이후 바로
"검색 계층이 실제로 깎은 케이스"와 "exact name은 틀렸지만 equivalent tool
surface로 맞은 케이스"를 분리한다.

`benchmarks.bfcl_tool_selection.llm_loop`는 graphify의
`build_tool_equivalence_groups(...)`를 사용해 candidate ambiguity 중 tool name,
Expand All @@ -131,6 +133,15 @@ case-level `target_equivalence_group_count`, summary-level
기록한다. `/tmp/gtc-xgen-equivalence-diagnostics.json` 기준 built-in suite 전체는
평균 equivalence group count `0.333333`, equivalence group case `5`건이다.

`2026-07-19`부터 BFCL model-loop report는
`equivalence_adjusted_exact_match`도 함께 남긴다. 이 값은 기존
`evaluator_exact_match`를 대체하지 않으며, BFCL leaderboard 점수로 사용하지
않는다. strict exact가 실패했더라도 `build_tool_equivalence_groups(...)` 기준
high-confidence equivalent surface이고 argument value가 맞을 때만 별도 credit을
준다. 위 4개 `near_duplicate_tool_surface` subset을 qwen3.6-27B로 재실행한
`/tmp/gtc-bfcl-neardup-adjusted-metric.json` 기준 strict/evaluator exact는
`0.00`, equivalence-adjusted exact는 `1.00`이다.

## 실행 타깃

```bash
Expand Down
Loading