Buckets:
tahamajs/Sysmem2_in_AI / ComputerAssignments /CA6_systematic_generalization /scripts /utility_scripts.py
| 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.