tahamajs's picture
download
raw
14.2 kB
import torch
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
from typing import Dict, List, Any, Optional
import json
import os
from pathlib import Path
import argparse
from rich.console import Console
from rich.table import Table
from rich.progress import Progress
from src.core.components import SystematicGeneralizationTask
from src.datasets.dataset_implementations import (
SCANDataset,
ArithmeticReasoningDataset,
VisualReasoningDataset,
)
from src.evaluation.evaluation_framework import SystematicGeneralizationEvaluator
def analyze_dataset_statistics(dataset_name: str, output_dir: str = "analysis"):
console = Console()
console.print(f"[blue]Analyzing {dataset_name} dataset statistics...[/blue]")
Path(output_dir).mkdir(parents=True, exist_ok=True)
if dataset_name.lower() == "scan":
dataset = SCANDataset()
elif dataset_name.lower() == "arithmetic":
dataset = ArithmeticReasoningDataset()
elif dataset_name.lower() == "visual":
dataset = VisualReasoningDataset()
else:
console.print(f"[red]Unknown dataset: {dataset_name}[/red]")
return
stats = dataset.get_statistics()
fig, axes = plt.subplots(2, 2, figsize=(15, 10))
fig.suptitle(f"{dataset_name} Dataset Statistics", fontsize=16)
if "complexity_distribution" in stats:
complexities = list(stats["complexity_distribution"].keys())
counts = list(stats["complexity_distribution"].values())
axes[0, 0].bar(complexities, counts)
axes[0, 0].set_xlabel("Complexity Level")
axes[0, 0].set_ylabel("Count")
axes[0, 0].set_title("Complexity Distribution")
axes[0, 0].grid(True, alpha=0.3)
lengths = []
for i in range(min(100, len(dataset))):
item = dataset[i]
if "command" in item:
length = (item["command"] != 0).sum().item()
elif "expression" in item:
length = (item["expression"] != 0).sum().item()
elif "description" in item:
length = (item["description"] != 0).sum().item()
else:
length = 0
lengths.append(length)
axes[0, 1].hist(lengths, bins=20, alpha=0.7)
axes[0, 1].set_xlabel("Sequence Length")
axes[0, 1].set_ylabel("Frequency")
axes[0, 1].set_title("Sequence Length Distribution")
axes[0, 1].grid(True, alpha=0.3)
if "compositional_examples" in stats:
comp_count = stats["compositional_examples"]
non_comp_count = stats["non_compositional_examples"]
axes[1, 0].pie(
[comp_count, non_comp_count],
labels=["Compositional", "Non-Compositional"],
autopct="%1.1f%%",
)
axes[1, 0].set_title("Compositional vs Non-Compositional")
if "primitive_counts" in stats:
primitives = list(stats["primitive_counts"].keys())[:10]
counts = list(stats["primitive_counts"].values())[:10]
axes[1, 1].bar(range(len(primitives)), counts)
axes[1, 1].set_xticks(range(len(primitives)))
axes[1, 1].set_xticklabels(primitives, rotation=45)
axes[1, 1].set_ylabel("Count")
axes[1, 1].set_title("Top 10 Primitives")
axes[1, 1].grid(True, alpha=0.3)
plt.tight_layout()
plt.savefig(
f"{output_dir}/{dataset_name}_statistics.png", dpi=300, bbox_inches="tight"
)
plt.show()
with open(f"{output_dir}/{dataset_name}_stats.json", "w") as f:
json.dump(stats, f, indent=2)
console.print(f"[green]Analysis completed! Results saved to {output_dir}/[/green]")
def compare_experiment_results(results_dir: str = "results"):
console = Console()
console.print("[blue]Comparing experiment results...[/blue]")
results_path = Path(results_dir) / "experiments"
result_files = list(results_path.glob("*_evaluation.json"))
if not result_files:
console.print("[red]No evaluation results found![/red]")
return
console.print(f"Found {len(result_files)} result files")
all_results = {}
for result_file in result_files:
experiment_name = result_file.stem.replace("_evaluation", "")
evaluator = SystematicGeneralizationEvaluator(results_dir)
results = evaluator.load_results(result_file.name)
all_results[experiment_name] = results
table = Table(title="Experiment Comparison")
table.add_column("Experiment", style="cyan")
table.add_column("Best Model", style="green")
table.add_column("Best Accuracy", justify="right")
table.add_column("Best Systematic Gap", justify="right")
table.add_column("Models Tested", justify="right")
for exp_name, exp_results in all_results.items():
best_model = None
best_score = -1
for model_name, model_results in exp_results.items():
avg_accuracy = np.mean([m.accuracy for m in model_results.values()])
avg_gap = np.mean(
[m.systematic_generalization_gap for m in model_results.values()]
)
score = avg_accuracy - avg_gap
if score > best_score:
best_score = score
best_model = model_name
if best_model:
best_model_results = exp_results[best_model]
best_accuracy = np.mean([m.accuracy for m in best_model_results.values()])
best_gap = np.mean(
[m.systematic_generalization_gap for m in best_model_results.values()]
)
table.add_row(
exp_name,
best_model,
f"{best_accuracy:.4f}",
f"{best_gap:.4f}",
str(len(exp_results)),
)
console.print(table)
fig, axes = plt.subplots(1, 2, figsize=(15, 6))
fig.suptitle("Experiment Comparison", fontsize=16)
experiments = list(all_results.keys())
best_accuracies = []
for exp_name in experiments:
exp_results = all_results[exp_name]
best_acc = 0
for model_results in exp_results.values():
avg_acc = np.mean([m.accuracy for m in model_results.values()])
best_acc = max(best_acc, avg_acc)
best_accuracies.append(best_acc)
axes[0].bar(experiments, best_accuracies)
axes[0].set_ylabel("Best Accuracy")
axes[0].set_title("Best Accuracy per Experiment")
axes[0].tick_params(axis="x", rotation=45)
axes[0].grid(True, alpha=0.3)
best_gaps = []
for exp_name in experiments:
exp_results = all_results[exp_name]
best_gap = float("inf")
for model_results in exp_results.values():
avg_gap = np.mean(
[m.systematic_generalization_gap for m in model_results.values()]
)
best_gap = min(best_gap, avg_gap)
best_gaps.append(best_gap)
axes[1].bar(experiments, best_gaps)
axes[1].set_ylabel("Best Systematic Gap")
axes[1].set_title("Best Systematic Gap per Experiment")
axes[1].tick_params(axis="x", rotation=45)
axes[1].grid(True, alpha=0.3)
plt.tight_layout()
plt.savefig(
f"{results_dir}/experiment_comparison.png", dpi=300, bbox_inches="tight"
)
plt.show()
console.print("[green]Comparison completed![/green]")
def generate_model_report(model_name: str, results_dir: str = "results"):
console = Console()
console.print(f"[blue]Generating report for {model_name}...[/blue]")
evaluator = SystematicGeneralizationEvaluator(results_dir)
result_files = list(Path(results_dir).glob("experiments/*_evaluation.json"))
if not result_files:
console.print("[red]No evaluation results found![/red]")
return
latest_results_file = max(result_files, key=lambda x: x.stat().st_mtime)
all_results = evaluator.load_results(latest_results_file.name)
if model_name not in all_results:
console.print(f"[red]Model {model_name} not found in results![/red]")
console.print(f"Available models: {list(all_results.keys())}")
return
model_results = all_results[model_name]
report_lines = []
report_lines.append(f"# Model Report: {model_name}")
report_lines.append(
f"Generated at: {pd.Timestamp.now().strftime('%Y-%m-%d %H:%M:%S')}"
)
report_lines.append("")
all_accuracies = [m.accuracy for m in model_results.values()]
all_gaps = [m.systematic_generalization_gap for m in model_results.values()]
all_comp_accs = [m.compositional_accuracy for m in model_results.values()]
report_lines.append("## Overall Statistics")
report_lines.append("")
report_lines.append(
f"- **Average Accuracy**: {np.mean(all_accuracies):.4f} ± {np.std(all_accuracies):.4f}"
)
report_lines.append(
f"- **Average Systematic Gap**: {np.mean(all_gaps):.4f} ± {np.std(all_gaps):.4f}"
)
report_lines.append(
f"- **Average Compositional Accuracy**: {np.mean(all_comp_accs):.4f} ± {np.std(all_comp_accs):.4f}"
)
report_lines.append("")
report_lines.append("## Per-Dataset Performance")
report_lines.append("")
report_lines.append(
"| Dataset | Accuracy | Systematic Gap | Compositional Acc | Inference Time |"
)
report_lines.append(
"|---------|----------|----------------|-------------------|----------------|"
)
for dataset_name, metrics in model_results.items():
report_lines.append(
f"| {dataset_name} | {metrics.accuracy:.4f} | "
f"{metrics.systematic_generalization_gap:.4f} | "
f"{metrics.compositional_accuracy:.4f} | "
f"{metrics.inference_time:.4f}s |"
)
report_lines.append("")
best_dataset = max(model_results.items(), key=lambda x: x[1].accuracy)
worst_dataset = min(model_results.items(), key=lambda x: x[1].accuracy)
report_lines.append("## Performance Analysis")
report_lines.append("")
report_lines.append(
f"- **Best Performance**: {best_dataset[0]} (Accuracy: {best_dataset[1].accuracy:.4f})"
)
report_lines.append(
f"- **Worst Performance**: {worst_dataset[0]} (Accuracy: {worst_dataset[1].accuracy:.4f})"
)
report_lines.append("")
report_lines.append("## Recommendations")
report_lines.append("")
if np.mean(all_gaps) > 0.1:
report_lines.append(
"- **High Systematic Gap**: Consider using more compositional training data or architectural improvements."
)
if np.mean(all_comp_accs) < 0.7:
report_lines.append(
"- **Low Compositional Accuracy**: Focus on improving compositional reasoning capabilities."
)
if np.mean(all_accuracies) > 0.8:
report_lines.append(
"- **Good Overall Performance**: Model shows strong performance across datasets."
)
report_lines.append("")
report_text = "\n".join(report_lines)
report_path = Path(results_dir) / "experiments" / f"{model_name}_report.md"
with open(report_path, "w") as f:
f.write(report_text)
console.print(f"[green]Report saved to {report_path}[/green]")
def cleanup_results(results_dir: str = "results", keep_latest: int = 5):
console = Console()
console.print(f"[blue]Cleaning up results in {results_dir}...[/blue]")
results_path = Path(results_dir)
exp_files = list(results_path.glob("experiments/*.json"))
exp_files.sort(key=lambda x: x.stat().st_mtime, reverse=True)
files_to_remove = exp_files[keep_latest:]
for file_path in files_to_remove:
file_path.unlink()
console.print(f"[yellow]Removed {file_path}[/yellow]")
model_files = list(results_path.glob("models/*.pt"))
model_files.sort(key=lambda x: x.stat().st_mtime, reverse=True)
files_to_remove = model_files[keep_latest:]
for file_path in files_to_remove:
file_path.unlink()
console.print(f"[yellow]Removed {file_path}[/yellow]")
log_files = list(results_path.glob("logs/*.log"))
log_files.sort(key=lambda x: x.stat().st_mtime, reverse=True)
files_to_remove = log_files[keep_latest:]
for file_path in files_to_remove:
file_path.unlink()
console.print(f"[yellow]Removed {file_path}[/yellow]")
console.print("[green]Cleanup completed![/green]")
def main():
parser = argparse.ArgumentParser(
description="Systematic Generalization Utility Scripts"
)
subparsers = parser.add_subparsers(dest="command", help="Available commands")
analyze_parser = subparsers.add_parser(
"analyze-dataset", help="Analyze dataset statistics"
)
analyze_parser.add_argument(
"dataset", choices=["scan", "arithmetic", "visual"], help="Dataset to analyze"
)
analyze_parser.add_argument(
"--output-dir", default="analysis", help="Output directory for analysis results"
)
compare_parser = subparsers.add_parser(
"compare-experiments", help="Compare experiment results"
)
compare_parser.add_argument(
"--results-dir", default="results", help="Results directory"
)
report_parser = subparsers.add_parser("model-report", help="Generate model report")
report_parser.add_argument("model", help="Model name to generate report for")
report_parser.add_argument(
"--results-dir", default="results", help="Results directory"
)
cleanup_parser = subparsers.add_parser("cleanup", help="Clean up old results")
cleanup_parser.add_argument(
"--results-dir", default="results", help="Results directory"
)
cleanup_parser.add_argument(
"--keep-latest", type=int, default=5, help="Number of latest files to keep"
)
args = parser.parse_args()
if args.command == "analyze-dataset":
analyze_dataset_statistics(args.dataset, args.output_dir)
elif args.command == "compare-experiments":
compare_experiment_results(args.results_dir)
elif args.command == "model-report":
generate_model_report(args.model, args.results_dir)
elif args.command == "cleanup":
cleanup_results(args.results_dir, args.keep_latest)
else:
parser.print_help()
if __name__ == "__main__":
main()

Xet Storage Details

Size:
14.2 kB
·
Xet hash:
050c92d77121025a8fbbd3890ead997fa7dfda9f7fdbabeb7151732f01772990

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.