tahamajs's picture
download
raw
18.8 kB
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from typing import List, Dict, Tuple, Any, Optional, Union
import numpy as np
import time
import json
import os
import yaml
from dataclasses import dataclass, asdict
from pathlib import Path
import logging
from rich.console import Console
from rich.progress import Progress, TaskID
from rich.table import Table
from rich.panel import Panel
from ..core.components import SystematicGeneralizationTask
from ..models.neural_architectures import (
ModularNetwork,
AttentionComposer,
GraphComposer,
MetaLearningComposer,
HierarchicalComposer,
)
from ..models.symbolic_systems import (
SymbolicReasoner,
ProgramSynthesizer,
GrammarBasedGenerator,
)
from ..models.neurosymbolic_systems import (
NeuralModuleNetwork,
NeuroSymbolicComposer,
HybridReasoningSystem,
)
from ..datasets.dataset_implementations import (
SCANDataset,
ArithmeticReasoningDataset,
VisualReasoningDataset,
SystematicSplitGenerator,
create_data_loader,
)
from ..evaluation.evaluation_framework import (
SystematicGeneralizationEvaluator,
EvaluationMetrics,
)
@dataclass
class ExperimentConfig:
experiment_name: str = "systematic_generalization"
random_seed: int = 42
embed_dim: int = 64
hidden_dim: int = 128
num_layers: int = 3
dropout: float = 0.1
batch_size: int = 32
learning_rate: float = 0.001
num_epochs: int = 100
patience: int = 10
max_length: int = 20
number_range: Tuple[int, int] = (1, 100)
test_split_ratio: float = 0.2
validation_split_ratio: float = 0.1
results_dir: str = "results"
save_models: bool = True
save_plots: bool = True
def to_dict(self) -> Dict[str, Any]:
return asdict(self)
@classmethod
def from_dict(cls, config_dict: Dict[str, Any]) -> "ExperimentConfig":
return cls(**config_dict)
def save(self, filepath: str):
with open(filepath, "w") as f:
yaml.dump(self.to_dict(), f, default_flow_style=False)
@classmethod
def load(cls, filepath: str) -> "ExperimentConfig":
with open(filepath, "r") as f:
config_dict = yaml.safe_load(f)
return cls.from_dict(config_dict)
class CompositionalTrainer:
def __init__(self, model: nn.Module, config: ExperimentConfig):
self.model = model
self.config = config
self.optimizer = optim.Adam(model.parameters(), lr=config.learning_rate)
self.criterion = nn.MSELoss()
self.train_losses = []
self.val_losses = []
self.train_accuracies = []
self.val_accuracies = []
self.console = Console()
self.setup_logging()
def setup_logging(self):
log_dir = Path(self.config.results_dir) / "logs"
log_dir.mkdir(parents=True, exist_ok=True)
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(levelname)s - %(message)s",
handlers=[
logging.FileHandler(log_dir / f"{self.config.experiment_name}.log"),
logging.StreamHandler(),
],
)
self.logger = logging.getLogger(__name__)
def train_epoch(self, train_loader: DataLoader) -> Tuple[float, float]:
self.model.train()
total_loss = 0.0
total_correct = 0
total_samples = 0
for batch in train_loader:
self.optimizer.zero_grad()
outputs = self._forward_pass(batch)
targets = self._get_targets(batch)
loss = self.criterion(outputs.squeeze(), targets.float())
loss.backward()
self.optimizer.step()
total_loss += loss.item()
predicted_classes = (outputs > 0.5).long().squeeze()
total_correct += (predicted_classes == targets.long()).sum().item()
total_samples += targets.size(0)
avg_loss = total_loss / len(train_loader)
accuracy = total_correct / total_samples
return avg_loss, accuracy
def evaluate(self, val_loader: DataLoader) -> Tuple[float, float]:
self.model.eval()
total_loss = 0.0
total_correct = 0
total_samples = 0
with torch.no_grad():
for batch in val_loader:
outputs = self._forward_pass(batch)
targets = self._get_targets(batch)
loss = self.criterion(outputs.squeeze(), targets.float())
total_loss += loss.item()
predicted_classes = (outputs > 0.5).long().squeeze()
total_correct += (predicted_classes == targets.long()).sum().item()
total_samples += targets.size(0)
avg_loss = total_loss / len(val_loader)
accuracy = total_correct / total_samples
return avg_loss, accuracy
def _forward_pass(self, batch: Dict[str, torch.Tensor]) -> torch.Tensor:
if isinstance(self.model, ModularNetwork):
return self.model(
batch["primitive1"], batch["operator"], batch["primitive2"]
)
elif isinstance(
self.model, (AttentionComposer, GraphComposer, HierarchicalComposer)
):
return self.model(batch["input"])
elif isinstance(self.model, NeuralModuleNetwork):
return self.model(batch["command"])
elif isinstance(self.model, NeuroSymbolicComposer):
outputs = self.model(batch["input"])
return outputs["values"]
else:
return self.model(batch["input"])
def _get_targets(self, batch: Dict[str, torch.Tensor]) -> torch.Tensor:
if "target" in batch:
return batch["target"]
elif "result" in batch:
return batch["result"]
else:
return torch.zeros(batch["input"].size(0))
def train(
self, train_loader: DataLoader, val_loader: DataLoader
) -> Dict[str, List[float]]:
best_val_loss = float("inf")
patience_counter = 0
self.logger.info(f"Starting training for {self.config.num_epochs} epochs")
with Progress() as progress:
task = progress.add_task("[green]Training...", total=self.config.num_epochs)
for epoch in range(self.config.num_epochs):
train_loss, train_acc = self.train_epoch(train_loader)
val_loss, val_acc = self.evaluate(val_loader)
self.train_losses.append(train_loss)
self.val_losses.append(val_loss)
self.train_accuracies.append(train_acc)
self.val_accuracies.append(val_acc)
if val_loss < best_val_loss:
best_val_loss = val_loss
patience_counter = 0
else:
patience_counter += 1
if patience_counter >= self.config.patience:
self.logger.info(f"Early stopping at epoch {epoch}")
break
if epoch % 10 == 0:
self.logger.info(
f"Epoch {epoch}: Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}, "
f"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}"
)
progress.update(task, advance=1)
return {
"train_losses": self.train_losses,
"val_losses": self.val_losses,
"train_accuracies": self.train_accuracies,
"val_accuracies": self.val_accuracies,
}
class SystematicGeneralizationExperiment:
def __init__(self, config: ExperimentConfig):
self.config = config
self.console = Console()
self.evaluator = SystematicGeneralizationEvaluator(config.results_dir)
self.setup_directories()
self.models = {}
self.datasets = {}
self.results = {}
self.setup_logging()
def setup_directories(self):
dirs = [
self.config.results_dir,
f"{self.config.results_dir}/experiments",
f"{self.config.results_dir}/plots",
f"{self.config.results_dir}/models",
f"{self.config.results_dir}/logs",
]
for dir_path in dirs:
Path(dir_path).mkdir(parents=True, exist_ok=True)
def setup_logging(self):
log_file = f"{self.config.results_dir}/logs/{self.config.experiment_name}.log"
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s - %(levelname)s - %(message)s",
handlers=[logging.FileHandler(log_file), logging.StreamHandler()],
)
self.logger = logging.getLogger(__name__)
def create_models(self) -> Dict[str, nn.Module]:
vocab_size = 100
models = {
"ModularNetwork": ModularNetwork(
primitive_vocab_size=vocab_size // 2,
operator_vocab_size=vocab_size // 2,
embed_dim=self.config.embed_dim,
hidden_dim=self.config.hidden_dim,
),
"AttentionComposer": AttentionComposer(
vocab_size=vocab_size,
embed_dim=self.config.embed_dim,
hidden_dim=self.config.hidden_dim,
),
"GraphComposer": GraphComposer(
vocab_size=vocab_size, embed_dim=self.config.embed_dim
),
"HierarchicalComposer": HierarchicalComposer(
vocab_size=vocab_size,
embed_dim=self.config.embed_dim,
hidden_dim=self.config.hidden_dim,
),
"NeuralModuleNetwork": NeuralModuleNetwork(
vocab_size=vocab_size, embed_dim=self.config.embed_dim
),
"NeuroSymbolicComposer": NeuroSymbolicComposer(
symbol_vocab_size=vocab_size, embed_dim=self.config.embed_dim
),
}
self.models = models
self.logger.info(f"Created {len(models)} models for comparison")
return models
def create_datasets(self) -> Dict[str, Any]:
datasets = {
"SCAN": SCANDataset(
split_type="simple",
max_length=self.config.max_length,
data_dir=f"{self.config.results_dir}/data",
),
"Arithmetic": ArithmeticReasoningDataset(
number_range=self.config.number_range,
data_dir=f"{self.config.results_dir}/data",
),
"Visual": VisualReasoningDataset(
data_dir=f"{self.config.results_dir}/data"
),
}
self.datasets = datasets
self.logger.info(f"Created {len(datasets)} datasets")
return datasets
def create_systematic_splits(
self, dataset: Any, split_type: str = "compositional"
) -> Tuple[DataLoader, DataLoader, DataLoader]:
if isinstance(dataset, SCANDataset):
train_indices, test_indices = (
SystematicSplitGenerator.create_compositional_split(dataset, ["thrice"])
)
elif isinstance(dataset, ArithmeticReasoningDataset):
train_indices, test_indices = (
SystematicSplitGenerator.create_complexity_split(
dataset, complexity_threshold=2
)
)
else:
train_indices, test_indices = SystematicSplitGenerator.create_length_split(
dataset, length_threshold=10
)
val_size = int(len(train_indices) * self.config.validation_split_ratio)
val_indices = train_indices[:val_size]
train_indices = train_indices[val_size:]
train_loader = create_data_loader(
dataset, train_indices, self.config.batch_size, shuffle=True
)
val_loader = create_data_loader(
dataset, val_indices, self.config.batch_size, shuffle=False
)
test_loader = create_data_loader(
dataset, test_indices, self.config.batch_size, shuffle=False
)
self.logger.info(
f"Split sizes - Train: {len(train_indices)}, Val: {len(val_indices)}, Test: {len(test_indices)}"
)
return train_loader, val_loader, test_loader
def run_single_experiment(
self, model: nn.Module, model_name: str, dataset: Any, dataset_name: str
) -> Dict[str, Any]:
self.logger.info(f"Running experiment: {model_name} on {dataset_name}")
train_loader, val_loader, test_loader = self.create_systematic_splits(dataset)
trainer = CompositionalTrainer(model, self.config)
start_time = time.time()
training_history = trainer.train(train_loader, val_loader)
training_time = time.time() - start_time
test_loss, test_acc = trainer.evaluate(test_loader)
metrics = EvaluationMetrics(
accuracy=test_acc,
precision=0.0,
recall=0.0,
f1_score=0.0,
systematic_generalization_gap=0.0,
compositional_accuracy=0.0,
length_generalization_accuracy=0.0,
complexity_generalization_accuracy=0.0,
training_time=training_time,
inference_time=0.0,
memory_usage=0.0,
)
if self.config.save_models:
model_path = (
f"{self.config.results_dir}/models/{model_name}_{dataset_name}.pt"
)
torch.save(model.state_dict(), model_path)
return {
"metrics": metrics,
"training_history": training_history,
"test_loss": test_loss,
"test_accuracy": test_acc,
}
def run_comprehensive_experiment(self) -> Dict[str, Any]:
self.console.print(
Panel.fit(
f"[bold blue]Systematic Generalization Experiment[/bold blue]\n"
f"Experiment: {self.config.experiment_name}\n"
f"Models: {len(self.models)} models\n"
f"Datasets: {len(self.datasets)} datasets",
title="Experiment Overview",
)
)
self.create_models()
self.create_datasets()
all_results = {}
with Progress() as progress:
total_tasks = len(self.models) * len(self.datasets)
main_task = progress.add_task(
"[green]Running experiments...", total=total_tasks
)
for model_name, model in self.models.items():
all_results[model_name] = {}
for dataset_name, dataset in self.datasets.items():
result = self.run_single_experiment(
model, model_name, dataset, dataset_name
)
all_results[model_name][dataset_name] = result
progress.update(main_task, advance=1)
self.logger.info("Evaluating all models...")
evaluation_results = {}
for model_name in self.models.keys():
evaluation_results[model_name] = {}
for dataset_name in self.datasets.keys():
_, _, test_loader = self.create_systematic_splits(
self.datasets[dataset_name]
)
model_results = self.evaluator.evaluate_model(
self.models[model_name], {dataset_name: test_loader}, model_name
)
evaluation_results[model_name][dataset_name] = model_results[
dataset_name
]
self.generate_reports(evaluation_results)
self.save_results(all_results, evaluation_results)
return {
"experiment_results": all_results,
"evaluation_results": evaluation_results,
}
def generate_reports(
self, evaluation_results: Dict[str, Dict[str, EvaluationMetrics]]
):
report = self.evaluator.generate_evaluation_report(evaluation_results)
report_path = f"{self.config.results_dir}/experiments/{self.config.experiment_name}_report.md"
with open(report_path, "w") as f:
f.write(report)
if self.config.save_plots:
self.evaluator.plot_evaluation_results(evaluation_results, save_plots=True)
self.display_summary_table(evaluation_results)
def display_summary_table(
self, evaluation_results: Dict[str, Dict[str, EvaluationMetrics]]
):
table = Table(title="Systematic Generalization Results Summary")
table.add_column("Model", style="cyan")
table.add_column("Dataset", style="magenta")
table.add_column("Accuracy", justify="right")
table.add_column("Systematic Gap", justify="right")
table.add_column("Compositional Acc", justify="right")
table.add_column("Inference Time", justify="right")
for model_name, model_results in evaluation_results.items():
for dataset_name, metrics in model_results.items():
table.add_row(
model_name,
dataset_name,
f"{metrics.accuracy:.4f}",
f"{metrics.systematic_generalization_gap:.4f}",
f"{metrics.compositional_accuracy:.4f}",
f"{metrics.inference_time:.4f}s",
)
self.console.print(table)
def save_results(
self,
experiment_results: Dict[str, Any],
evaluation_results: Dict[str, Dict[str, EvaluationMetrics]],
):
exp_results_path = f"{self.config.results_dir}/experiments/{self.config.experiment_name}_results.json"
with open(exp_results_path, "w") as f:
json.dump(experiment_results, f, indent=2, default=str)
self.evaluator.save_results(
evaluation_results, f"{self.config.experiment_name}_evaluation.json"
)
config_path = f"{self.config.results_dir}/experiments/{self.config.experiment_name}_config.yaml"
self.config.save(config_path)
self.logger.info(f"Results saved to {self.config.results_dir}/experiments/")
def run_systematic_generalization_experiment(config_path: str = None) -> Dict[str, Any]:
if config_path and os.path.exists(config_path):
config = ExperimentConfig.load(config_path)
else:
config = ExperimentConfig()
torch.manual_seed(config.random_seed)
np.random.seed(config.random_seed)
experiment = SystematicGeneralizationExperiment(config)
results = experiment.run_comprehensive_experiment()
return results
if __name__ == "__main__":
results = run_systematic_generalization_experiment()
print("Experiment completed successfully!")

Xet Storage Details

Size:
18.8 kB
·
Xet hash:
49fcae18554813cfca5a4204134f7dc8747d74943ee446205db7c67c7411a92f

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