Skip to content

Commit 8a96eb9

Browse files
committed
style(benchmarks): apply black formatting for truthfulqa PR
1 parent 1abc780 commit 8a96eb9

6 files changed

Lines changed: 29 additions & 40 deletions

File tree

cascadeflow/agent.py

Lines changed: 7 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1725,19 +1725,11 @@ async def _execute_direct_with_timing(
17251725
prompt_tokens = response.metadata.get("prompt_tokens")
17261726
completion_tokens = response.metadata.get("completion_tokens")
17271727
total_tokens = response.metadata.get("total_tokens")
1728-
if (
1729-
total_tokens is None
1730-
and prompt_tokens is not None
1731-
and completion_tokens is not None
1732-
):
1728+
if total_tokens is None and prompt_tokens is not None and completion_tokens is not None:
17331729
total_tokens = prompt_tokens + completion_tokens
17341730

17351731
cost = None
1736-
if (
1737-
LITELLM_AVAILABLE
1738-
and prompt_tokens is not None
1739-
and completion_tokens is not None
1740-
):
1732+
if LITELLM_AVAILABLE and prompt_tokens is not None and completion_tokens is not None:
17411733
try:
17421734
provider = LiteLLMCostProvider()
17431735
cost = provider.calculate_cost(
@@ -2222,7 +2214,11 @@ def __init__(
22222214
self.speedup = 1.0
22232215

22242216
token_total = total_tokens
2225-
if token_total is None and prompt_tokens is not None and completion_tokens is not None:
2217+
if (
2218+
token_total is None
2219+
and prompt_tokens is not None
2220+
and completion_tokens is not None
2221+
):
22262222
token_total = prompt_tokens + completion_tokens
22272223
if token_total is None:
22282224
token_total = int(len(content.split()) * 1.3)

tests/benchmarks/base.py

Lines changed: 14 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -429,7 +429,9 @@ def _generate_summary(self) -> BenchmarkSummary:
429429

430430
# Cascade metrics
431431
direct_routed = sum(1 for r in valid_results if r.direct_routed)
432-
drafter_accepted = sum(1 for r in valid_results if r.routing_strategy == "cascade" and r.accepted)
432+
drafter_accepted = sum(
433+
1 for r in valid_results if r.routing_strategy == "cascade" and r.accepted
434+
)
433435
escalated = sum(1 for r in valid_results if r.verifier_rejected)
434436
cascade_total = drafter_accepted + escalated
435437

@@ -447,14 +449,18 @@ def _generate_summary(self) -> BenchmarkSummary:
447449
p95_latency = latencies[p95_idx]
448450
cascadeflow_latencies = [r.cascadeflow_latency_ms for r in valid_results]
449451
avg_cascadeflow_latency = (
450-
sum(cascadeflow_latencies) / len(cascadeflow_latencies) if cascadeflow_latencies else 0.0
452+
sum(cascadeflow_latencies) / len(cascadeflow_latencies)
453+
if cascadeflow_latencies
454+
else 0.0
451455
)
452456

453457
# Quality metrics
454458
correct = sum(1 for r in valid_results if r.is_correct)
455459
accuracy = (correct / len(valid_results) * 100) if valid_results else 0.0
456460

457-
drafter_results = [r for r in valid_results if r.routing_strategy == "cascade" and r.accepted]
461+
drafter_results = [
462+
r for r in valid_results if r.routing_strategy == "cascade" and r.accepted
463+
]
458464
drafter_correct = sum(1 for r in drafter_results if r.is_correct)
459465
drafter_accuracy = (
460466
(drafter_correct / len(drafter_results) * 100) if drafter_results else 0.0
@@ -468,9 +474,7 @@ def _generate_summary(self) -> BenchmarkSummary:
468474

469475
direct_results = [r for r in valid_results if r.direct_routed]
470476
direct_correct = sum(1 for r in direct_results if r.is_correct)
471-
direct_accuracy = (
472-
(direct_correct / len(direct_results) * 100) if direct_results else 0.0
473-
)
477+
direct_accuracy = (direct_correct / len(direct_results) * 100) if direct_results else 0.0
474478

475479
# Token usage
476480
total_input = sum(r.tokens_input for r in valid_results)
@@ -483,7 +487,9 @@ def _generate_summary(self) -> BenchmarkSummary:
483487
failed_tests=failed,
484488
drafter_accepted=drafter_accepted,
485489
escalated_to_verifier=escalated,
486-
acceptance_rate_pct=(drafter_accepted / cascade_total * 100) if cascade_total > 0 else 0.0,
490+
acceptance_rate_pct=(
491+
(drafter_accepted / cascade_total * 100) if cascade_total > 0 else 0.0
492+
),
487493
escalation_rate_pct=(escalated / cascade_total * 100) if cascade_total > 0 else 0.0,
488494
direct_routed=direct_routed,
489495
direct_routing_pct=(direct_routed / successful * 100) if successful > 0 else 0.0,
@@ -524,9 +530,7 @@ def _print_summary(self, summary: BenchmarkSummary) -> None:
524530
print(
525531
f" Escalated: {summary.escalated_to_verifier} ({summary.escalation_rate_pct:.1f}%)"
526532
)
527-
print(
528-
f" Direct Routed: {summary.direct_routed} ({summary.direct_routing_pct:.1f}%)"
529-
)
533+
print(f" Direct Routed: {summary.direct_routed} ({summary.direct_routing_pct:.1f}%)")
530534

531535
print("\nCOST ANALYSIS:")
532536
print(f" Total Cost: ${summary.total_cost:.6f}")

tests/benchmarks/bfcl/bfcl_full_benchmark.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -394,7 +394,9 @@ def _check_function_correct(
394394
counts[func_key] = counts.get(func_key, 0) + 1
395395

396396
for func, expected_count in counts.items():
397-
func_mentions = response_lower.count(func) + response_lower.count(func.replace("_", " "))
397+
func_mentions = response_lower.count(func) + response_lower.count(
398+
func.replace("_", " ")
399+
)
398400
if func_mentions < expected_count:
399401
return False, True
400402
return True, True
@@ -404,9 +406,7 @@ def _check_function_correct(
404406
if expected_func is None:
405407
# No tool should be used
406408
func_correct = (
407-
found_func is None
408-
or "don't need" in response_lower
409-
or "no tool" in response_lower
409+
found_func is None or "don't need" in response_lower or "no tool" in response_lower
410410
)
411411
return func_correct, True
412412

tests/benchmarks/metrics.py

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -145,9 +145,7 @@ def from_results(results: list[BenchmarkResult]) -> "QualityMetrics":
145145
direct_results = [r for r in results if r.direct_routed]
146146
direct_correct = sum(1 for r in direct_results if r.is_correct)
147147
direct_incorrect = len(direct_results) - direct_correct
148-
direct_accuracy = (
149-
(direct_correct / len(direct_results) * 100) if direct_results else 0.0
150-
)
148+
direct_accuracy = (direct_correct / len(direct_results) * 100) if direct_results else 0.0
151149

152150
return QualityMetrics(
153151
overall_accuracy=overall_accuracy,

tests/benchmarks/mmlu/mmlu_full_benchmark.py

Lines changed: 2 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -681,9 +681,7 @@ async def process_problem(index: int, problem: dict) -> tuple[int, BenchmarkResu
681681
)
682682
return index, benchmark_result, output
683683
except Exception as e:
684-
output = (
685-
f"[{index+1}/{len(problems)}] {problem['subject'][:20]:20s}: ⚠️ ERROR: {e}"
686-
)
684+
output = f"[{index+1}/{len(problems)}] {problem['subject'][:20]:20s}: ⚠️ ERROR: {e}"
687685
benchmark_result = BenchmarkResult(
688686
query=problem["question"][:50] + "...",
689687
subject=problem["subject"],
@@ -707,10 +705,7 @@ async def process_problem(index: int, problem: dict) -> tuple[int, BenchmarkResu
707705
for start in range(0, len(problems), concurrency):
708706
batch = problems[start : start + concurrency]
709707
batch_results = await asyncio.gather(
710-
*(
711-
process_problem(start + offset, problem)
712-
for offset, problem in enumerate(batch)
713-
)
708+
*(process_problem(start + offset, problem) for offset, problem in enumerate(batch))
714709
)
715710
for index, benchmark_result, output in batch_results:
716711
results[index] = benchmark_result

tests/benchmarks/tool_calls_realworld.py

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -480,11 +480,7 @@ async def run_cascade(self, query: Any) -> dict[str, Any]:
480480
"cascadeflow_latency_ms": (
481481
(result.complexity_detection_ms or 0)
482482
+ (result.metadata.get("domain_detection_ms", 0) if result.metadata else 0)
483-
+ (
484-
result.metadata.get("tool_complexity_analysis_ms", 0)
485-
if result.metadata
486-
else 0
487-
)
483+
+ (result.metadata.get("tool_complexity_analysis_ms", 0) if result.metadata else 0)
488484
+ (result.quality_verification_ms or 0)
489485
),
490486
"tokens_input": result.metadata.get("prompt_tokens", 0),

0 commit comments

Comments
 (0)