tahamajs's picture
download
raw
28.5 kB
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.