import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import StandardScaler
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt

# Specify columns for inputs (features) and output (target). Index starts at 0.
start_col_fetures = 3
end_col_features  = 7  # Note index excluded in slicing with .iloc
col_number_target = 7
csv_path = "WT_Dummy_Fatigue.csv"

# ==========================================
# DATA PREPARATION PIPELINE
# ==========================================
class TabularCSVDataset(Dataset):
    def __init__(self, X, y):
        # Convert arrays to PyTorch Tensors
        self.X = torch.tensor(X, dtype=torch.float32)
        self.y = torch.tensor(y, dtype=torch.float32).unsqueeze(1)

    def __len__(self):
        return len(self.X)

    def __getitem__(self, idx):
        return self.X[idx], self.y[idx]

def prepare_data(csv_file_path):
    # Load raw CSV file
    df = pd.read_csv(csv_file_path)
    
    # Extract features (first 4 columns) and target variable (last column)
    X_raw = df.iloc[:, start_col_fetures:end_col_features].values
    y_raw = df.iloc[:, col_number_target].values
    
    # Perform an 80% Training and 20% Testing data split
    X_train, X_test, y_train, y_test = train_test_split(
        X_raw, y_raw, test_size=0.20, random_state=42
    )
    
    # Normalize features for smoother and faster weight convergence
    scaler = StandardScaler()
    X_train_scaled = scaler.fit_transform(X_train)
    X_test_scaled = scaler.transform(X_test)
    
    # Package into custom Dataset objects
    train_dataset = TabularCSVDataset(X_train_scaled, y_train)
    test_dataset = TabularCSVDataset(X_test_scaled, y_test)
    
    # Create DataLoaders for mini-batch generation
    train_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)
    
    return train_loader, test_dataset, X_test_scaled, y_test

# ==========================================
# DEFINE THE NETWORK ARCHITECTURES
# ==========================================
class ShallowANN(nn.Module):
    def __init__(self, input_dim=4, hidden_dim=32, output_dim=1):
        super(ShallowANN, self).__init__()
        self.hidden_layer = nn.Linear(input_dim, hidden_dim)
        self.relu = nn.ReLU()
        self.output_layer = nn.Linear(hidden_dim, output_dim)
        
    def forward(self, x):
        return self.output_layer(self.relu(self.hidden_layer(x)))

class DeepDNN(nn.Module):
    def __init__(self, input_dim=4, output_dim=1):
        super(DeepDNN, self).__init__()
        self.network = nn.Sequential(
            nn.Linear(input_dim, 64),
            nn.ReLU(),
            nn.Linear(64, 32),
            nn.ReLU(),
            nn.Linear(32, 16),
            nn.ReLU(),
            nn.Linear(16, output_dim)
        )
        
    def forward(self, x):
        return self.network(x)

# ==========================================
# TRAINING LOOP FUNCTION
# ==========================================
def train_model(model, train_loader, epochs=50, lr=0.01):
    # Setting up standard MSE Loss calculation and Adam Optimizer
    criterion = nn.MSELoss()
    optimizer = optim.Adam(model.parameters(), lr=lr)
    
    model.train()
    for epoch in range(epochs):
        epoch_loss = 0.0
        for features, targets in train_loader:
            optimizer.zero_grad()            # 1. Clear previous gradients
            predictions = model(features)     # 2. Forward Propagation
            loss = criterion(predictions, targets) # 3. Calculate current error
            loss.backward()                  # 4. Backpropagation (compute gradients)
            optimizer.step()                 # 5. Optimize internal weights
            
            epoch_loss += loss.item() * features.size(0)
    return model

# ==========================================
# VISUALIZATION ENGINE
# ==========================================
def plot_results(y_true, ann_preds, dnn_preds):
    plt.figure(figsize=(12, 5))
    
    # Plot 1: Shallow ANN Performance
    plt.subplot(1, 2, 1)
    plt.scatter(y_true, ann_preds, color='teal', alpha=0.6, label='Predicted')
    plt.plot([y_true.min(), y_true.max()], [y_true.min(), y_true.max()], 'k--', lw=2, label='Ideal Perfect Fit')
    plt.title('Shallow ANN: Predicted vs Actual')
    plt.xlabel('Actual Target Values')
    plt.ylabel('Predicted Values')
    plt.legend()
    plt.grid(True)
    
    # Plot 2: Deep DNN Performance
    plt.subplot(1, 2, 2)
    plt.scatter(y_true, dnn_preds, color='orangered', alpha=0.6, label='Predicted')
    plt.plot([y_true.min(), y_true.max()], [y_true.min(), y_true.max()], 'k--', lw=2, label='Ideal Perfect Fit')
    plt.title('Deep DNN: Predicted vs Actual')
    plt.xlabel('Actual Target Values')
    plt.ylabel('Predicted Values')
    plt.legend()
    plt.grid(True)
    
    plt.tight_layout()
    plt.show()

# ==========================================
# EXECUTION ENTRY POINT
# ==========================================
if __name__ == "__main__":
    
    try:
        # Preprocess data and enforce the 80/20 split
        train_loader, test_dataset, X_test, y_test = prepare_data(csv_path)
        
        # Instantiate both architectures
        ann_model = ShallowANN(input_dim=4, output_dim=1)
        dnn_model = DeepDNN(input_dim=4, output_dim=1)
        
        # Train models through full weight updates loops
        print("--- Training Shallow ANN...\n")
        ann_model = train_model(ann_model, train_loader, epochs=60)
        
        print("--- Training Deep DNN...\n")
        dnn_model = train_model(dnn_model, train_loader, epochs=60)
        
        # Evaluate both on unseen 20% test data split
        ann_model.eval()
        dnn_model.eval()
        with torch.no_grad():
            X_test_tensor = torch.tensor(X_test, dtype=torch.float32)
            ann_predictions = ann_model(X_test_tensor).numpy().flatten()
            images_predictions = dnn_model(X_test_tensor).numpy().flatten()
            
        # Graph the fit vectors
        plot_results(y_test, ann_predictions, images_predictions)
        
    except FileNotFoundError:
        print(f"Please replace '{csv_path}' with a correct path / data.")

