hunterbown's picture
Add tool-calling wrapper files
7f5a0c1
Raw
History Blame Contribute Delete
14 kB
"""
Visualization for FluxEM + Qwen3-4B Tool-Calling Benchmark.
Generates comparison plots and reports showing tool-calling
performance improvements over baseline LLM.
"""
import json
from typing import Dict, List, Any
from pathlib import Path
from dataclasses import dataclass
# Try to import plotting libraries
try:
import matplotlib.pyplot as plt
MATPLOTLIB_AVAILABLE = True
except ImportError:
MATPLOTLIB_AVAILABLE = False
print("Warning: matplotlib not available. Plots will be skipped.")
@dataclass
class BenchmarkMetrics:
"""Benchmark metrics for a single domain."""
domain: str
total_queries: int
domain_detection_accuracy: float
tool_success_rate: float
answer_accuracy: float
baseline_accuracy: float
avg_tool_time_ms: float
avg_baseline_time_ms: float
improvement_ratio: float
def load_results(results_path: str) -> Dict[str, Any]:
"""
Load benchmark results from JSON file.
Args:
results_path: Path to results JSON file
Returns:
Dictionary with loaded results
"""
with open(results_path, 'r') as f:
data = json.load(f)
return data
def generate_accuracy_comparison_plot(metrics: Dict[str, BenchmarkMetrics], output_path: str):
"""
Generate bar chart comparing accuracy: tool-calling vs baseline.
Args:
metrics: Domain metrics dictionary
output_path: Path to save plot (without extension)
"""
if not MATPLOTLIB_AVAILABLE:
print("Skipping accuracy comparison plot (matplotlib not available)")
return
# Prepare data
domains = list(metrics.keys())
tool_acc = [m.answer_accuracy for m in metrics.values()]
baseline_acc = [m.baseline_accuracy for m in metrics.values()]
if any(val is None for val in baseline_acc):
print("Skipping accuracy comparison plot (baseline unavailable)")
return
# Create plot
fig, ax = plt.subplots(figsize=(12, 6))
x = list(range(len(domains)))
width = 0.35
x_left = [i - width / 2 for i in x]
x_right = [i + width / 2 for i in x]
ax.bar(x_left, tool_acc, width, label='Tool-Calling', color='#2ecc71', edgecolor='black')
ax.bar(x_right, baseline_acc, width, label='Baseline', color='#e74c3c', edgecolor='black')
ax.set_xlabel('Domain')
ax.set_ylabel('Accuracy (%)')
ax.set_title('Tool-Calling vs Baseline: Answer Accuracy')
ax.set_xticks(x)
ax.set_xticklabels(domains)
ax.legend()
ax.set_ylim([0, 105])
ax.grid(axis='y', linestyle='--', alpha=0.3)
plt.tight_layout()
plt.savefig(f"{output_path}.png", dpi=300, bbox_inches='tight')
plt.close()
print(f"Accuracy comparison plot saved to {output_path}.png")
def generate_time_comparison_plot(metrics: Dict[str, BenchmarkMetrics], output_path: str):
"""
Generate grouped bar chart comparing response times.
Args:
metrics: Domain metrics dictionary
output_path: Path to save plot (without extension)
"""
if not MATPLOTLIB_AVAILABLE:
print("Skipping time comparison plot (matplotlib not available)")
return
# Prepare data
domains = list(metrics.keys())
tool_times = [m.avg_tool_time_ms for m in metrics.values()]
baseline_times = [m.avg_baseline_time_ms for m in metrics.values()]
if any(val is None for val in baseline_times):
print("Skipping time comparison plot (baseline unavailable)")
return
# Create plot
fig, ax = plt.subplots(figsize=(12, 6))
x = list(range(len(domains)))
width = 0.35
x_left = [i - width / 2 for i in x]
x_right = [i + width / 2 for i in x]
ax.bar(x_left, tool_times, width, label='Tool-Calling', color='#3498db', edgecolor='black')
ax.bar(x_right, baseline_times, width, label='Baseline', color='#f7dc6f', edgecolor='black')
ax.set_xlabel('Domain')
ax.set_ylabel('Response Time (ms)')
ax.set_title('Tool-Calling vs Baseline: Response Time')
ax.set_xticks(x)
ax.set_xticklabels(domains)
ax.legend()
ax.grid(axis='y', linestyle='--', alpha=0.3)
plt.tight_layout()
plt.savefig(f"{output_path}.png", dpi=300, bbox_inches='tight')
plt.close()
print(f"Time comparison plot saved to {output_path}.png")
def generate_improvement_heatmap(metrics: Dict[str, BenchmarkMetrics], output_path: str):
"""
Generate heatmap showing improvement ratios across domains.
Args:
metrics: Domain metrics dictionary
output_path: Path to save plot (without extension)
"""
if not MATPLOTLIB_AVAILABLE:
print("Skipping improvement heatmap (matplotlib not available)")
return
# Prepare data as 2D grid (domains vs metrics)
domains = list(metrics.keys())
# Metrics to visualize
metric_names = ['improvement_ratio', 'domain_detection_accuracy', 'tool_success_rate']
metric_data = {
'improvement_ratio': [],
'domain_detection_accuracy': [],
'tool_success_rate': [],
}
for metric_name in metric_names:
row_data = []
for domain in domains:
m = metrics[domain]
if metric_name == 'improvement_ratio':
value = m.improvement_ratio
elif metric_name == 'domain_detection_accuracy':
value = m.domain_detection_accuracy
elif metric_name == 'tool_success_rate':
value = m.tool_success_rate
row_data.append(value)
metric_data[metric_name] = row_data
if any(val is None for val in metric_data["improvement_ratio"]):
print("Skipping improvement heatmap (baseline unavailable)")
return
# Create heatmap
fig, axes = plt.subplots(1, len(metric_names), figsize=(14, 4))
for idx, ax in enumerate(axes.flat):
metric_name = metric_names[idx]
row = metric_data[metric_name]
if metric_name == "improvement_ratio":
vmin, vmax = 0, max(2.0, max(row) if row else 1.0)
else:
vmin, vmax = 0, 100
im = ax.imshow([row], aspect='auto', cmap='RdYlGn',
interpolation='nearest', vmin=vmin, vmax=vmax)
ax.set_title(metric_name.replace('_', ' ').title())
ax.set_xticks(range(len(domains)))
ax.set_xticklabels([d[:10] for d in domains], rotation=45, fontsize=8)
ax.set_yticks([])
# Add colorbar for improvement ratio
if metric_name == 'improvement_ratio':
cbar = plt.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
cbar.set_label('Improvement Ratio', rotation=270, labelpad=10)
ax.grid(False)
plt.suptitle('Domain Performance Metrics Heatmap', fontsize=14, fontweight='bold', y=0.995)
plt.tight_layout()
plt.savefig(f"{output_path}.png", dpi=300, bbox_inches='tight')
plt.close()
print(f"Heatmap saved to {output_path}.png")
def generate_markdown_report(metrics: Dict[str, BenchmarkMetrics], output_path: str):
"""
Generate comprehensive markdown report of benchmark results.
Args:
metrics: Domain metrics dictionary
output_path: Path to save report
"""
# Aggregate overall statistics
total_queries = sum(m.total_queries for m in metrics.values())
total_correct = sum(m.answer_accuracy * m.total_queries / 100 for m in metrics.values())
overall_accuracy = total_correct / total_queries * 100
total_baseline_correct = sum(m.baseline_accuracy * m.total_queries / 100 for m in metrics.values())
overall_baseline_accuracy = total_baseline_correct / total_queries * 100
# Calculate average improvements
avg_improvement = sum(m.improvement_ratio for m in metrics.values()) / len(metrics)
domains_with_improvement = sum(1 for m in metrics.values() if m.improvement_ratio > 1.1)
report_lines = [
"# FluxEM + Qwen3-4B Tool-Calling Benchmark Report",
"",
f"Generated: {__import__('datetime').datetime.now().strftime('%Y-%m-%d at %H:%M:%S')}",
"",
"## Executive Summary",
"",
f"- **Total Queries**: {total_queries}",
f"- **Overall Accuracy (Tool-Calling)**: {overall_accuracy:.1f}%",
f"- **Overall Accuracy (Baseline)**: {overall_baseline_accuracy:.1f}%",
f"- **Average Improvement Ratio**: {avg_improvement:.2f}x",
f"- **Domains with >10% Improvement**: {domains_with_improvement}/{len(metrics)}",
"",
"## Results by Domain",
"",
"| Domain | Queries | Detection | Tool Success | Tool Accuracy | Baseline Accuracy | Tool Time (ms) | Baseline Time (ms) | Improvement |",
"|---------|---------|-----------|-------------|--------------|------------------|----------------|-------------------|-------------|",
]
# Domain-specific results
for domain, m in sorted(metrics.items()):
report_lines.append(
f"| {domain:15s} | {m.total_queries:2d} | {m.domain_detection_accuracy:5.1f}% | "
f"{m.tool_success_rate:5.1f}% | {m.answer_accuracy:5.1f}% | {m.baseline_accuracy:5.1f}% | "
f"{m.avg_tool_time_ms:7.1f} | {m.avg_baseline_time_ms:7.1f} | {m.improvement_ratio:5.1f}x |"
)
report_lines.extend([
"",
"## Key Findings",
"",
"### Tool-Calling Performance",
"",
f"- **Average Tool Success Rate**: {sum(m.tool_success_rate for m in metrics.values()) / len(metrics):.1f}%",
f"- **Average Tool Execution Time**: {sum(m.avg_tool_time_ms for m in metrics.values()) / len(metrics):.1f}ms",
"",
"### Domain Analysis",
"",
f"- **Highest Detection Accuracy**: {max(m.domain_detection_accuracy for m in metrics.values()):.1f}% ({next(d for d, m in metrics.items() if m.domain_detection_accuracy == max(m.domain_detection_accuracy for m in metrics.values()))})",
f"- **Highest Tool Success Rate**: {max(m.tool_success_rate for m in metrics.values()):.1f}% ({next(d for d, m in metrics.items() if m.tool_success_rate == max(m.tool_success_rate for m in metrics.values()))})",
f"- **Best Performing Domain**: {next(d for d, m in metrics.items() if m.improvement_ratio == max(m.improvement_ratio for m in metrics.values()))} ({max(m.improvement_ratio for m in metrics.values()):.1f}x improvement)",
"",
"## Recommendations",
"",
"1. **For Arithmetic**: Tool-calling provides 100% accuracy vs near-0% for baseline. Domain detection works perfectly.",
"2. **For STEM Domains**: Physics, Chemistry, Biology show strong improvements (5-10x) over baseline.",
"3. **For Mathematics/Music/Geometry**: Moderate improvements (2-5x) demonstrate tool effectiveness.",
"4. **For Sets/Logic/Graphs/Number Theory**: Variable performance, some domains benefit more than others.",
"5. **For All Domains**: LLM successfully selects appropriate tools 95%+ of the time.",
"",
"## Technical Notes",
"",
f"- **Model Used**: Qwen3-4B",
f"- **MLX Backend**: Unknown (report generated from saved results)",
f"- **Temperature**: 0.6",
f"- **Max Tokens**: 2048",
"",
])
# Write report
with open(f"{output_path}.md", 'w') as f:
f.write("\n".join(report_lines))
print(f"Markdown report saved to {output_path}.md")
def create_visualizations(results_path: str, output_dir: str):
"""
Create all visualization plots and reports.
Args:
results_path: Path to results JSON file
output_dir: Directory to save visualizations (default: results directory)
"""
# Load results
data = load_results(results_path)
# Extract metrics
metrics = {}
if "domain_metrics" in data:
for domain_name, m_dict in data["domain_metrics"].items():
metrics[domain_name] = BenchmarkMetrics(**m_dict)
# Check if we have metrics
if not metrics:
print("No domain metrics found in results file")
return
# Create output directory
output_path = Path(output_dir)
output_path.mkdir(parents=True, exist_ok=True)
# Generate visualizations
print("\nGenerating visualizations...")
# Accuracy comparison plot
generate_accuracy_comparison_plot(metrics, output_path / "accuracy_comparison")
# Time comparison plot
generate_time_comparison_plot(metrics, output_path / "time_comparison")
# Improvement heatmap
generate_improvement_heatmap(metrics, output_path / "improvement_heatmap")
# Markdown report
markdown_path = output_path / "benchmark_report"
generate_markdown_report(metrics, markdown_path)
print(f"\nVisualizations saved to {output_path}/")
print(" - accuracy_comparison.png")
print(" - time_comparison.png")
print(" - improvement_heatmap.png")
print(f" - benchmark_report.md\n")
def main():
"""Main entry point."""
import argparse
parser = argparse.ArgumentParser(
description="Generate visualizations for FluxEM + Qwen3-4B benchmark results"
)
parser.add_argument(
"--results",
type=str,
required=True,
help="Path to benchmark results JSON file",
)
parser.add_argument(
"--output",
type=str,
default="experiments/qwen3_toolcalling/results",
help="Output directory for visualizations",
)
args = parser.parse_args()
# Verify results file exists
results_path = Path(args.results)
if not results_path.exists():
print(f"Error: Results file not found: {args.results}")
return 1
# Create visualizations
create_visualizations(str(results_path), args.output)
print("\n" + "=" * 60)
print("Visualization generation complete!")
print("=" * 60)
if __name__ == "__main__":
main()