import numpy as np
import torch
import torch.nn as nn
import matplotlib.pyplot as plt

# Set seed for reproducibility
np.random.seed(42)
torch.manual_seed(42)

# Generate synthetic time series data
f = lambda mu, s, n: np.hstack([np.random.normal(mu[i], s[i], n[i]) for i in range(len(n))])

mu = np.random.rand(100)
s = np.random.rand(100)
n = np.random.randint(100, 200, 100)
y = f(mu, s, n)

# Create time series dataset
class TimeSeriesDataset(torch.utils.data.Dataset):
    def __init__(self, data, seq_length=20, pred_horizon=5):
        self.seq_length = seq_length
        self.pred_horizon = pred_horizon
        self.data = torch.FloatTensor(data).unsqueeze(1)  # (N, 1)
        
    def __len__(self):
        return len(self.data) - self.seq_length - self.pred_horizon
    
    def __getitem__(self, idx):
        x = self.data[idx:idx + self.seq_length]
        y = self.data[idx + self.seq_length:idx + self.seq_length + self.pred_horizon]
        return x, y.squeeze(1)

# LSTM Forecasting Model
class LSTMForecaster(nn.Module):
    def __init__(self, input_size=1, hidden_size=64, num_layers=2, pred_horizon=5):
        super().__init__()
        self.lstm = nn.LSTM(input_size, hidden_size, num_layers, batch_first=True, dropout=0.2)
        self.fc = nn.Sequential(
            nn.Linear(hidden_size, 32),
            nn.ReLU(),
            nn.Linear(32, pred_horizon)
        )
        
    def forward(self, x):
        lstm_out, _ = self.lstm(x)
        out = self.fc(lstm_out[:, -1, :])  # Use last output
        return out

# Normalize data
class Normalizer:
    def __init__(self):
        self.min, self.max = None, None
        
    def fit(self, data):
        self.min, self.max = data.min(), data.max()
        
    def transform(self, data):
        return (data - self.min) / (self.max - self.min + 1e-8)
    
    def inverse_transform(self, data):
        return data * (self.max - self.min + 1e-8) + self.min

# Training
def train_model(model, train_loader, val_loader, epochs=50, lr=0.001):
    optimizer = torch.optim.Adam(model.parameters(), lr=lr)
    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=5, factor=0.5)
    criterion = nn.MSELoss()
    
    history = {'train_loss': [], 'val_loss': []}
    
    for epoch in range(epochs):
        model.train()
        train_loss = 0
        for x, y in train_loader:
            pred = model(x)
            loss = criterion(pred, y)
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            train_loss += loss.item()
        
        model.eval()
        val_loss = 0
        with torch.no_grad():
            for x, y in val_loader:
                pred = model(x)
                val_loss += criterion(pred, y).item()
        
        train_loss /= len(train_loader)
        val_loss /= len(val_loader)
        scheduler.step(val_loss)
        
        history['train_loss'].append(train_loss)
        history['val_loss'].append(val_loss)
        
        if epoch % 10 == 0:
            print(f"Epoch {epoch}: Train Loss = {train_loss:.4f}, Val Loss = {val_loss:.4f}")
    
    return history

# Split data
split = int(0.8 * len(y))
train_data, test_data = y[:split], y[split:]

# Normalize
normalizer = Normalizer()
normalizer.fit(train_data)
train_norm = normalizer.transform(train_data)
test_norm = normalizer.transform(test_data)

# Create datasets
seq_length, pred_horizon = 20, 5
train_dataset = TimeSeriesDataset(train_norm, seq_length, pred_horizon)
test_dataset = TimeSeriesDataset(test_norm, seq_length, pred_horizon)

train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=64, shuffle=True)
test_loader = torch.utils.data.DataLoader(test_dataset, batch_size=64)

# Initialize and train
model = LSTMForecaster(input_size=1, hidden_size=64, num_layers=2, pred_horizon=pred_horizon)
history = train_model(model, train_loader, test_loader, epochs=50)

# Forecast future values
def forecast(model, last_sequence, steps=10):
    model.eval()
    predictions = []
    seq = last_sequence.clone()
    
    with torch.no_grad():
        for _ in range(steps):
            pred = model(seq.unsqueeze(0))
            predictions.append(pred.item())
            seq = torch.cat([seq[1:], pred.view(1, 1)], dim=0)
    
    return np.array(predictions)

# Visualize results
fig, axes = plt.subplots(2, 2, figsize=(14, 8))

# Training loss
axes[0, 0].plot(history['train_loss'], label='Train')
axes[0, 0].plot(history['val_loss'], label='Validation')
axes[0, 0].set_title('Training Progress')
axes[0, 0].set_xlabel('Epoch')
axes[0, 0].legend()

# Sample predictions
model.eval()
sample_idx = 100
sample_input = torch.FloatTensor(train_dataset.data[sample_idx:sample_idx+seq_length])
actual_future = test_data[sample_idx:sample_idx+pred_horizon]
pred_future = forecast(model, sample_input, pred_horizon)

axes[0, 1].plot(range(seq_length), normalizer.transform(train_data[sample_idx:sample_idx+seq_length]), label='Input')
axes[0, 1].plot(range(seq_length, seq_length+pred_horizon), normalizer.transform(actual_future), 'g-', label='Actual')
axes[0, 1].plot(range(seq_length, seq_length+pred_horizon), pred_future, 'r--', label='Predicted')
axes[0, 1].set_title('Sample Forecast')
axes[0, 1].legend()

# Original data distribution
axes[1, 0].hist(y, 50, edgecolor='black')
axes[1, 0].set_title('Data Distribution')

# Forecast comparison
full_pred = []
model.eval()
with torch.no_grad():
    for i in range(0, len(test_norm) - seq_length, pred_horizon):
        x = torch.FloatTensor(test_dataset.data[i:i+seq_length])
        pred = model(x.unsqueeze(0)).squeeze().numpy()
        full_pred.extend(pred)

axes[1, 1].plot(normalizer.inverse_transform(test_norm[seq_length:]), label='Actual', alpha=0.7)
axes[1, 1].plot(normalizer.inverse_transform(np.array(full_pred[:len(test_norm)-seq_length])), label='Forecast', alpha=0.7)
axes[1, 1].set_title('Test Set Forecast vs Actual')
axes[1, 1].legend()

plt.tight_layout()
plt.savefig('forecast_results.png', dpi=150)
plt.show()
