import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
import numpy as np

# ==========================================
# 1. RANDOM PROJECTION ENCODER
# ==========================================
class RandomProjectionEncoder(nn.Module):
    def __init__(self, in_features, out_features):
        super().__init__()
        # Using a better distribution for the projection matrix
        proj_matrix = torch.randn(in_features, out_features) * (1.0 / np.sqrt(in_features))
        self.register_buffer('proj_matrix', proj_matrix)

    def forward(self, x):
        return torch.matmul(x, self.proj_matrix)

# ==========================================
# 2. CAUSAL-RESONANCE DOT PRODUCT LAYER
# ==========================================
class CRDPLinear(nn.Module):
    def __init__(self, in_features, out_features, memory_size=5):
        super().__init__()
        self.in_features = in_features
        self.out_features = out_features
        self.memory_size = memory_size

        # Kaiming-style initialization to prevent signal death
        self.weight = nn.Parameter(torch.randn(out_features, in_features) * np.sqrt(2.0 / in_features))
        self.bias = nn.Parameter(torch.zeros(out_features))
        
        # Memory kernel should be normalized
        self.memory_kernel = nn.Parameter(torch.randn(1, memory_size) * 0.1)
        self.resonance = nn.Parameter(torch.zeros(out_features))
        self.trajectory_buffer = None

    def forward(self, x):
        batch_size = x.size(0)
        current_proj = torch.matmul(x, self.weight.t()) + self.bias
        
        if self.trajectory_buffer is None or self.trajectory_buffer.size(0) != batch_size:
            self.trajectory_buffer = torch.zeros(batch_size, self.memory_size, self.out_features, device=x.device)
        
        # Detach to prevent infinite BPTT graphs
        self.trajectory_buffer = torch.cat([self.trajectory_buffer[:, 1:, :].detach(), current_proj.unsqueeze(1)], dim=1)
        
        kernel_expanded = self.memory_kernel.unsqueeze(0).expand(batch_size, -1, -1)
        memory_contribution = torch.bmm(kernel_expanded, self.trajectory_buffer).squeeze(1)

        return current_proj + memory_contribution + self.resonance

# ==========================================
# 3. T2C SUPERVISOR
# ==========================================
class T2CSupervisor:
    def __init__(self, state_dim):
        self.state_dim = state_dim
        self.correction_map = {} 

    def build_map(self, model, loader):
        model.eval()
        print("Building T2C Correction Map from errors...")
        with torch.no_grad():
            for data, target in loader:
                x = data.view(-1, 784)
                x = model.encoder(x)
                x = torch.relu(model.layer1(x))
                x = torch.relu(model.layer2(x))
                
                logits = model.layer3(x)
                preds = logits.argmax(dim=1)

                for i in range(data.size(0)):
                    if preds[i] != target[i]:
                        wrong_state = x[i].cpu().numpy() 
                        target_vec = model.layer3.weight[target[i]].cpu().numpy()
                        correction = target_vec - wrong_state
                        # Lower precision for the hash to increase "resonance" matches
                        state_hash = tuple(np.round(wrong_state, 0)) 
                        self.correction_map[state_hash] = torch.from_numpy(correction).float()

    def apply_correction(self, x):
        corrected_x = x.clone()
        for i in range(x.size(0)):
            state_hash = tuple(np.round(x[i].cpu().numpy(), 0))
            if state_hash in self.correction_map:
                corrected_x[i] += self.correction_map[state_hash].to(x.device)
        return corrected_x

# ==========================================
# 4. INTEGRATED MODEL
# ==========================================
class T2C_CCT_Model(nn.Module):
    def __init__(self):
        super().__init__()
        self.encoder = RandomProjectionEncoder(784, 128) # Increased to 128 for better signal
        self.layer1 = CRDPLinear(128, 64)
        self.layer2 = CRDPLinear(64, 32)
        self.layer3 = CRDPLinear(32, 10)
        self.relu = nn.ReLU()

    def forward(self, x, supervisor=None):
        x = x.view(-1, 784)
        x = self.encoder(x)
        x = self.relu(self.layer1(x))
        x = self.relu(self.layer2(x))
        if supervisor is not None:
            x = supervisor.apply_correction(x)
        x = self.layer3(x)
        return x

# ==========================================
# 5. RUN EXPERIMENT
# ==========================================
def run_experiment():
    # Higher learning rate to break the 2.30 loss plateau
    batch_size, lr, epochs = 64, 0.005, 10 
    transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,))])
    train_loader = DataLoader(datasets.MNIST('../data', train=True, download=True, transform=transform), batch_size=batch_size, shuffle=True)
    test_loader = DataLoader(datasets.MNIST('../data', train=False, transform=transform), batch_size=batch_size, shuffle=False)

    model = T2C_CCT_Model()
    criterion = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.parameters(), lr=lr)

    model.train()
    print("Training with Resonance...")
    for epoch in range(epochs):
        total_loss = 0
        for data, target in train_loader:
            optimizer.zero_grad()
            output = model(data)
            loss = criterion(output, target)
            loss.backward()
            optimizer.step()
            total_loss += loss.item()
        print(f"Epoch {epoch+1}, Loss: {total_loss/len(train_loader):.4f}")

    supervisor = T2CSupervisor(state_dim=32)
    supervisor.build_map(model, train_loader)

    model.eval()
    correct = 0
    with torch.no_grad():
        for data, target in test_loader:
            output = model(data, supervisor=supervisor)
            pred = output.argmax(dim=1, keepdim=True)
            correct += pred.eq(target.view_as(pred)).sum().item()

    print(f"\nFinal Test Accuracy with T2C Correction: {100. * correct / len(test_loader.dataset):.2f}%")

if __name__ == "__main__":
    run_experiment()
