import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torchdiffeq import odeint

# ------------------- Spatial message passing -------------------
class SpatialPropagator(nn.Module):
    def __init__(self, latent_dim, hidden_dim=64):
        super().__init__()
        self.message_mlp = nn.Sequential(
            nn.Linear(latent_dim * 2, hidden_dim),
            nn.ReLU(),
            nn.Linear(hidden_dim, latent_dim)
        )
        self.update_mlp = nn.Sequential(
            nn.Linear(latent_dim * 2, hidden_dim),
            nn.ReLU(),
            nn.Linear(hidden_dim, latent_dim)
        )

    def forward(self, h):
        B, C, H, W = h.shape
        h_pad = F.pad(h, (1,1,1,1), mode='circular')
        h_left  = h_pad[:,:,1:H+1,0:W]
        h_right = h_pad[:,:,1:H+1,2:W+2]
        h_up    = h_pad[:,:,0:H,  1:W+1]
        h_down  = h_pad[:,:,2:H+2,1:W+1]

        def message(src):
            cat = torch.cat([h, src], dim=1).permute(0,2,3,1).reshape(-1, 2*C)
            msg = self.message_mlp(cat).reshape(B,H,W,C).permute(0,3,1,2)
            return msg

        m = message(h_left) + message(h_right) + message(h_up) + message(h_down)
        cat = torch.cat([h, m], dim=1).permute(0,2,3,1).reshape(-1, 2*C)
        h_new = self.update_mlp(cat).reshape(B,H,W,C).permute(0,3,1,2)
        return h_new

# ------------------- Neural ODE function -------------------
class ODEFunc(nn.Module):
    def __init__(self, latent_dim, hidden_dim=128):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(latent_dim + 1, hidden_dim),
            nn.Tanh(),
            nn.Linear(hidden_dim, latent_dim),
            nn.Tanh(),                       # <<< bound dh/dt to [-1,1]
        )
        for m in self.modules():
            if isinstance(m, nn.Linear):
                nn.init.normal_(m.weight, std=0.05)
                nn.init.zeros_(m.bias)

    def forward(self, t, h):
        B, C, H, W = h.shape
        h_flat = h.permute(0,2,3,1).reshape(-1, C)
        t_tensor = float(t) * torch.ones(B*H*W, 1, device=h.device)
        dstate = self.net(torch.cat([t_tensor, h_flat], dim=1))
        print(dstate)
        return dstate.reshape(B, H, W, C).permute(0,3,1,2)
                
# ------------------- Prior model -------------------
class PriorDynamics:
    def __init__(self, latent_dim, grid_shape):
        self.latent_dim = latent_dim
        self.grid_shape = grid_shape
        self.prior_mean = None
        self.prior_logvar = None

    def set_prior(self, mean, logvar):
        self.prior_mean = mean
        self.prior_logvar = logvar

    def kl_divergence(self, h):
        if self.prior_mean is None:
            return 0.0
        var_prior = torch.exp(self.prior_logvar)
        kl = 0.5 * (var_prior + (h - self.prior_mean)**2 - 1 - self.prior_logvar)
        return kl.sum(dim=[1,2,3]).mean()

