import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
import numpy as np
import pandas as pd
from statsmodels.tsa.stattools import grangercausalitytests

# --- 1. Define Model & Gravity Mechanism ---
class GravityNet(nn.Module):
    def __init__(self):
        super(GravityNet, self).__init__()
        self.features = nn.Sequential(
            nn.Conv2d(1, 16, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Flatten(),
            nn.Linear(16 * 14 * 14, 10)
        )
    def forward(self, x):
        return self.features(x)

def apply_gravity(model, strength=0.01):
    """
    Applies a persistent directional pull (gravity/debt) 
    to the vector field of gradients.
    """
    with torch.no_grad():
        for param in model.parameters():
            if param.grad is not None:
                # Gravity pulls weights toward a fixed base structure 
                # (In this case, a downward structural bias toward a negative offset)
                param.grad.add_(strength * torch.sign(param.data))

# --- 2. Training and Evaluation Function ---
def run_experiment(use_gravity, gravity_strength=0.02):
    transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,))])
    train_loader = DataLoader(datasets.MNIST('../data', train=True, download=True, transform=transform), batch_size=128, shuffle=True)
    test_loader = DataLoader(datasets.MNIST('../data', train=False, transform=transform), batch_size=1000, shuffle=False)
    
    model = GravityNet()
    optimizer = optim.SGD(model.parameters(), lr=0.01)
    criterion = nn.CrossEntropyLoss()
    
    accuracy_history = []
    
    # Run 1 quick epoch but sample accuracy over short steps to build a time-series trajectory
    model.train()
    for batch_idx, (data, target) in enumerate(train_loader):

        for i in range(10):
            optimizer.zero_grad()
            output = model(data)
            loss = criterion(output, target)
            loss.backward()
            
            if use_gravity:
                apply_gravity(model, strength=gravity_strength)
                
            optimizer.step()
        
        # Track training accuracy tracking dynamically
        if batch_idx % 1 == 0:
            pred = output.argmax(dim=1, keepdim=True)
            acc = pred.eq(target.view_as(pred)).sum().item() / len(data)
            accuracy_history.append(acc)
            print("acc", acc)
               
        if batch_idx >= 200: # Limit steps for a fast test
            break
            
    # Final Test Evaluation
    model.eval()
    correct = 0
    with torch.no_grad():
        for data, target in test_loader:
            output = model(data)
            pred = output.argmax(dim=1, keepdim=True)
            correct += pred.eq(target.view_as(pred)).sum().item()
            
    final_test_acc = correct / len(test_loader.dataset)
    return accuracy_history, final_test_acc

# --- 3. Compute Granger Causality ---
def calculate_causality(control_history, gravity_history):
    """
    Uses Granger Causality to test if the gravity intervention 
    causally forces a shift in accuracy trajectories.
    """
    print("\n--- Causality Analysis ---")
    # Structure data as a time series dataframe
    df = pd.DataFrame({
        'Control_Acc': control_history,
        'Gravity_Acc': gravity_history
    }).diff().dropna() # Stationarize the time series
    
    # Test if Gravity_Acc Granger-causes Control_Acc (checks landscape interaction)
    # Maxlag 2 checking past trajectory steps
    try:
        gc_res = grangercausalitytests(df[['Control_Acc', 'Gravity_Acc']], maxlag=[2])
        p_value = gc_res[2][0]['ssr_ftest'][1]
        
        print(f"Granger Causality p-value: {p_value:.5f}")
        if p_value < 0.05:
            print("Verdict: CAUSAL EFFECT DETECTED. The gravitational vector field fundamentally altered the trajectory of thought.")
        else:
            print("Verdict: No statistically significant causal divergence found in this short run.")
    except Exception as e:
        print(f"Could not calculate causality matrix: {e}. Try increasing training steps.")

# --- 4. Execution ---
if __name__ == "__main__":
  
    print("Running Experiment (With Gravity Force Applied to Landscape)...")
    gravity_history, gravity_test_acc = run_experiment(use_gravity=True, gravity_strength=0.0001) #0.04

    print("Running Baseline Control (No Gravity)...")
    control_history, control_test_acc = run_experiment(use_gravity=False)
    
    print("\n--- Final Results ---")
    print(f"Control (No Gravity) Test Accuracy: {control_test_acc * 100:.2f}%")
    print(f"Gravity Model Test Accuracy:       {gravity_test_acc * 100:.2f}%")
    
    calculate_causality(control_history, gravity_history)
