import torch
import torchvision.models as models
import numpy as np
import zlib
import sys
import os

# --- Your Compression Functions (from MNIST experiment) ---
def quantize_tensor(tensor, num_levels=256):
    # ... (Your implementation here) ...
    min_val = tensor.min().item()
    max_val = tensor.max().item()
    if min_val == max_val:
        # Handle constant tensor
        return torch.zeros_like(tensor, dtype=torch.uint8), min_val, max_val, 0.0, 0
    scale = (max_val - min_val) / (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, 0

def compress_weights(model, num_levels=256):
    # ... (Your implementation here) ...
    compressed_data = {}
    total_original_bytes = 0
    total_compressed_bytes = 0
    for name, param in model.named_parameters():
        if param.requires_grad:
            q, min_val, max_val, scale, zero_point = quantize_tensor(param.detach().cpu(), num_levels)
            q_np = q.numpy().tobytes()
            comp = zlib.compress(q_np, level=9)
            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()
            total_compressed_bytes += len(comp)
    ratio = total_compressed_bytes / total_original_bytes
    print(f"Compression ratio: {ratio:.4f}")
    return compressed_data, ratio

# --- 1. Load an Open-Weight Model (e.g., ResNet18) ---
print("Loading pre-trained ResNet18...")
model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)
model.eval()

# --- 2. Compress its weights ---
print("\nCompressing weights...")
compressed_data, ratio = compress_weights(model, num_levels=256)

# Optional: Save the compressed dictionary to disk for later use
import pickle
with open("resnet18_compressed.pkl", "wb") as f:
    pickle.dump(compressed_data, f)
print(f"Compressed model saved to 'resnet18_compressed.pkl'")

# Check actual disk usage of the saved file
disk_size = os.path.getsize("resnet18_compressed.pkl") / (1024 * 1024)
print(f"Compressed file on disk: {disk_size:.2f} MB")