# ------------------- Main model -------------------
class NeuralODEFieldCollapser(nn.Module):
    def __init__(self, input_channels, latent_dim, hidden_dim=128, num_propagate=2):
        super().__init__()
        self.latent_dim = latent_dim
        self.num_propagate = num_propagate
        self.h_norm = nn.LayerNorm(latent_dim)
        
        # Encoder: [B,T,C,H,W] -> [B,C,1,H,W] -> squeeze -> [B,C,H,W]
        self.encoder = nn.Sequential(
            nn.Conv3d(input_channels, 32, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.AdaptiveAvgPool3d((1, None, None)),  # pool time axis to 1
            nn.Conv3d(32, latent_dim, kernel_size=1),
            nn.ReLU()
        )

        self.propagator = SpatialPropagator(latent_dim, hidden_dim)
        self.ode_func = ODEFunc(latent_dim, hidden_dim)

        self.decoder = nn.Sequential(
            nn.Conv2d(latent_dim, 64, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.Conv2d(64, 1, kernel_size=3, padding=1)
        )

        self.prior = PriorDynamics(latent_dim, (1,1))
        self.beta = 0.01

    def set_prior_from_data(self, dataset):
        prior_mean = torch.zeros(1, self.latent_dim, 1, 1)
        prior_logvar = torch.zeros(1, self.latent_dim, 1, 1)
        self.prior.set_prior(prior_mean, prior_logvar)

    def forward(self, x_past, t_span):
        """
        x_past: [B, T_in, C_in, H, W]
        t_span: 1D tensor of evaluation times (length T_out)
        Returns: predictions [B, T_out, 1, H, W], latent_states [T_out, B, C, H, W]
        """
        B, T_in, C_in, H, W = x_past.shape  # now exactly 5D
        # Convert to [B, C, T, H, W] for Conv3d
        x = x_past.permute(0, 2, 1, 3, 4)   # [B, C_in, T_in, H, W]
        h0 = self.encoder(x)                 # [B, latent_dim, 1, H, W]
        h0 = h0.squeeze(2)                   # [B, latent_dim, H, W]

        self.prior.grid_shape = (H, W)

        # Spatial propagation
        for _ in range(self.num_propagate):
            h0 = self.propagator(h0)

        h0 = self.h_norm(h0.permute(0,2,3,1)).permute(0,3,1,2)
        
        # Neural ODE
        # in NeuralODEFieldCollapser.forward
        h_t = odeint(self.ode_func, h0, t_span, method='rk4',
                     options={'step_size': float(t_span[1] - t_span[0])})

                     
        # Decode each time step
        preds = [self.decoder(h_t[i]).unsqueeze(1) for i in range(h_t.size(0))]
        predictions = torch.cat(preds, dim=1)  # [B, T_out, 1, H, W]
        return predictions, h_t

    def collapse_loss(self, h_t, predictions, targets):
        mse = F.mse_loss(predictions, targets)
        kl = sum(self.prior.kl_divergence(h_t[i]) for i in range(h_t.size(0))) / h_t.size(0)
        var = predictions.var(dim=1).mean()
        diff_h = (h_t[1:] - h_t[:-1]).pow(2).mean()
        loss = mse + self.beta * kl + 0.001 * var + 0.0001 * diff_h
        return loss, {'mse': mse.item(), 'kl': kl.item(), 'var': var.item(), 'diff': diff_h.item()}

# ------------------- Synthetic data (correct shape) -------------------
def generate_heat_equation_data(num_samples=500, T=20, dt=0.1, grid_size=16):
    """Returns [N, T, C=1, H, W] (5D)"""
    dx = 1.0 / (grid_size - 1)
    data = []
    for _ in range(num_samples):
        u = torch.rand(1, grid_size, grid_size) * 0.5   # [C=1, H, W]
        seq = [u.clone()]
        for _ in range(1, T):
            u_pad = F.pad(u, (1,1,1,1), mode='circular')
            lap = (u_pad[:, 2:, 1:-1] + u_pad[:, :-2, 1:-1] +
                   u_pad[:, 1:-1, 2:] + u_pad[:, 1:-1, :-2] - 4*u) / (dx**2)
            u = u + dt * (lap + 0.1 * u * (1-u)) + 0.01 * torch.randn_like(u)
            seq.append(u.clone())
        # stack time steps: [T, C=1, H, W] – no extra unsqueeze
        data.append(torch.stack(seq, dim=0))
    return torch.stack(data, dim=0)  # [N, T, 1, H, W]

# ------------------- Training -------------------
def train():
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    batch_size = 16
    latent_dim = 32
    T_past = 10
    T_future = 10
    grid_size = 16

    data = generate_heat_equation_data(500, T_past+T_future, grid_size=grid_size).to(device)
    train_data, val_data = data[:400], data[400:]

    model = NeuralODEFieldCollapser(input_channels=1, latent_dim=latent_dim).to(device)
    model.set_prior_from_data(None)
    optimizer = optim.Adam(model.parameters(), lr=1e-3)

    dt = 0.1
    t_span = torch.arange(0, T_future*dt, dt, device=device)

    for epoch in range(50):
        model.train()
        idx = torch.randperm(len(train_data))
        for start in range(0, len(train_data), batch_size):
            end = start + batch_size
            batch_idx = idx[start:end]
            x_past = train_data[batch_idx, :T_past]      # [B, T_past, 1, H, W]
            y_future = train_data[batch_idx, T_past:]    # [B, T_future, 1, H, W]

            optimizer.zero_grad()
            preds, h_t = model(x_past, t_span)
            loss, loss_dict = model.collapse_loss(h_t, preds, y_future)
            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)   # <<<
            optimizer.step()
            
        # Validation
        with torch.no_grad():
            x_val = val_data[:, :T_past]
            y_val = val_data[:, T_past:]
            pred_val, _ = model(x_val, t_span)
            val_mse = F.mse_loss(pred_val, y_val).item()
        print(f"Epoch {epoch:3d} | Train loss {loss.item():.4f} | Val MSE {val_mse:.4f}")

    return model

if __name__ == "__main__":
    model = train()
