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

np.random.seed(42)
torch.manual_seed(42)

# ============================================
# YOUR ORIGINAL DATA STRUCTURE
# ============================================
mu_true = np.random.rand(100)
s_true = np.random.rand(100)
n_true = np.random.randint(100, 200, 100)

f_original = lambda mu, s, n: np.hstack(
    [np.random.normal(mu[i], s[i], n[i]) for i in range(len(n))]
)
y_sequential = f_original(mu_true, s_true, n_true)

# ============================================
# SEGMENT INTO INDIVIDUAL SERIES
# ============================================
def segment_by_lengths(data, lengths):
    series_list = []
    idx = 0
    for length in lengths:
        series_list.append(data[idx:idx + length])
        idx += length
    return series_list

series_list = segment_by_lengths(y_sequential, n_true)

# ============================================
# PADDED DATASET (Fixed shape handling)
# ============================================
class PaddedTimeSeriesDataset(torch.utils.data.Dataset):
    def __init__(self, series_list, seq_length=20, pred_horizon=5):
        self.seq_length = seq_length
        self.pred_horizon = pred_horizon
        
        max_len = max(len(s) for s in series_list)
        self.padded_data = []
        self.mask = []
        
        for series in series_list:
            padded = np.zeros(max_len, dtype=np.float32)
            padded[:len(series)] = series
            
            mask = np.zeros(max_len, dtype=np.float32)
            mask[:len(series)] = 1.0
            
            self.padded_data.append(padded)
            self.mask.append(mask)
        
        self.padded_data = torch.FloatTensor(np.array(self.padded_data)).unsqueeze(-1)  # (N, L, 1)
        self.mask = torch.FloatTensor(np.array(self.mask))
        
    def __len__(self):
        total = 0
        for i in range(len(self.padded_data)):
            valid_len = int(self.mask[i].sum().item())
            total += max(0, valid_len - self.seq_length - self.pred_horizon)
        return total
    
    def __getitem__(self, idx):
        for series_idx in range(len(self.padded_data)):
            valid_len = int(self.mask[series_idx].sum().item())
            available = max(0, valid_len - self.seq_length - self.pred_horizon)
            
            if idx < available:
                start = idx
                break
            idx -= available
        
        x = self.padded_data[series_idx, start:start + self.seq_length]       # (seq_len, 1)
        y = self.padded_data[series_idx, start + self.seq_length:start + self.seq_length + self.pred_horizon]  # (pred_horizon, 1)
        m = self.mask[series_idx, start + self.seq_length:start + self.seq_length + self.pred_horizon]  # (pred_horizon,)
        
        return x, y.squeeze(-1), m  # Return y as (pred_horizon,) not (pred_horizon, 1)

# ============================================
# MASKED LSTM MODEL
# ============================================
class MaskedLSTM(nn.Module):
    def __init__(self, input_size=1, hidden_size=64, num_layers=2, pred_horizon=5):
        super().__init__()
        self.hidden_size = hidden_size
        self.pred_horizon = pred_horizon
        
        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, mask=None):
        # x: (batch, seq_len, 1)
        lstm_out, _ = self.lstm(x)  # (batch, seq_len, hidden)
        
        # Get last valid hidden state
        if mask is not None:
            last_valid_idx = mask.sum(dim=1).long() - 1
            batch_size = x.size(0)
            last_hidden = lstm_out[torch.arange(batch_size), last_valid_idx]
        else:
            last_hidden = lstm_out[:, -1, :]
        
        return self.fc(last_hidden)  # (batch, pred_horizon)

# ============================================
# TRAINING (Fixed loss calculation)
# ============================================
seq_length = 20
pred_horizon = 5
dataset = PaddedTimeSeriesDataset(series_list, seq_length, pred_horizon)
train_loader = torch.utils.data.DataLoader(dataset, batch_size=32, shuffle=True)

model = MaskedLSTM(input_size=1, hidden_size=64, num_layers=2, pred_horizon=pred_horizon)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)

