import torch
import torch.nn as nn
import torch.optim as optim
import torchvision
import torchvision.transforms as transforms
import numpy as np
import zlib
import struct
from pathlib import Path
from typing import Dict, Tuple

# -------------------------------
# 1. Define MNIST CNN Model
# -------------------------------
class SimpleCNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 32, 3, 1)
        self.conv2 = nn.Conv2d(32, 64, 3, 1)
        self.dropout1 = nn.Dropout2d(0.25)
        self.dropout2 = nn.Dropout2d(0.5)
        self.fc1 = nn.Linear(9216, 128)
        self.fc2 = nn.Linear(128, 10)

    def forward(self, x):
        x = self.conv1(x)
        x = nn.functional.relu(x)
        x = self.conv2(x)
        x = nn.functional.relu(x)
        x = nn.functional.max_pool2d(x, 2)
        x = self.dropout1(x)
        x = torch.flatten(x, 1)
        x = self.fc1(x)
        x = nn.functional.relu(x)
        x = self.dropout2(x)
        x = self.fc2(x)
        return nn.functional.log_softmax(x, dim=1)

# -------------------------------
# 2. Triple Integration / Differentiation
# -------------------------------
def triple_cumsum(arr: np.ndarray) -> np.ndarray:
    """Apply cumulative sum three times along flattened view."""
    flat = arr.ravel().astype(np.float64)
    for _ in range(3):
        flat = np.cumsum(flat)
    return flat.reshape(arr.shape).astype(np.float64)

def triple_diff(arr: np.ndarray) -> np.ndarray:
    """Inverse of triple_cumsum: three successive differences."""
    flat = arr.ravel().astype(np.float64)
    for _ in range(3):
        flat = np.diff(flat, prepend=0)   # prepend 0 to restore length
    return flat.reshape(arr.shape).astype(np.float64)

def compress_tensor(tensor: torch.Tensor) -> bytes:
    """Apply triple cumsum, then zlib compress."""
    np_arr = tensor.detach().cpu().numpy()
    transformed = triple_cumsum(np_arr)
    return zlib.compress(transformed.tobytes(), level=9)

def decompress_tensor(compressed: bytes, original_shape: Tuple, dtype=np.float64) -> torch.Tensor:
    """Decompress and invert triple diff to restore original tensor."""
    raw = zlib.decompress(compressed)
    transformed = np.frombuffer(raw, dtype=dtype).reshape(original_shape)
    restored = triple_diff(transformed)
    return torch.from_numpy(restored.astype(np.float32))

def compression_ratio(original_size: int, compressed_size: int) -> float:
    return original_size / compressed_size

# -------------------------------
# 3. Training & Evaluation
# -------------------------------
def train(model, device, train_loader, optimizer, epoch):
    model.train()
    for batch_idx, (data, target) in enumerate(train_loader):
        data, target = data.to(device), target.to(device)
        optimizer.zero_grad()
        output = model(data)
        loss = nn.functional.nll_loss(output, target)
        loss.backward()
        optimizer.step()
        if batch_idx % 100 == 0:
            print(f'Train Epoch {epoch} [{batch_idx * len(data)}/{len(train_loader.dataset)} '
                  f'({100. * batch_idx / len(train_loader):.0f}%)]\tLoss: {loss.item():.6f}')

def test(model, device, test_loader):
    model.eval()
    test_loss = 0
    correct = 0
    with torch.no_grad():
        for data, target in test_loader:
            data, target = data.to(device), target.to(device)
            output = model(data)
            test_loss += nn.functional.nll_loss(output, target, reduction='sum').item()
            pred = output.argmax(dim=1, keepdim=True)
            correct += pred.eq(target.view_as(pred)).sum().item()
    test_loss /= len(test_loader.dataset)
    accuracy = 100. * correct / len(test_loader.dataset)
    print(f'Test set: Average loss: {test_loss:.4f}, Accuracy: {correct}/{len(test_loader.dataset)} ({accuracy:.2f}%)\n')
    return accuracy

