Buckets:
tahamajs/Sysmem2_in_AI / ComputerAssignments /CA6_systematic_generalization /src /visualization /visualization_tools.py
| import torch | |
| import torch.nn as nn | |
| import numpy as np | |
| import matplotlib.pyplot as plt | |
| import seaborn as sns | |
| import pandas as pd | |
| from typing import List, Dict, Tuple, Any, Optional, Union | |
| import plotly.graph_objects as go | |
| import plotly.express as px | |
| from plotly.subplots import make_subplots | |
| import networkx as nx | |
| from sklearn.manifold import TSNE | |
| from sklearn.decomposition import PCA | |
| import json | |
| import os | |
| from pathlib import Path | |
| from rich.console import Console | |
| from rich.table import Table | |
| from rich.panel import Panel | |
| from rich.progress import Progress | |
| from ..core.components import SystematicGeneralizationTask, CompositeExpression | |
| from ..evaluation.evaluation_framework import EvaluationMetrics | |
| class SystematicGeneralizationVisualizer: | |
| def __init__(self, results_dir: str = "results"): | |
| self.results_dir = results_dir | |
| self.console = Console() | |
| plt.style.use("seaborn-v0_8") | |
| sns.set_palette("husl") | |
| self.plots_dir = Path(results_dir) / "plots" | |
| self.plots_dir.mkdir(parents=True, exist_ok=True) | |
| def plot_model_comparison( | |
| self, results: Dict[str, Dict[str, EvaluationMetrics]], save_plots: bool = True | |
| ) -> None: | |
| model_names = list(results.keys()) | |
| datasets = list(results[model_names[0]].keys()) | |
| fig, axes = plt.subplots(2, 3, figsize=(18, 12)) | |
| fig.suptitle( | |
| "Comprehensive Model Comparison for Systematic Generalization", | |
| fontsize=16, | |
| fontweight="bold", | |
| ) | |
| self._plot_accuracy_comparison(axes[0, 0], results, model_names, datasets) | |
| self._plot_systematic_gap(axes[0, 1], results, model_names, datasets) | |
| self._plot_compositional_accuracy(axes[0, 2], results, model_names, datasets) | |
| self._plot_inference_time(axes[1, 0], results, model_names, datasets) | |
| self._plot_length_generalization(axes[1, 1], results, model_names, datasets) | |
| self._plot_complexity_generalization(axes[1, 2], results, model_names, datasets) | |
| plt.tight_layout() | |
| if save_plots: | |
| plt.savefig( | |
| self.plots_dir / "comprehensive_model_comparison.png", | |
| dpi=300, | |
| bbox_inches="tight", | |
| ) | |
| plt.show() | |
| def _plot_accuracy_comparison( | |
| self, | |
| ax, | |
| results: Dict[str, Dict[str, EvaluationMetrics]], | |
| model_names: List[str], | |
| datasets: List[str], | |
| ): | |
| x = np.arange(len(datasets)) | |
| width = 0.8 / len(model_names) | |
| for i, model_name in enumerate(model_names): | |
| accuracies = [results[model_name][dataset].accuracy for dataset in datasets] | |
| ax.bar(x + i * width, accuracies, width, label=model_name, alpha=0.8) | |
| ax.set_xlabel("Datasets") | |
| ax.set_ylabel("Accuracy") | |
| ax.set_title("Model Accuracy Comparison") | |
| ax.set_xticks(x + width * (len(model_names) - 1) / 2) | |
| ax.set_xticklabels(datasets, rotation=45) | |
| ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left") | |
| ax.grid(True, alpha=0.3) | |
| ax.set_ylim(0, 1) | |
| def _plot_systematic_gap( | |
| self, | |
| ax, | |
| results: Dict[str, Dict[str, EvaluationMetrics]], | |
| model_names: List[str], | |
| datasets: List[str], | |
| ): | |
| x = np.arange(len(datasets)) | |
| width = 0.8 / len(model_names) | |
| colors = [ | |
| "red" if gap > 0.1 else "orange" if gap > 0.05 else "green" | |
| for model_name in model_names | |
| for gap in [ | |
| results[model_name][dataset].systematic_generalization_gap | |
| for dataset in datasets | |
| ] | |
| ] | |
| for i, model_name in enumerate(model_names): | |
| gaps = [ | |
| results[model_name][dataset].systematic_generalization_gap | |
| for dataset in datasets | |
| ] | |
| bars = ax.bar(x + i * width, gaps, width, label=model_name, alpha=0.8) | |
| for j, bar in enumerate(bars): | |
| gap = gaps[j] | |
| if gap > 0.1: | |
| bar.set_color("red") | |
| elif gap > 0.05: | |
| bar.set_color("orange") | |
| else: | |
| bar.set_color("green") | |
| ax.set_xlabel("Datasets") | |
| ax.set_ylabel("Systematic Generalization Gap") | |
| ax.set_title("Systematic Generalization Gap\n(Lower is Better)") | |
| ax.set_xticks(x + width * (len(model_names) - 1) / 2) | |
| ax.set_xticklabels(datasets, rotation=45) | |
| ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left") | |
| ax.grid(True, alpha=0.3) | |
| def _plot_compositional_accuracy( | |
| self, | |
| ax, | |
| results: Dict[str, Dict[str, EvaluationMetrics]], | |
| model_names: List[str], | |
| datasets: List[str], | |
| ): | |
| x = np.arange(len(datasets)) | |
| width = 0.8 / len(model_names) | |
| for i, model_name in enumerate(model_names): | |
| comp_accs = [ | |
| results[model_name][dataset].compositional_accuracy | |
| for dataset in datasets | |
| ] | |
| ax.bar(x + i * width, comp_accs, width, label=model_name, alpha=0.8) | |
| ax.set_xlabel("Datasets") | |
| ax.set_ylabel("Compositional Accuracy") | |
| ax.set_title("Compositional Accuracy Comparison") | |
| ax.set_xticks(x + width * (len(model_names) - 1) / 2) | |
| ax.set_xticklabels(datasets, rotation=45) | |
| ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left") | |
| ax.grid(True, alpha=0.3) | |
| ax.set_ylim(0, 1) | |
| def _plot_inference_time( | |
| self, | |
| ax, | |
| results: Dict[str, Dict[str, EvaluationMetrics]], | |
| model_names: List[str], | |
| datasets: List[str], | |
| ): | |
| x = np.arange(len(datasets)) | |
| width = 0.8 / len(model_names) | |
| for i, model_name in enumerate(model_names): | |
| times = [ | |
| results[model_name][dataset].inference_time for dataset in datasets | |
| ] | |
| ax.bar(x + i * width, times, width, label=model_name, alpha=0.8) | |
| ax.set_xlabel("Datasets") | |
| ax.set_ylabel("Inference Time (seconds)") | |
| ax.set_title("Inference Time Comparison\n(Lower is Better)") | |
| ax.set_xticks(x + width * (len(model_names) - 1) / 2) | |
| ax.set_xticklabels(datasets, rotation=45) | |
| ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left") | |
| ax.grid(True, alpha=0.3) | |
| def _plot_length_generalization( | |
| self, | |
| ax, | |
| results: Dict[str, Dict[str, EvaluationMetrics]], | |
| model_names: List[str], | |
| datasets: List[str], | |
| ): | |
| x = np.arange(len(datasets)) | |
| width = 0.8 / len(model_names) | |
| for i, model_name in enumerate(model_names): | |
| length_accs = [ | |
| results[model_name][dataset].length_generalization_accuracy | |
| for dataset in datasets | |
| ] | |
| ax.bar(x + i * width, length_accs, width, label=model_name, alpha=0.8) | |
| ax.set_xlabel("Datasets") | |
| ax.set_ylabel("Length Generalization Accuracy") | |
| ax.set_title("Length Generalization Performance") | |
| ax.set_xticks(x + width * (len(model_names) - 1) / 2) | |
| ax.set_xticklabels(datasets, rotation=45) | |
| ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left") | |
| ax.grid(True, alpha=0.3) | |
| ax.set_ylim(0, 1) | |
| def _plot_complexity_generalization( | |
| self, | |
| ax, | |
| results: Dict[str, Dict[str, EvaluationMetrics]], | |
| model_names: List[str], | |
| datasets: List[str], | |
| ): | |
| x = np.arange(len(datasets)) | |
| width = 0.8 / len(model_names) | |
| for i, model_name in enumerate(model_names): | |
| complexity_accs = [ | |
| results[model_name][dataset].complexity_generalization_accuracy | |
| for dataset in datasets | |
| ] | |
| ax.bar(x + i * width, complexity_accs, width, label=model_name, alpha=0.8) | |
| ax.set_xlabel("Datasets") | |
| ax.set_ylabel("Complexity Generalization Accuracy") | |
| ax.set_title("Complexity Generalization Performance") | |
| ax.set_xticks(x + width * (len(model_names) - 1) / 2) | |
| ax.set_xticklabels(datasets, rotation=45) | |
| ax.legend(bbox_to_anchor=(1.05, 1), loc="upper left") | |
| ax.grid(True, alpha=0.3) | |
| ax.set_ylim(0, 1) | |
| def plot_training_dynamics( | |
| self, | |
| training_history: Dict[str, List[float]], | |
| model_name: str = "Model", | |
| save_plots: bool = True, | |
| ) -> None: | |
| fig, axes = plt.subplots(2, 2, figsize=(15, 10)) | |
| fig.suptitle(f"Training Dynamics: {model_name}", fontsize=16, fontweight="bold") | |
| epochs = range(1, len(training_history["train_losses"]) + 1) | |
| axes[0, 0].plot( | |
| epochs, | |
| training_history["train_losses"], | |
| label="Training Loss", | |
| linewidth=2, | |
| color="blue", | |
| ) | |
| axes[0, 0].plot( | |
| epochs, | |
| training_history["val_losses"], | |
| label="Validation Loss", | |
| linewidth=2, | |
| color="red", | |
| linestyle="--", | |
| ) | |
| axes[0, 0].set_xlabel("Epoch") | |
| axes[0, 0].set_ylabel("Loss") | |
| axes[0, 0].set_title("Loss Curves") | |
| axes[0, 0].legend() | |
| axes[0, 0].grid(True, alpha=0.3) | |
| axes[0, 1].plot( | |
| epochs, | |
| training_history["train_accuracies"], | |
| label="Training Accuracy", | |
| linewidth=2, | |
| color="green", | |
| ) | |
| axes[0, 1].plot( | |
| epochs, | |
| training_history["val_accuracies"], | |
| label="Validation Accuracy", | |
| linewidth=2, | |
| color="orange", | |
| linestyle="--", | |
| ) | |
| axes[0, 1].set_xlabel("Epoch") | |
| axes[0, 1].set_ylabel("Accuracy") | |
| axes[0, 1].set_title("Accuracy Curves") | |
| axes[0, 1].legend() | |
| axes[0, 1].grid(True, alpha=0.3) | |
| axes[0, 1].set_ylim(0, 1) | |
| generalization_gap = [ | |
| t - v | |
| for t, v in zip( | |
| training_history["train_accuracies"], training_history["val_accuracies"] | |
| ) | |
| ] | |
| axes[1, 0].plot(epochs, generalization_gap, linewidth=2, color="purple") | |
| axes[1, 0].set_xlabel("Epoch") | |
| axes[1, 0].set_ylabel("Generalization Gap") | |
| axes[1, 0].set_title("Generalization Gap Over Time") | |
| axes[1, 0].grid(True, alpha=0.3) | |
| axes[1, 0].axhline(y=0, color="black", linestyle="-", alpha=0.3) | |
| if "learning_rates" in training_history: | |
| axes[1, 1].plot( | |
| epochs, training_history["learning_rates"], linewidth=2, color="brown" | |
| ) | |
| axes[1, 1].set_xlabel("Epoch") | |
| axes[1, 1].set_ylabel("Learning Rate") | |
| axes[1, 1].set_title("Learning Rate Schedule") | |
| axes[1, 1].grid(True, alpha=0.3) | |
| else: | |
| convergence_window = 10 | |
| if len(training_history["val_losses"]) >= convergence_window: | |
| rolling_mean = ( | |
| pd.Series(training_history["val_losses"]) | |
| .rolling(window=convergence_window) | |
| .mean() | |
| ) | |
| axes[1, 1].plot(epochs, rolling_mean, linewidth=2, color="brown") | |
| axes[1, 1].set_xlabel("Epoch") | |
| axes[1, 1].set_ylabel( | |
| f"Rolling Mean Loss (window={convergence_window})" | |
| ) | |
| axes[1, 1].set_title("Convergence Analysis") | |
| axes[1, 1].grid(True, alpha=0.3) | |
| plt.tight_layout() | |
| if save_plots: | |
| plt.savefig( | |
| self.plots_dir | |
| / f'training_dynamics_{model_name.lower().replace(" ", "_")}.png', | |
| dpi=300, | |
| bbox_inches="tight", | |
| ) | |
| plt.show() | |
| def plot_systematic_generalization_analysis( | |
| self, results: Dict[str, Dict[str, EvaluationMetrics]], save_plots: bool = True | |
| ) -> None: | |
| fig, axes = plt.subplots(2, 2, figsize=(16, 12)) | |
| fig.suptitle( | |
| "Systematic Generalization Analysis", fontsize=16, fontweight="bold" | |
| ) | |
| model_names = list(results.keys()) | |
| datasets = list(results[model_names[0]].keys()) | |
| accuracies = [] | |
| gaps = [] | |
| colors = [] | |
| for i, model_name in enumerate(model_names): | |
| for dataset in datasets: | |
| metrics = results[model_name][dataset] | |
| accuracies.append(metrics.accuracy) | |
| gaps.append(metrics.systematic_generalization_gap) | |
| colors.append(i) | |
| scatter = axes[0, 0].scatter( | |
| accuracies, gaps, c=colors, cmap="tab10", alpha=0.7, s=100 | |
| ) | |
| axes[0, 0].set_xlabel("Accuracy") | |
| axes[0, 0].set_ylabel("Systematic Generalization Gap") | |
| axes[0, 0].set_title("Accuracy vs Systematic Gap") | |
| axes[0, 0].grid(True, alpha=0.3) | |
| for i, model_name in enumerate(model_names): | |
| axes[0, 0].scatter([], [], c=plt.cm.tab10(i), label=model_name) | |
| axes[0, 0].legend() | |
| comp_acc_matrix = np.zeros((len(model_names), len(datasets))) | |
| for i, model_name in enumerate(model_names): | |
| for j, dataset in enumerate(datasets): | |
| comp_acc_matrix[i, j] = results[model_name][ | |
| dataset | |
| ].compositional_accuracy | |
| im = axes[0, 1].imshow( | |
| comp_acc_matrix, cmap="RdYlGn", aspect="auto", vmin=0, vmax=1 | |
| ) | |
| axes[0, 1].set_xticks(range(len(datasets))) | |
| axes[0, 1].set_xticklabels(datasets, rotation=45) | |
| axes[0, 1].set_yticks(range(len(model_names))) | |
| axes[0, 1].set_yticklabels(model_names) | |
| axes[0, 1].set_title("Compositional Accuracy Heatmap") | |
| plt.colorbar(im, ax=axes[0, 1]) | |
| complexities = [] | |
| performances = [] | |
| for model_name in model_names: | |
| for dataset in datasets: | |
| metrics = results[model_name][dataset] | |
| complexities.append(metrics.systematic_generalization_gap) | |
| performances.append(metrics.accuracy) | |
| axes[1, 0].scatter(complexities, performances, alpha=0.7, s=100) | |
| axes[1, 0].set_xlabel("Systematic Generalization Gap (Complexity)") | |
| axes[1, 0].set_ylabel("Overall Performance") | |
| axes[1, 0].set_title("Performance vs Complexity") | |
| axes[1, 0].grid(True, alpha=0.3) | |
| model_scores = {} | |
| for model_name in model_names: | |
| scores = [] | |
| for dataset in datasets: | |
| metrics = results[model_name][dataset] | |
| score = metrics.accuracy - metrics.systematic_generalization_gap | |
| scores.append(score) | |
| model_scores[model_name] = np.mean(scores) | |
| sorted_models = sorted(model_scores.items(), key=lambda x: x[1], reverse=True) | |
| model_names_sorted = [item[0] for item in sorted_models] | |
| scores_sorted = [item[1] for item in sorted_models] | |
| bars = axes[1, 1].bar( | |
| range(len(model_names_sorted)), | |
| scores_sorted, | |
| color=plt.cm.viridis(np.linspace(0, 1, len(model_names_sorted))), | |
| ) | |
| axes[1, 1].set_xticks(range(len(model_names_sorted))) | |
| axes[1, 1].set_xticklabels(model_names_sorted, rotation=45) | |
| axes[1, 1].set_ylabel("Combined Score (Accuracy - Systematic Gap)") | |
| axes[1, 1].set_title("Model Ranking by Systematic Generalization") | |
| axes[1, 1].grid(True, alpha=0.3) | |
| for i, (bar, score) in enumerate(zip(bars, scores_sorted)): | |
| axes[1, 1].text( | |
| bar.get_x() + bar.get_width() / 2, | |
| bar.get_height() + 0.01, | |
| f"{score:.3f}", | |
| ha="center", | |
| va="bottom", | |
| ) | |
| plt.tight_layout() | |
| if save_plots: | |
| plt.savefig( | |
| self.plots_dir / "systematic_generalization_analysis.png", | |
| dpi=300, | |
| bbox_inches="tight", | |
| ) | |
| plt.show() | |
| def create_interactive_dashboard( | |
| self, results: Dict[str, Dict[str, EvaluationMetrics]] | |
| ) -> None: | |
| data = [] | |
| for model_name, model_results in results.items(): | |
| for dataset_name, metrics in model_results.items(): | |
| data.append( | |
| { | |
| "Model": model_name, | |
| "Dataset": dataset_name, | |
| "Accuracy": metrics.accuracy, | |
| "Systematic Gap": metrics.systematic_generalization_gap, | |
| "Compositional Accuracy": metrics.compositional_accuracy, | |
| "Length Generalization": metrics.length_generalization_accuracy, | |
| "Complexity Generalization": metrics.complexity_generalization_accuracy, | |
| "Inference Time": metrics.inference_time, | |
| } | |
| ) | |
| df = pd.DataFrame(data) | |
| fig = make_subplots( | |
| rows=2, | |
| cols=2, | |
| subplot_titles=( | |
| "Accuracy Comparison", | |
| "Systematic Generalization Gap", | |
| "Compositional Accuracy", | |
| "Inference Time", | |
| ), | |
| specs=[ | |
| [{"type": "bar"}, {"type": "bar"}], | |
| [{"type": "bar"}, {"type": "bar"}], | |
| ], | |
| ) | |
| for model_name in df["Model"].unique(): | |
| model_data = df[df["Model"] == model_name] | |
| fig.add_trace( | |
| go.Bar( | |
| name=model_name, | |
| x=model_data["Dataset"], | |
| y=model_data["Accuracy"], | |
| showlegend=True, | |
| ), | |
| row=1, | |
| col=1, | |
| ) | |
| fig.add_trace( | |
| go.Bar( | |
| name=model_name, | |
| x=model_data["Dataset"], | |
| y=model_data["Systematic Gap"], | |
| showlegend=False, | |
| ), | |
| row=1, | |
| col=2, | |
| ) | |
| fig.add_trace( | |
| go.Bar( | |
| name=model_name, | |
| x=model_data["Dataset"], | |
| y=model_data["Compositional Accuracy"], | |
| showlegend=False, | |
| ), | |
| row=2, | |
| col=1, | |
| ) | |
| fig.add_trace( | |
| go.Bar( | |
| name=model_name, | |
| x=model_data["Dataset"], | |
| y=model_data["Inference Time"], | |
| showlegend=False, | |
| ), | |
| row=2, | |
| col=2, | |
| ) | |
| fig.update_layout( | |
| title_text="Systematic Generalization Interactive Dashboard", | |
| height=800, | |
| showlegend=True, | |
| ) | |
| fig.write_html(str(self.plots_dir / "interactive_dashboard.html")) | |
| fig.show() | |
| def generate_comprehensive_report( | |
| self, results: Dict[str, Dict[str, EvaluationMetrics]], output_file: str = None | |
| ) -> str: | |
| report_lines = [] | |
| report_lines.append("# Comprehensive Systematic Generalization Analysis Report") | |
| report_lines.append( | |
| f"Generated at: {pd.Timestamp.now().strftime('%Y-%m-%d %H:%M:%S')}" | |
| ) | |
| report_lines.append("") | |
| report_lines.append("## Executive Summary") | |
| report_lines.append("") | |
| all_accuracies = [] | |
| all_gaps = [] | |
| all_comp_accs = [] | |
| for model_results in results.values(): | |
| for metrics in model_results.values(): | |
| all_accuracies.append(metrics.accuracy) | |
| all_gaps.append(metrics.systematic_generalization_gap) | |
| all_comp_accs.append(metrics.compositional_accuracy) | |
| 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("## Model Rankings") | |
| report_lines.append("") | |
| model_scores = {} | |
| for model_name, model_results in results.items(): | |
| scores = [] | |
| for metrics in model_results.values(): | |
| score = metrics.accuracy - metrics.systematic_generalization_gap | |
| scores.append(score) | |
| model_scores[model_name] = np.mean(scores) | |
| sorted_models = sorted(model_scores.items(), key=lambda x: x[1], reverse=True) | |
| report_lines.append("| Rank | Model | Combined Score |") | |
| report_lines.append("|------|-------|----------------|") | |
| for i, (model_name, score) in enumerate(sorted_models, 1): | |
| report_lines.append(f"| {i} | {model_name} | {score:.4f} |") | |
| report_lines.append("") | |
| report_lines.append("## Detailed Analysis") | |
| report_lines.append("") | |
| for model_name, model_results in results.items(): | |
| report_lines.append(f"### {model_name}") | |
| report_lines.append("") | |
| 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()] | |
| ) | |
| avg_comp_acc = np.mean( | |
| [m.compositional_accuracy for m in model_results.values()] | |
| ) | |
| report_lines.append(f"- **Average Accuracy**: {avg_accuracy:.4f}") | |
| report_lines.append(f"- **Average Systematic Gap**: {avg_gap:.4f}") | |
| report_lines.append( | |
| f"- **Average Compositional Accuracy**: {avg_comp_acc:.4f}" | |
| ) | |
| 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( | |
| 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("") | |
| best_model = sorted_models[0][0] | |
| worst_gap_model = min( | |
| results.items(), | |
| key=lambda x: np.mean( | |
| [m.systematic_generalization_gap for m in x[1].values()] | |
| ), | |
| )[0] | |
| report_lines.append( | |
| f"1. **Best Overall Model**: {best_model} shows the best balance of accuracy and systematic generalization." | |
| ) | |
| report_lines.append( | |
| f"2. **Lowest Systematic Gap**: {worst_gap_model} demonstrates the best systematic generalization capabilities." | |
| ) | |
| report_lines.append( | |
| "3. **Architecture Insights**: Models with explicit compositional structure tend to perform better on systematic generalization tasks." | |
| ) | |
| report_lines.append( | |
| "4. **Training Recommendations**: Consider using systematic data augmentation and compositional loss functions." | |
| ) | |
| report_lines.append("") | |
| report_text = "\n".join(report_lines) | |
| if output_file: | |
| with open(output_file, "w") as f: | |
| f.write(report_text) | |
| return report_text | |
| class CompositionalPatternAnalyzer: | |
| def __init__(self, results_dir: str = "results"): | |
| self.results_dir = results_dir | |
| self.console = Console() | |
| def analyze_compositional_patterns( | |
| self, task: SystematicGeneralizationTask | |
| ) -> Dict[str, Any]: | |
| analysis = { | |
| "task_statistics": task.get_statistics(), | |
| "compositional_analysis": {}, | |
| "pattern_analysis": {}, | |
| } | |
| all_examples = task.training_examples + task.test_examples | |
| component_counts = {} | |
| for example in all_examples: | |
| for component in example.components: | |
| comp_name = component.name | |
| if comp_name not in component_counts: | |
| component_counts[comp_name] = {"count": 0, "examples": []} | |
| component_counts[comp_name]["count"] += 1 | |
| component_counts[comp_name]["examples"].append(example.structure) | |
| analysis["component_frequency"] = component_counts | |
| pattern_counts = {} | |
| for example in all_examples: | |
| pattern = tuple(comp.type.value for comp in example.components) | |
| if pattern not in pattern_counts: | |
| pattern_counts[pattern] = 0 | |
| pattern_counts[pattern] += 1 | |
| analysis["pattern_frequency"] = pattern_counts | |
| complexity_distribution = {} | |
| for example in all_examples: | |
| complexity = example.complexity | |
| if complexity not in complexity_distribution: | |
| complexity_distribution[complexity] = 0 | |
| complexity_distribution[complexity] += 1 | |
| analysis["complexity_distribution"] = complexity_distribution | |
| return analysis | |
| def visualize_compositional_patterns( | |
| self, analysis: Dict[str, Any], save_plots: bool = True | |
| ) -> None: | |
| fig, axes = plt.subplots(2, 2, figsize=(16, 12)) | |
| fig.suptitle("Compositional Pattern Analysis", fontsize=16, fontweight="bold") | |
| component_counts = analysis["component_frequency"] | |
| components = list(component_counts.keys()) | |
| counts = [component_counts[comp]["count"] for comp in components] | |
| axes[0, 0].bar(range(len(components)), counts) | |
| axes[0, 0].set_xticks(range(len(components))) | |
| axes[0, 0].set_xticklabels(components, rotation=45) | |
| axes[0, 0].set_ylabel("Frequency") | |
| axes[0, 0].set_title("Component Frequency Distribution") | |
| axes[0, 0].grid(True, alpha=0.3) | |
| pattern_counts = analysis["pattern_frequency"] | |
| patterns = list(pattern_counts.keys()) | |
| pattern_labels = [str(p) for p in patterns] | |
| pattern_values = list(pattern_counts.values()) | |
| axes[0, 1].bar(range(len(patterns)), pattern_values) | |
| axes[0, 1].set_xticks(range(len(patterns))) | |
| axes[0, 1].set_xticklabels(pattern_labels, rotation=45) | |
| axes[0, 1].set_ylabel("Frequency") | |
| axes[0, 1].set_title("Pattern Frequency Distribution") | |
| axes[0, 1].grid(True, alpha=0.3) | |
| complexity_dist = analysis["complexity_distribution"] | |
| complexities = list(complexity_dist.keys()) | |
| complexity_values = list(complexity_dist.values()) | |
| axes[1, 0].bar(complexities, complexity_values) | |
| axes[1, 0].set_xlabel("Complexity Level") | |
| axes[1, 0].set_ylabel("Frequency") | |
| axes[1, 0].set_title("Complexity Distribution") | |
| axes[1, 0].grid(True, alpha=0.3) | |
| components = list(component_counts.keys()) | |
| cooccurrence_matrix = np.zeros((len(components), len(components))) | |
| for i, comp1 in enumerate(components): | |
| for j, comp2 in enumerate(components): | |
| if i != j: | |
| examples1 = set(component_counts[comp1]["examples"]) | |
| examples2 = set(component_counts[comp2]["examples"]) | |
| cooccurrence = len(examples1.intersection(examples2)) | |
| cooccurrence_matrix[i, j] = cooccurrence | |
| im = axes[1, 1].imshow(cooccurrence_matrix, cmap="Blues", aspect="auto") | |
| axes[1, 1].set_xticks(range(len(components))) | |
| axes[1, 1].set_xticklabels(components, rotation=45) | |
| axes[1, 1].set_yticks(range(len(components))) | |
| axes[1, 1].set_yticklabels(components) | |
| axes[1, 1].set_title("Component Co-occurrence Matrix") | |
| plt.colorbar(im, ax=axes[1, 1]) | |
| plt.tight_layout() | |
| if save_plots: | |
| plots_dir = Path(self.results_dir) / "plots" | |
| plots_dir.mkdir(parents=True, exist_ok=True) | |
| plt.savefig( | |
| plots_dir / "compositional_patterns.png", dpi=300, bbox_inches="tight" | |
| ) | |
| plt.show() | |
Xet Storage Details
- Size:
- 28.5 kB
- Xet hash:
- 0309247a2d51a12db87fad8388f6dcb51e278608b583ac3a096f75e8edab5ad8
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.