print("=== Training on Mixed-Length Time Series ===")
for epoch in range(50):
    total_loss = 0
    num_batches = 0
    
    for x, y, mask in train_loader:
        # x: (batch, seq_len, 1), y: (batch, pred_horizon), mask: (batch, pred_horizon)
        
        pred = model(x)  # (batch, pred_horizon)
        
        # Compute MSE only on valid (non-padded) positions
        loss = ((pred - y) ** 2)  # (batch, pred_horizon)
        loss = (loss * mask).sum() / mask.sum()  # Apply mask, average over valid
        
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

        print(loss.item())
                
        total_loss += loss.item()
        num_batches += 1
    
    if epoch % 10 == 0:
        print(f"Epoch {epoch}: Loss = {total_loss / num_batches:.4f}")

print("\nTraining complete!")

# ============================================
# FORECASTING FUNCTION
# ============================================
def forecast_next(model, context, steps=5):
    model.eval()
    with torch.no_grad():
        context_tensor = torch.FloatTensor(context).view(1, -1, 1)  # (1, seq_len, 1)
        predictions = []
        
        for _ in range(steps):
            pred = model(context_tensor)  # (1, pred_horizon)
            predictions.append(pred[0, 0].item())
            
            # Roll context window
            context_tensor = torch.roll(context_tensor, -1, dims=1)
            context_tensor[0, -1, 0] = pred[0, 0].item()
        
        return np.array(predictions)

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

# Original concatenated data
axes[0, 0].plot(y_sequential[:2000], linewidth=0.5)
axes[0, 0].set_title("Original Data (concatenated different-length series)")
axes[0, 0].set_xlabel("Time Index")

# Individual series
for i in range(5):
    offset = sum(n_true[:i])
    axes[0, 1].plot(np.arange(n_true[i]) + offset, series_list[i], linewidth=0.7, alpha=0.7)
axes[0, 1].set_title("Individual Series (first 5)")

# Sample forecast
sample_idx = 50
context = series_list[sample_idx][:seq_length]
actual = series_list[sample_idx][seq_length:seq_length + pred_horizon]
forecasted = forecast_next(model, context, pred_horizon)

axes[1, 0].plot(range(seq_length), context, 'b-', label='Context', linewidth=2)
axes[1, 0].plot(range(seq_length, seq_length + pred_horizon), actual, 'g-', label='Actual', linewidth=2)
axes[1, 0].plot(range(seq_length, seq_length + pred_horizon), forecasted, 'r--', label='Forecast', linewidth=2)
axes[1, 0].axvline(x=seq_length, color='gray', linestyle='--', alpha=0.5)
axes[1, 0].set_title(f"Sample Forecast (series {sample_idx})")
axes[1, 0].legend()
axes[1, 0].grid(True, alpha=0.3)

# Distribution comparison
all_preds = []
model.eval()
with torch.no_grad():
    for series in series_list:
        if len(series) >= seq_length + pred_horizon:
            context = series[:seq_length]
            pred = forecast_next(model, context, pred_horizon)
            all_preds.extend(pred)

axes[1, 1].hist(y_sequential, 50, alpha=0.5, density=True, label='Original')
if all_preds:
    axes[1, 1].hist(all_preds, 30, alpha=0.5, density=True, label='Forecasts')
axes[1, 1].set_title("Distribution: Original vs Forecasts")
axes[1, 1].legend()

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

# ============================================
# EVALUATION METRICS
# ============================================
def evaluate_forecast(model, series_list, seq_length, pred_horizon):
    mse_list, mae_list = [], []
    
    for series in series_list:
        if len(series) >= seq_length + pred_horizon:
            context = series[:seq_length]
            actual = series[seq_length:seq_length + pred_horizon]
            pred = forecast_next(model, context, pred_horizon)
            
            mse_list.append(np.mean((pred - actual) ** 2))
            mae_list.append(np.mean(np.abs(pred - actual)))
    
    return np.mean(mse_list), np.mean(mae_list)

mse, mae = evaluate_forecast(model, series_list, seq_length, pred_horizon)
print(f"\n=== Evaluation Metrics ===")
print(f"Mean MSE: {mse:.4f}")
print(f"Mean MAE: {mae:.4f}")
