import torch
import numpy as np
import zlib
import pickle
import gc
from transformers import AutoModelForCausalLM, AutoTokenizer

# -------------------------------
# Compression/decompression functions
# -------------------------------
def quantize_tensor(tensor, num_levels=256):
    # Convert to float32 for quantization to avoid FP16 issues
    tensor = tensor.float()
    min_val = tensor.min().item()
    max_val = tensor.max().item()
    if min_val == max_val:
        q = torch.zeros_like(tensor, dtype=torch.uint8)
        scale = 0.0
        return q, min_val, max_val, scale
    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

def compress_weights(model, num_levels=256):
    compressed_data = {}
    total_original_bytes = 0
    total_compressed_bytes = 0
    seen_tensors = set()
    for name, param in model.named_parameters():
        if param.requires_grad and id(param) not in seen_tensors:
            seen_tensors.add(id(param))
            q, min_val, max_val, scale = 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

def decompress_and_apply(model, compressed_data):
    applied_count = 0
    
    # Handle tied weights first - set lm_head to point to embed_tokens
    if hasattr(model, 'lm_head') and hasattr(model, 'model') and hasattr(model.model, 'embed_tokens'):
        if 'model.embed_tokens.weight' in compressed_data:
            with torch.no_grad():
                model.lm_head.weight.copy_(model.model.embed_tokens.weight)
            applied_count += 1
    
    # Decompress and apply other weights
    with torch.no_grad():
        for name, param in model.named_parameters():
            if name in compressed_data:
                info = compressed_data[name]
                q_bytes = zlib.decompress(info['compressed'])
                q_np = np.frombuffer(q_bytes, dtype=np.uint8).reshape(info['shape']).copy()
                q = torch.from_numpy(q_np)
                
                min_val = info['min_val']
                max_val = info['max_val']
                num_levels = info['num_levels']
                
                if num_levels > 1:
                    scale = (max_val - min_val) / (num_levels - 1)
                    dequantized = (min_val + scale * q).to(param.dtype)
                else:
                    dequantized = torch.full(q.shape, min_val, dtype=param.dtype)
                
                param.copy_(dequantized)
                applied_count += 1
                
    print(f"Applied {applied_count} decompressed weight tensors")
    return model

# -------------------------------
# 1. Load original Qwen3.5-0.8B LM model
# -------------------------------
model_name = "Qwen/Qwen3.5-0.8B"
print(f"Loading {model_name}...")
model = AutoModelForCausalLM.from_pretrained(model_name, torch_dtype=torch.float16, device_map="cpu")
model.eval()

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

# Save compressed dict to disk
compressed_file = "qwen35_08b_compressed.pkl"
with open(compressed_file, "wb") as f:
    pickle.dump(compressed, f)
print(f"Compressed model saved to {compressed_file}")

# -------------------------------
# 3. Load compressed weights from disk and decompress INTO SAME MODEL
# -------------------------------
print("\nLoading compressed weights from disk...")
with open(compressed_file, "rb") as f:
    compressed_loaded = pickle.load(f)

print("Decompressing and applying weights to model...")
decompress_and_apply(model, compressed_loaded)
model.eval()

# Clear compressed data from memory
del compressed_loaded, compressed
gc.collect()

# -------------------------------
# 4. Run compressed model FIRST
# -------------------------------
print("\nRunning compressed model...")
tokenizer = AutoTokenizer.from_pretrained(model_name)
tokenizer.pad_token = tokenizer.eos_token
input_text = "The future of artificial intelligence is"
input_ids = tokenizer.encode(input_text, return_tensors="pt")

with torch.no_grad():
    output_decomp = model(input_ids)
    logits_decomp = output_decomp.logits.clone()
    gen_decomp = model.generate(input_ids, max_length=40, do_sample=False)

print(f"Compressed output: {tokenizer.decode(gen_decomp[0], skip_special_tokens=True)}")

# Clear some memory
del output_decomp
gc.collect()

# -------------------------------
# 5. Reload original model to compare
# -------------------------------
print("\nReloading original model for comparison...")
del model
gc.collect()

model_orig = AutoModelForCausalLM.from_pretrained(model_name, torch_dtype=torch.float16, device_map="cpu")
model_orig.eval()

with torch.no_grad():
    output_orig = model_orig(input_ids)
    logits_orig = output_orig.logits
    gen_orig = model_orig.generate(input_ids, max_length=40, do_sample=False)

print(f"Original output:     {tokenizer.decode(gen_orig[0], skip_special_tokens=True)}")

# -------------------------------
# 6. Verify output similarity
# -------------------------------
logits_orig = logits_orig.float()
logits_decomp = logits_decomp.float()

mse = torch.nn.functional.mse_loss(logits_orig, logits_decomp).item()
cosine_sim = torch.nn.functional.cosine_similarity(logits_orig.flatten(), logits_decomp.flatten(), dim=0).item()
max_diff = (logits_orig - logits_decomp).abs().max().item()

print("\n--- Verification Results ---")
print(f"MSE between logits: {mse:.6e}")
print(f"Cosine similarity: {cosine_sim:.8f}")
print(f"Maximum absolute difference: {max_diff:.6f}")
