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
from torch.utils.data import DataLoader

# -------------------------------
# 1. Load MNIST data
# -------------------------------
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 = DataLoader(train_set, batch_size=64, shuffle=True)
test_loader  = DataLoader(test_set, batch_size=1000, shuffle=False)

# -------------------------------
# 2. Define a simple MLP model
# -------------------------------
class SimpleMLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(28*28, 256)
        self.fc2 = nn.Linear(256, 128)
        self.fc3 = nn.Linear(128, 10)
        self.relu = nn.ReLU()

    def forward(self, x):
        x = x.view(-1, 28*28)
        x = self.relu(self.fc1(x))
        x = self.relu(self.fc2(x))
        x = self.fc3(x)
        return x

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = SimpleMLP().to(device)

# -------------------------------
# 3. Train the model (full precision)
# -------------------------------
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)

def train(epochs=5):
    model.train()
    for epoch in range(epochs):
        for images, labels in train_loader:
            images, labels = images.to(device), labels.to(device)
            optimizer.zero_grad()
            outputs = model(images)
            loss = criterion(outputs, labels)
            loss.backward()
            optimizer.step()
        print(f'Epoch {epoch+1} completed')

def test():
    model.eval()
    correct = 0
    total = 0
    with torch.no_grad():
        for images, labels in test_loader:
            images, labels = images.to(device), labels.to(device)
            outputs = model(images)
            _, predicted = torch.max(outputs, 1)
            total += labels.size(0)
            correct += (predicted == labels).sum().item()
    acc = 100 * correct / total
    print(f'Test accuracy: {acc:.2f}%')
    return acc

print("Training full-precision model...")
train(epochs=5)
full_acc = test()

# -------------------------------
# 4. Weight compression functions
# -------------------------------
def quantize_tensor(tensor, num_levels=256):
    """Uniformly quantize a float tensor to num_levels integers."""
    # Compute min and max (per tensor)
    min_val = tensor.min().item()
    max_val = tensor.max().item()
    if min_val == max_val:
        # All values identical -> quantize to 0
        quantized = torch.zeros_like(tensor, dtype=torch.uint8)
        scale = 0.0
        zero_point = 0
        return quantized, min_val, max_val, scale, zero_point
    # Scale and zero-point for affine quantization: q = round((x - min) / scale) + zero_point
    scale = (max_val - min_val) / (num_levels - 1)
    zero_point = 0  # we keep zero_point = 0 for simplicity (symmetric? no, asymmetric but zero_point fixed)
    # Actually standard: q = clamp(round((x - min_val) / scale), 0, num_levels-1)
    q = torch.round((tensor - min_val) / scale).clamp(0, num_levels-1).to(torch.uint8)
    return q, min_val, max_val, scale, zero_point

def compress_weights(model, num_levels=256):
    """Compress all parameters of the model using quantization + zlib."""
    compressed_data = {}
    total_original_bytes = 0
    total_compressed_bytes = 0
    for name, param in model.named_parameters():
        if param.requires_grad:
            # Quantize
            q, min_val, max_val, scale, zero_point = quantize_tensor(param.detach().cpu(), num_levels)
            # Convert quantized tensor to bytes
            q_np = q.numpy().tobytes()
            # Lossless compress with zlib
            comp = zlib.compress(q_np, level=9)
            # Store metadata + compressed bytes
            compressed_data[name] = {
                'compressed': comp,
                'shape': param.shape,
                'min_val': min_val,
                'max_val': max_val,
                'num_levels': num_levels,
                'dtype': param.dtype
            }
            total_original_bytes += param.numel() * param.element_size()  # float32: 4 bytes each
            total_compressed_bytes += len(comp)
    ratio = total_compressed_bytes / total_original_bytes
    print(f"Compression ratio (compressed/original) = {ratio:.4f}")
    print(f"Original size: {total_original_bytes/1024:.2f} KB, Compressed size: {total_compressed_bytes/1024:.2f} KB")
    return compressed_data, ratio

def decompress_weights(compressed_data):
    """Decompress and reconstruct parameters as torch tensors."""
    state_dict = {}
    for name, info in compressed_data.items():
        # Decompress bytes
        q_bytes = zlib.decompress(info['compressed'])
        # Convert to numpy array
        q_np = np.frombuffer(q_bytes, dtype=np.uint8).reshape(info['shape'])
        q = torch.from_numpy(q_np).float()
        # Dequantize
        min_val = info['min_val']
        max_val = info['max_val']
        num_levels = info['num_levels']
        scale = (max_val - min_val) / (num_levels - 1) if num_levels > 1 else 0
        dequantized = min_val + scale * q
        state_dict[name] = dequantized.to(info['dtype'])
    return state_dict

# -------------------------------
# 5. Compress the trained model weights
# -------------------------------
print("\nCompressing weights (quantization to 256 levels + zlib)...")
compressed_data, ratio = compress_weights(model, num_levels=256)

# -------------------------------
# 6. Decompress and evaluate accuracy
# -------------------------------
# Create a new model and load decompressed weights
decompressed_state = decompress_weights(compressed_data)
model_quant = SimpleMLP().to(device)
model_quant.load_state_dict(decompressed_state, strict=True)

# Test accuracy after compression/decompression
print("\nEvaluating model after decompression (quantized weights):")
quant_acc = test()

# Optional: also test with only quantization (no zlib) – just for comparison
# (We already have that because decompress recovers the quantized values)

print(f"\nSummary:")
print(f"Full precision test accuracy: {full_acc:.2f}%")
print(f"Quantized (256 levels) test accuracy: {quant_acc:.2f}%")
print(f"Compression ratio (zlib on quantized indices): {ratio:.4f}")