# -------------------------------
# 4. Main: Train, Compress, Measure
# -------------------------------
def main():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"Using {device}")

    # Data loaders
    transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,))])
    train_set = torchvision.datasets.MNIST(root='../data', train=True, download=True, transform=transform)
    test_set = torchvision.datasets.MNIST(root='../data', train=False, download=True, transform=transform)
    train_loader = torch.utils.data.DataLoader(train_set, batch_size=64, shuffle=True)
    test_loader = torch.utils.data.DataLoader(test_set, batch_size=1000, shuffle=False)

    # Model, optimizer
    model = SimpleCNN().to(device)
    optimizer = optim.Adam(model.parameters(), lr=0.001)

    # Train
    for epoch in range(1, 4):  # 3 epochs for quick demo
        train(model, device, train_loader, optimizer, epoch)
        test(model, device, test_loader)

    # Save original model (standard .pth)
    torch.save(model.state_dict(), "mnist_model_original.pth")
    print("\n✅ Original model saved as 'mnist_model_original.pth'")

    # -------------------------------
    # 5. Compression metrics per parameter
    # -------------------------------
    total_original_bytes = 0
    total_compressed_raw_bytes = 0
    total_compressed_triple_bytes = 0
    param_stats = []

    for name, param in model.named_parameters():
        if param.requires_grad:
            # Original size (float32)
            orig_bytes = param.numel() * 4
            total_original_bytes += orig_bytes

            # Raw compression (baseline)
            raw_bytes = param.detach().cpu().numpy().astype(np.float32).tobytes()
            comp_raw = zlib.compress(raw_bytes, level=9)
            total_compressed_raw_bytes += len(comp_raw)

            # Triple-cumsum compression
            comp_triple = compress_tensor(param)
            total_compressed_triple_bytes += len(comp_triple)

            # Verify perfect reconstruction (lossless)
            restored = decompress_tensor(comp_triple, param.shape)
            max_err = (param.cpu() - restored).abs().max().item()
            param_stats.append({
                "name": name,
                "orig_bytes": orig_bytes,
                "raw_compressed": len(comp_raw),
                "triple_compressed": len(comp_triple),
                "reconstruction_error": max_err
            })

    # Print per‑parameter table
    print("\n📊 Compression Metrics per Parameter")
    print(f"{'Parameter':<25} {'Original (B)':<12} {'Raw zlib (B)':<12} {'Triple+zlib (B)':<15} {'Error':<10}")
    print("-" * 75)
    for s in param_stats:
        print(f"{s['name']:<25} {s['orig_bytes']:<12} {s['raw_compressed']:<12} {s['triple_compressed']:<15} {s['reconstruction_error']:.2e}")

    # Total ratios
    ratio_raw = total_original_bytes / total_compressed_raw_bytes
    ratio_triple = total_original_bytes / total_compressed_triple_bytes
    print("\n📈 Overall Compression Ratios (original / compressed)")
    print(f"Raw weights + zlib       : {ratio_raw:.2f}x  ({total_compressed_raw_bytes/total_original_bytes*100:.1f}% of original)")
    print(f"Triple‑cumsum + zlib     : {ratio_triple:.2f}x  ({total_compressed_triple_bytes/total_original_bytes*100:.1f}% of original)")
    print(f"Space saving vs raw zlib : {(1 - total_compressed_triple_bytes/total_compressed_raw_bytes)*100:.1f}%")

    # Optional: save the triple‑compressed state (as a dict of bytes)
    compressed_state = {}
    for name, param in model.named_parameters():
        if param.requires_grad:
            compressed_state[name] = compress_tensor(param)
    # You could save this dict as a .pth file (but it holds bytes, not tensors)
    # torch.save(compressed_state, "mnist_model_triple_compressed.pth")
    # print("\n💾 Triple‑compressed model state saved as 'mnist_model_triple_compressed.pth'")

if __name__ == "__main__":
    main()
