import torch
import torch.nn as nn
import numpy as np
from sklearn.metrics import mutual_info_score
from scipy.optimize import curve_fit
import matplotlib.pyplot as plt
import hashlib

# ---------- 2D Local Attention (Convolutional) ----------
class LocalAttention2D(nn.Module):
    def __init__(self, d_model, kernel_size=3):
        super().__init__()
        self.kernel_size = kernel_size
        self.conv = nn.Conv2d(d_model, d_model, kernel_size, padding=kernel_size//2, groups=d_model)
        # Simplified: each channel attends locally via depthwise conv
        
    def forward(self, h):
        # h: (batch, seq_len, d_model); we reshape to 2D: (batch, d_model, height, width)
        batch, seq, d = h.shape
        size = int(np.sqrt(seq))
        assert size*size == seq, "seq_len must be perfect square"
        h_2d = h.permute(0,2,1).reshape(batch, d, size, size)
        # Local interaction via convolution
        h_local = self.conv(h_2d)
        # No softmax over whole sequence – each position's update is local
        h_out = h_local.reshape(batch, d, seq).permute(0,2,1)
        return h_out, None  # no attention weights needed

# Modified STT with local 2D attention
class LocalSTT(nn.Module):
    def __init__(self, d_model=64, seq_len=16, temp=1.0, kernel_size=3):
        super().__init__()
        self.d_model = d_model
        self.seq_len = seq_len
        self.token_embed = nn.Parameter(torch.randn(1, seq_len, d_model) * 0.02)
        self.pos_embed = nn.Parameter(torch.randn(1, seq_len, d_model) * 0.02)
        # Use several local attention layers
        self.local_attns = nn.ModuleList([LocalAttention2D(d_model, kernel_size) for _ in range(4)])
        # Collapse layers (same as before)
        xi_hash = hashlib.md5(b"pi_anchor:e_anchor").hexdigest()
        xi_tensor = self._hash_to_tensor(xi_hash, d_model)
        self.collapse_layers = nn.ModuleList([
            CollapseSelfTunable(d_model, xi_tensor, threshold=0.5, temp=temp)
            for _ in range(4)
        ])
        
    def _hash_to_tensor(self, hash_str, d_model):
        bytes_data = bytes.fromhex(hash_str)[:d_model]
        tensor = torch.tensor([b / 255.0 for b in bytes_data], dtype=torch.float)
        return tensor / tensor.norm()
    
    def forward(self, steps=20, return_history=False):
        h = self.token_embed + self.pos_embed
        history = []
        for _ in range(steps):
            for attn, collapse in zip(self.local_attns, self.collapse_layers):
                h, _ = attn(h)
                h, _ = collapse(h)
            history.append(h.detach().clone())
        if return_history:
            return torch.stack(history, dim=0)
        return h
    
    def compute_phi(self, steps=20):
        H = self.forward(steps=steps, return_history=True).squeeze(1)  # (steps, seq_len, d_model)
        half = self.d_model // 2
        X = H[:, :, :half].reshape(-1, half * H.shape[1]).numpy()
        Y = H[:, :, half:].reshape(-1, half * H.shape[1]).numpy()
        Xd = (X > 0).astype(int)
        Yd = (Y > 0).astype(int)
        I_whole = mutual_info_score(
            [f"{x}{y}" for x,y in zip(Xd[:-1].flatten(), Yd[:-1].flatten())],
            [f"{x}{y}" for x,y in zip(Xd[1:].flatten(), Yd[1:].flatten())]
        )
        I_X = mutual_info_score(Xd[:-1].flatten(), Xd[1:].flatten())
        I_Y = mutual_info_score(Yd[:-1].flatten(), Yd[1:].flatten())
        return max(0.0, I_whole - (I_X + I_Y))

# Sweep temperature for local STT
temperatures = np.linspace(0.2, 2.0, 15)
phi_local = []
Tc_est = 0.9  # rough estimate (we'll refine)

for T in temperatures:
    model = LocalSTT(d_model=64, seq_len=16, temp=T, kernel_size=3)
    with torch.no_grad():
        phi = model.compute_phi(steps=25)
        phi_local.append(phi)

# Fit power law below estimated Tc (~0.85-0.9)
T_below = temperatures[temperatures < 0.9]
phi_below = phi_local[:len(T_below)]
def phi_power(T, A, beta, Tc):
    return A * (Tc - T)**beta
popt, _ = curve_fit(phi_power, T_below, phi_below, p0=[0.5, 0.125, 0.9])
A_fit, beta_fit, Tc_fit = popt
print(f"Local STT: Estimated Tc = {Tc_fit:.3f}, β = {beta_fit:.3f}")
